Enrich Dex logs with real IP and request ID (#3661)

Signed-off-by: m.nabokikh <maksim.nabokikh@flant.com>
Signed-off-by: Maksim Nabokikh <max.nabokih@gmail.com>
Co-authored-by: Márk Sági-Kazár <sagikazarmark@users.noreply.github.com>
This commit is contained in:
Maksim Nabokikh
2024-08-01 21:37:35 +02:00
committed by GitHub
co-authored by Márk Sági-Kazár
parent 6ceb26509b
commit 225660785c
11 changed files with 333 additions and 187 deletions
+33 -9
View File
@@ -6,6 +6,7 @@ import (
"fmt"
"log/slog"
"net/http"
"net/netip"
"os"
"slices"
"strings"
@@ -182,15 +183,38 @@ type OAuth2 struct {
// Web is the config format for the HTTP server.
type Web struct {
HTTP string `json:"http"`
HTTPS string `json:"https"`
Headers Headers `json:"headers"`
TLSCert string `json:"tlsCert"`
TLSKey string `json:"tlsKey"`
TLSMinVersion string `json:"tlsMinVersion"`
TLSMaxVersion string `json:"tlsMaxVersion"`
AllowedOrigins []string `json:"allowedOrigins"`
AllowedHeaders []string `json:"allowedHeaders"`
HTTP string `json:"http"`
HTTPS string `json:"https"`
Headers Headers `json:"headers"`
TLSCert string `json:"tlsCert"`
TLSKey string `json:"tlsKey"`
TLSMinVersion string `json:"tlsMinVersion"`
TLSMaxVersion string `json:"tlsMaxVersion"`
AllowedOrigins []string `json:"allowedOrigins"`
AllowedHeaders []string `json:"allowedHeaders"`
ClientRemoteIP ClientRemoteIP `json:"clientRemoteIP"`
}
type ClientRemoteIP struct {
Header string `json:"header"`
TrustedProxies []string `json:"trustedProxies"`
}
func (cr *ClientRemoteIP) ParseTrustedProxies() ([]netip.Prefix, error) {
if cr == nil {
return nil, nil
}
trusted := make([]netip.Prefix, 0, len(cr.TrustedProxies))
for _, cidr := range cr.TrustedProxies {
ipNet, err := netip.ParsePrefix(cidr)
if err != nil {
return nil, fmt.Errorf("failed to parse CIDR %q: %v", cidr, err)
}
trusted = append(trusted, ipNet)
}
return trusted, nil
}
type Headers struct {
+67
View File
@@ -0,0 +1,67 @@
package main
import (
"context"
"fmt"
"log/slog"
"os"
"strings"
"github.com/dexidp/dex/server"
)
var logFormats = []string{"json", "text"}
func newLogger(level slog.Level, format string) (*slog.Logger, error) {
var handler slog.Handler
switch strings.ToLower(format) {
case "", "text":
handler = slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{
Level: level,
})
case "json":
handler = slog.NewJSONHandler(os.Stderr, &slog.HandlerOptions{
Level: level,
})
default:
return nil, fmt.Errorf("log format is not one of the supported values (%s): %s", strings.Join(logFormats, ", "), format)
}
return slog.New(newRequestContextHandler(handler)), nil
}
var _ slog.Handler = requestContextHandler{}
type requestContextHandler struct {
handler slog.Handler
}
func newRequestContextHandler(handler slog.Handler) slog.Handler {
return requestContextHandler{
handler: handler,
}
}
func (h requestContextHandler) Enabled(ctx context.Context, level slog.Level) bool {
return h.handler.Enabled(ctx, level)
}
func (h requestContextHandler) Handle(ctx context.Context, record slog.Record) error {
if v, ok := ctx.Value(server.RequestKeyRemoteIP).(string); ok {
record.AddAttrs(slog.String(string(server.RequestKeyRemoteIP), v))
}
if v, ok := ctx.Value(server.RequestKeyRequestID).(string); ok {
record.AddAttrs(slog.String(string(server.RequestKeyRequestID), v))
}
return h.handler.Handle(ctx, record)
}
func (h requestContextHandler) WithAttrs(attrs []slog.Attr) slog.Handler {
return requestContextHandler{h.handler.WithAttrs(attrs)}
}
func (h requestContextHandler) WithGroup(name string) slog.Handler {
return h.handler.WithGroup(name)
}
+7 -20
View File
@@ -348,6 +348,13 @@ func runServe(options serveOptions) error {
}
serverConfig.RefreshTokenPolicy = refreshTokenPolicy
serverConfig.RealIPHeader = c.Web.ClientRemoteIP.Header
serverConfig.TrustedRealIPCIDRs, err = c.Web.ClientRemoteIP.ParseTrustedProxies()
if err != nil {
return fmt.Errorf("failed to parse client remote IP settings: %v", err)
}
serv, err := server.NewServer(context.Background(), serverConfig)
if err != nil {
return fmt.Errorf("failed to initialize server: %v", err)
@@ -528,26 +535,6 @@ func runServe(options serveOptions) error {
return nil
}
var logFormats = []string{"json", "text"}
func newLogger(level slog.Level, format string) (*slog.Logger, error) {
var handler slog.Handler
switch strings.ToLower(format) {
case "", "text":
handler = slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{
Level: level,
})
case "json":
handler = slog.NewJSONHandler(os.Stderr, &slog.HandlerOptions{
Level: level,
})
default:
return nil, fmt.Errorf("log format is not one of the supported values (%s): %s", strings.Join(logFormats, ", "), format)
}
return slog.New(handler), nil
}
func applyConfigOverrides(options serveOptions, config *Config) {
if options.webHTTPAddr != "" {
config.Web.HTTP = options.webHTTPAddr
+4 -1
View File
@@ -58,7 +58,10 @@ web:
# X-XSS-Protection: "1; mode=block"
# Content-Security-Policy: "default-src 'self'"
# Strict-Transport-Security: "max-age=31536000; includeSubDomains"
# clientRemoteIP:
# header: X-Forwarded-For
# trustedProxies:
# - 10.0.0.0/8
# Configuration for dex appearance
# frontend:
+19 -19
View File
@@ -48,7 +48,7 @@ func (s *Server) handleDeviceExchange(w http.ResponseWriter, r *http.Request) {
invalidAttempt = false
}
if err := s.templates.device(r, w, s.getDeviceVerificationURI(), userCode, invalidAttempt); err != nil {
s.logger.Error("server template error", "err", err)
s.logger.ErrorContext(r.Context(), "server template error", "err", err)
s.renderError(r, w, http.StatusNotFound, "Page not found")
}
default:
@@ -64,7 +64,7 @@ func (s *Server) handleDeviceCode(w http.ResponseWriter, r *http.Request) {
case http.MethodPost:
err := r.ParseForm()
if err != nil {
s.logger.Error("could not parse Device Request body", "err", err)
s.logger.ErrorContext(r.Context(), "could not parse Device Request body", "err", err)
s.tokenErrHelper(w, errInvalidRequest, "", http.StatusNotFound)
return
}
@@ -85,7 +85,7 @@ func (s *Server) handleDeviceCode(w http.ResponseWriter, r *http.Request) {
return
}
s.logger.Info("received device request", "client_id", clientID, "scoped", scopes)
s.logger.InfoContext(r.Context(), "received device request", "client_id", clientID, "scoped", scopes)
// Make device code
deviceCode := storage.NewDeviceCode()
@@ -107,7 +107,7 @@ func (s *Server) handleDeviceCode(w http.ResponseWriter, r *http.Request) {
}
if err := s.storage.CreateDeviceRequest(ctx, deviceReq); err != nil {
s.logger.Error("failed to store device request", "err", err)
s.logger.ErrorContext(r.Context(), "failed to store device request", "err", err)
s.tokenErrHelper(w, errInvalidRequest, "", http.StatusInternalServerError)
return
}
@@ -126,14 +126,14 @@ func (s *Server) handleDeviceCode(w http.ResponseWriter, r *http.Request) {
}
if err := s.storage.CreateDeviceToken(ctx, deviceToken); err != nil {
s.logger.Error("failed to store device token", "err", err)
s.logger.ErrorContext(r.Context(), "failed to store device token", "err", err)
s.tokenErrHelper(w, errInvalidRequest, "", http.StatusInternalServerError)
return
}
u, err := url.Parse(s.issuerURL.String())
if err != nil {
s.logger.Error("could not parse issuer URL", "err", err)
s.logger.ErrorContext(r.Context(), "could not parse issuer URL", "err", err)
s.tokenErrHelper(w, errInvalidRequest, "", http.StatusInternalServerError)
return
}
@@ -211,7 +211,7 @@ func (s *Server) handleDeviceToken(w http.ResponseWriter, r *http.Request) {
deviceToken, err := s.storage.GetDeviceToken(deviceCode)
if err != nil {
if err != storage.ErrNotFound {
s.logger.Error("failed to get device code", "err", err)
s.logger.ErrorContext(r.Context(), "failed to get device code", "err", err)
}
s.tokenErrHelper(w, errInvalidRequest, "Invalid Device code.", http.StatusBadRequest)
return
@@ -241,7 +241,7 @@ func (s *Server) handleDeviceToken(w http.ResponseWriter, r *http.Request) {
}
// Update device token last request time in storage
if err := s.storage.UpdateDeviceToken(deviceCode, updater); err != nil {
s.logger.Error("failed to update device token", "err", err)
s.logger.ErrorContext(r.Context(), "failed to update device token", "err", err)
s.renderError(r, w, http.StatusInternalServerError, "")
return
}
@@ -258,7 +258,7 @@ func (s *Server) handleDeviceToken(w http.ResponseWriter, r *http.Request) {
case providedCodeVerifier != "" && codeChallengeFromStorage != "":
calculatedCodeChallenge, err := s.calculateCodeChallenge(providedCodeVerifier, deviceToken.PKCE.CodeChallengeMethod)
if err != nil {
s.logger.Error("failed to calculate code challenge", "err", err)
s.logger.ErrorContext(r.Context(), "failed to calculate code challenge", "err", err)
s.tokenErrHelper(w, errServerError, "", http.StatusInternalServerError)
return
}
@@ -303,7 +303,7 @@ func (s *Server) handleDeviceCallback(w http.ResponseWriter, r *http.Request) {
if err != nil || s.now().After(authCode.Expiry) {
errCode := http.StatusBadRequest
if err != nil && err != storage.ErrNotFound {
s.logger.Error("failed to get auth code", "err", err)
s.logger.ErrorContext(r.Context(), "failed to get auth code", "err", err)
errCode = http.StatusInternalServerError
}
s.renderError(r, w, errCode, "Invalid or expired auth code.")
@@ -315,7 +315,7 @@ func (s *Server) handleDeviceCallback(w http.ResponseWriter, r *http.Request) {
if err != nil || s.now().After(deviceReq.Expiry) {
errCode := http.StatusBadRequest
if err != nil && err != storage.ErrNotFound {
s.logger.Error("failed to get device code", "err", err)
s.logger.ErrorContext(r.Context(), "failed to get device code", "err", err)
errCode = http.StatusInternalServerError
}
s.renderError(r, w, errCode, "Invalid or expired user code.")
@@ -325,7 +325,7 @@ func (s *Server) handleDeviceCallback(w http.ResponseWriter, r *http.Request) {
client, err := s.storage.GetClient(deviceReq.ClientID)
if err != nil {
if err != storage.ErrNotFound {
s.logger.Error("failed to get client", "err", err)
s.logger.ErrorContext(r.Context(), "failed to get client", "err", err)
s.tokenErrHelper(w, errServerError, "", http.StatusInternalServerError)
} else {
s.tokenErrHelper(w, errInvalidClient, "Invalid client credentials.", http.StatusUnauthorized)
@@ -339,7 +339,7 @@ func (s *Server) handleDeviceCallback(w http.ResponseWriter, r *http.Request) {
resp, err := s.exchangeAuthCode(ctx, w, authCode, client)
if err != nil {
s.logger.Error("could not exchange auth code for clien", "client_id", deviceReq.ClientID, "err", err)
s.logger.ErrorContext(r.Context(), "could not exchange auth code for clien", "client_id", deviceReq.ClientID, "err", err)
s.renderError(r, w, http.StatusInternalServerError, "Failed to exchange auth code.")
return
}
@@ -349,7 +349,7 @@ func (s *Server) handleDeviceCallback(w http.ResponseWriter, r *http.Request) {
if err != nil || s.now().After(old.Expiry) {
errCode := http.StatusBadRequest
if err != nil && err != storage.ErrNotFound {
s.logger.Error("failed to get device token", "err", err)
s.logger.ErrorContext(r.Context(), "failed to get device token", "err", err)
errCode = http.StatusInternalServerError
}
s.renderError(r, w, errCode, "Invalid or expired device code.")
@@ -362,7 +362,7 @@ func (s *Server) handleDeviceCallback(w http.ResponseWriter, r *http.Request) {
}
respStr, err := json.MarshalIndent(resp, "", " ")
if err != nil {
s.logger.Error("failed to marshal device token response", "err", err)
s.logger.ErrorContext(r.Context(), "failed to marshal device token response", "err", err)
s.renderError(r, w, http.StatusInternalServerError, "")
return old, err
}
@@ -374,13 +374,13 @@ func (s *Server) handleDeviceCallback(w http.ResponseWriter, r *http.Request) {
// Update refresh token in the storage, store the token and mark as complete
if err := s.storage.UpdateDeviceToken(deviceReq.DeviceCode, updater); err != nil {
s.logger.Error("failed to update device token", "err", err)
s.logger.ErrorContext(r.Context(), "failed to update device token", "err", err)
s.renderError(r, w, http.StatusBadRequest, "")
return
}
if err := s.templates.deviceSuccess(r, w, client.Name); err != nil {
s.logger.Error("Server template error", "err", err)
s.logger.ErrorContext(r.Context(), "Server template error", "err", err)
s.renderError(r, w, http.StatusNotFound, "Page not found")
}
@@ -412,10 +412,10 @@ func (s *Server) verifyUserCode(w http.ResponseWriter, r *http.Request) {
deviceRequest, err := s.storage.GetDeviceRequest(userCode)
if err != nil || s.now().After(deviceRequest.Expiry) {
if err != nil && err != storage.ErrNotFound {
s.logger.Error("failed to get device request", "err", err)
s.logger.ErrorContext(r.Context(), "failed to get device request", "err", err)
}
if err := s.templates.device(r, w, s.getDeviceVerificationURI(), userCode, true); err != nil {
s.logger.Error("Server template error", "err", err)
s.logger.ErrorContext(r.Context(), "Server template error", "err", err)
s.renderError(r, w, http.StatusNotFound, "Page not found")
}
return
+95 -94
View File
File diff suppressed because it is too large Load Diff
+11 -10
View File
@@ -179,7 +179,7 @@ func (s *Server) getTokenFromRequest(r *http.Request) (string, TokenTypeEnum, er
token := r.PostForm.Get("token")
tokenType, err := s.guessTokenType(r.Context(), token)
if err != nil {
s.logger.Error("failed to guess token type", "err", err)
s.logger.ErrorContext(r.Context(), "failed to guess token type", "err", err)
return "", 0, newIntrospectInternalServerError()
}
@@ -193,7 +193,7 @@ func (s *Server) getTokenFromRequest(r *http.Request) (string, TokenTypeEnum, er
return token, tokenType, nil
}
func (s *Server) introspectRefreshToken(_ context.Context, token string) (*Introspection, error) {
func (s *Server) introspectRefreshToken(ctx context.Context, token string) (*Introspection, error) {
rToken := new(internal.RefreshToken)
if err := internal.Unmarshal(token, rToken); err != nil {
// For backward compatibility, assume the refresh_token is a raw refresh token ID
@@ -205,19 +205,19 @@ func (s *Server) introspectRefreshToken(_ context.Context, token string) (*Intro
rToken = &internal.RefreshToken{RefreshId: token, Token: ""}
}
rCtx, err := s.getRefreshTokenFromStorage(nil, rToken)
rCtx, err := s.getRefreshTokenFromStorage(ctx, nil, rToken)
if err != nil {
if errors.Is(err, invalidErr) || errors.Is(err, expiredErr) {
return nil, newIntrospectInactiveTokenError()
}
s.logger.Error("failed to get refresh token", "err", err)
s.logger.ErrorContext(ctx, "failed to get refresh token", "err", err)
return nil, newIntrospectInternalServerError()
}
subjectString, sErr := genSubject(rCtx.storageToken.Claims.UserID, rCtx.storageToken.ConnectorID)
if sErr != nil {
s.logger.Error("failed to marshal offline session ID", "err", err)
s.logger.ErrorContext(ctx, "failed to marshal offline session ID", "err", err)
return nil, newIntrospectInternalServerError()
}
@@ -253,19 +253,19 @@ func (s *Server) introspectAccessToken(ctx context.Context, token string) (*Intr
var claims IntrospectionExtra
if err := idToken.Claims(&claims); err != nil {
s.logger.Error("error while fetching token claims", "err", err.Error())
s.logger.ErrorContext(ctx, "error while fetching token claims", "err", err.Error())
return nil, newIntrospectInternalServerError()
}
clientID, err := getClientID(idToken.Audience, claims.AuthorizingParty)
if err != nil {
s.logger.Error("error while fetching client_id from token:", "err", err.Error())
s.logger.ErrorContext(ctx, "error while fetching client_id from token:", "err", err.Error())
return nil, newIntrospectInternalServerError()
}
client, err := s.storage.GetClient(clientID)
if err != nil {
s.logger.Error("error while fetching client from storage", "err", err.Error())
s.logger.ErrorContext(ctx, "error while fetching client from storage", "err", err.Error())
return nil, newIntrospectInternalServerError()
}
@@ -299,7 +299,7 @@ func (s *Server) handleIntrospect(w http.ResponseWriter, r *http.Request) {
introspect, err = s.introspectRefreshToken(ctx, token)
default:
// Token type is neither handled token types.
s.logger.Error("unknown token type", "token_type", tokenType)
s.logger.ErrorContext(r.Context(), "unknown token type", "token_type", tokenType)
introspectInactiveErr(w)
return
}
@@ -309,7 +309,7 @@ func (s *Server) handleIntrospect(w http.ResponseWriter, r *http.Request) {
if intErr, ok := err.(*introspectionError); ok {
s.introspectErrHelper(w, intErr.typ, intErr.desc, intErr.code)
} else {
s.logger.Error("an unknown error occurred", "err", err.Error())
s.logger.ErrorContext(r.Context(), "an unknown error occurred", "err", err.Error())
s.introspectErrHelper(w, errServerError, "An unknown error occurred", http.StatusInternalServerError)
}
@@ -332,6 +332,7 @@ func (s *Server) introspectErrHelper(w http.ResponseWriter, typ string, descript
}
if err := tokenErr(w, typ, description, statusCode); err != nil {
// TODO(nabokihms): error with context
s.logger.Error("introspect error response", "err", err)
}
}
+1 -1
View File
@@ -259,7 +259,7 @@ func TestHandleIntrospect(t *testing.T) {
mockTestStorage(t, s.storage)
activeAccessToken, expiry, err := s.newIDToken("test", storage.Claims{
activeAccessToken, expiry, err := s.newIDToken(ctx, "test", storage.Claims{
UserID: "1",
Username: "jane",
Email: "jane.doe@example.com",
+13 -13
View File
@@ -303,8 +303,8 @@ type federatedIDClaims struct {
UserID string `json:"user_id,omitempty"`
}
func (s *Server) newAccessToken(clientID string, claims storage.Claims, scopes []string, nonce, connID string) (accessToken string, expiry time.Time, err error) {
return s.newIDToken(clientID, claims, scopes, nonce, storage.NewID(), "", connID)
func (s *Server) newAccessToken(ctx context.Context, clientID string, claims storage.Claims, scopes []string, nonce, connID string) (accessToken string, expiry time.Time, err error) {
return s.newIDToken(ctx, clientID, claims, scopes, nonce, storage.NewID(), "", connID)
}
func getClientID(aud audience, azp string) (string, error) {
@@ -350,10 +350,10 @@ func genSubject(userID string, connID string) (string, error) {
return internal.Marshal(sub)
}
func (s *Server) newIDToken(clientID string, claims storage.Claims, scopes []string, nonce, accessToken, code, connID string) (idToken string, expiry time.Time, err error) {
func (s *Server) newIDToken(ctx context.Context, clientID string, claims storage.Claims, scopes []string, nonce, accessToken, code, connID string) (idToken string, expiry time.Time, err error) {
keys, err := s.storage.GetKeys()
if err != nil {
s.logger.Error("failed to get keys", "err", err)
s.logger.ErrorContext(ctx, "failed to get keys", "err", err)
return "", expiry, err
}
@@ -371,7 +371,7 @@ func (s *Server) newIDToken(clientID string, claims storage.Claims, scopes []str
subjectString, err := genSubject(claims.UserID, connID)
if err != nil {
s.logger.Error("failed to marshal offline session ID", "err", err)
s.logger.ErrorContext(ctx, "failed to marshal offline session ID", "err", err)
return "", expiry, fmt.Errorf("failed to marshal offline session ID: %v", err)
}
@@ -386,7 +386,7 @@ func (s *Server) newIDToken(clientID string, claims storage.Claims, scopes []str
if accessToken != "" {
atHash, err := accessTokenHash(signingAlg, accessToken)
if err != nil {
s.logger.Error("error computing at_hash", "err", err)
s.logger.ErrorContext(ctx, "error computing at_hash", "err", err)
return "", expiry, fmt.Errorf("error computing at_hash: %v", err)
}
tok.AccessTokenHash = atHash
@@ -395,7 +395,7 @@ func (s *Server) newIDToken(clientID string, claims storage.Claims, scopes []str
if code != "" {
cHash, err := accessTokenHash(signingAlg, code)
if err != nil {
s.logger.Error("error computing c_hash", "err", err)
s.logger.ErrorContext(ctx, "error computing c_hash", "err", err)
return "", expiry, fmt.Errorf("error computing c_hash: #{err}")
}
tok.CodeHash = cHash
@@ -423,7 +423,7 @@ func (s *Server) newIDToken(clientID string, claims storage.Claims, scopes []str
// initial auth request.
continue
}
isTrusted, err := s.validateCrossClientTrust(clientID, peerID)
isTrusted, err := s.validateCrossClientTrust(ctx, clientID, peerID)
if err != nil {
return "", expiry, err
}
@@ -482,7 +482,7 @@ func (s *Server) parseAuthorizationRequest(r *http.Request) (*storage.AuthReques
if err == storage.ErrNotFound {
return nil, newDisplayedErr(http.StatusNotFound, "Invalid client_id (%q).", clientID)
}
s.logger.Error("failed to get client", "err", err)
s.logger.ErrorContext(r.Context(), "failed to get client", "err", err)
return nil, newDisplayedErr(http.StatusInternalServerError, "Database error.")
}
@@ -501,7 +501,7 @@ func (s *Server) parseAuthorizationRequest(r *http.Request) (*storage.AuthReques
if connectorID != "" {
connectors, err := s.storage.ListConnectors()
if err != nil {
s.logger.Error("failed to list connectors", "err", err)
s.logger.ErrorContext(r.Context(), "failed to list connectors", "err", err)
return nil, newRedirectedErr(errServerError, "Unable to retrieve connectors")
}
if !validateConnectorID(connectors, connectorID) {
@@ -537,7 +537,7 @@ func (s *Server) parseAuthorizationRequest(r *http.Request) (*storage.AuthReques
continue
}
isTrusted, err := s.validateCrossClientTrust(clientID, peerID)
isTrusted, err := s.validateCrossClientTrust(r.Context(), clientID, peerID)
if err != nil {
return nil, newRedirectedErr(errServerError, "Internal server error.")
}
@@ -630,14 +630,14 @@ func parseCrossClientScope(scope string) (peerID string, ok bool) {
return
}
func (s *Server) validateCrossClientTrust(clientID, peerID string) (trusted bool, err error) {
func (s *Server) validateCrossClientTrust(ctx context.Context, clientID, peerID string) (trusted bool, err error) {
if peerID == clientID {
return true, nil
}
peer, err := s.storage.GetClient(peerID)
if err != nil {
if err != storage.ErrNotFound {
s.logger.Error("failed to get client", "err", err)
s.logger.ErrorContext(ctx, "failed to get client", "err", err)
return false, err
}
return false, nil
+20 -20
View File
@@ -80,14 +80,14 @@ type refreshContext struct {
}
// getRefreshTokenFromStorage checks that refresh token is valid and exists in the storage and gets its info
func (s *Server) getRefreshTokenFromStorage(clientID *string, token *internal.RefreshToken) (*refreshContext, *refreshError) {
func (s *Server) getRefreshTokenFromStorage(ctx context.Context, clientID *string, token *internal.RefreshToken) (*refreshContext, *refreshError) {
refreshCtx := refreshContext{requestToken: token}
// Get RefreshToken
refresh, err := s.storage.GetRefresh(token.RefreshId)
if err != nil {
if err != storage.ErrNotFound {
s.logger.Error("failed to get refresh token", "err", err)
s.logger.ErrorContext(ctx, "failed to get refresh token", "err", err)
return nil, newInternalServerError()
}
return nil, invalidErr
@@ -95,7 +95,7 @@ func (s *Server) getRefreshTokenFromStorage(clientID *string, token *internal.Re
// Only check ClientID if it was provided;
if clientID != nil && (refresh.ClientID != *clientID) {
s.logger.Error("trying to claim token for different client", "client_id", clientID, "refresh_client_id", refresh.ClientID)
s.logger.ErrorContext(ctx, "trying to claim token for different client", "client_id", clientID, "refresh_client_id", refresh.ClientID)
// According to https://datatracker.ietf.org/doc/html/rfc6749#section-5.2 Dex should respond with an
// invalid grant error if token has already been claimed by another client.
return nil, &refreshError{msg: errInvalidGrant, desc: invalidErr.desc, code: http.StatusBadRequest}
@@ -108,18 +108,18 @@ func (s *Server) getRefreshTokenFromStorage(clientID *string, token *internal.Re
case refresh.ObsoleteToken != token.Token:
fallthrough
case refresh.ObsoleteToken == "":
s.logger.Error("refresh token claimed twice", "token_id", refresh.ID)
s.logger.ErrorContext(ctx, "refresh token claimed twice", "token_id", refresh.ID)
return nil, invalidErr
}
}
if s.refreshTokenPolicy.CompletelyExpired(refresh.CreatedAt) {
s.logger.Error("refresh token expired", "token_id", refresh.ID)
s.logger.ErrorContext(ctx, "refresh token expired", "token_id", refresh.ID)
return nil, expiredErr
}
if s.refreshTokenPolicy.ExpiredBecauseUnused(refresh.LastUsed) {
s.logger.Error("refresh token expired due to inactivity", "token_id", refresh.ID)
s.logger.ErrorContext(ctx, "refresh token expired due to inactivity", "token_id", refresh.ID)
return nil, expiredErr
}
@@ -128,7 +128,7 @@ func (s *Server) getRefreshTokenFromStorage(clientID *string, token *internal.Re
// Get Connector
refreshCtx.connector, err = s.getConnector(refresh.ConnectorID)
if err != nil {
s.logger.Error("connector not found", "connector_id", refresh.ConnectorID, "err", err)
s.logger.ErrorContext(ctx, "connector not found", "connector_id", refresh.ConnectorID, "err", err)
return nil, newInternalServerError()
}
@@ -137,7 +137,7 @@ func (s *Server) getRefreshTokenFromStorage(clientID *string, token *internal.Re
switch {
case err != nil:
if err != storage.ErrNotFound {
s.logger.Error("failed to get offline session", "err", err)
s.logger.ErrorContext(ctx, "failed to get offline session", "err", err)
return nil, newInternalServerError()
}
case len(refresh.ConnectorData) > 0:
@@ -195,7 +195,7 @@ func (s *Server) refreshWithConnector(ctx context.Context, rCtx *refreshContext,
newIdent, err := refreshConn.Refresh(ctx, parseScopes(rCtx.scopes), ident)
if err != nil {
s.logger.Error("failed to refresh identity", "err", err)
s.logger.ErrorContext(ctx, "failed to refresh identity", "err", err)
return ident, newInternalServerError()
}
@@ -205,7 +205,7 @@ func (s *Server) refreshWithConnector(ctx context.Context, rCtx *refreshContext,
}
// updateOfflineSession updates offline session in the storage
func (s *Server) updateOfflineSession(refresh *storage.RefreshToken, ident connector.Identity, lastUsed time.Time) *refreshError {
func (s *Server) updateOfflineSession(ctx context.Context, refresh *storage.RefreshToken, ident connector.Identity, lastUsed time.Time) *refreshError {
offlineSessionUpdater := func(old storage.OfflineSessions) (storage.OfflineSessions, error) {
if old.Refresh[refresh.ClientID].ID != refresh.ID {
return old, errors.New("refresh token invalid")
@@ -216,7 +216,7 @@ func (s *Server) updateOfflineSession(refresh *storage.RefreshToken, ident conne
old.ConnectorData = ident.ConnectorData
}
s.logger.Debug("saved connector data", "user_id", ident.UserID, "connector_data", ident.ConnectorData)
s.logger.DebugContext(ctx, "saved connector data", "user_id", ident.UserID, "connector_data", ident.ConnectorData)
return old, nil
}
@@ -225,7 +225,7 @@ func (s *Server) updateOfflineSession(refresh *storage.RefreshToken, ident conne
// in offline session for the user.
err := s.storage.UpdateOfflineSessions(refresh.Claims.UserID, refresh.ConnectorID, offlineSessionUpdater)
if err != nil {
s.logger.Error("failed to update offline session", "err", err)
s.logger.ErrorContext(ctx, "failed to update offline session", "err", err)
return newInternalServerError()
}
@@ -316,11 +316,11 @@ func (s *Server) updateRefreshToken(ctx context.Context, rCtx *refreshContext) (
// Update refresh token in the storage.
err := s.storage.UpdateRefreshToken(rCtx.storageToken.ID, refreshTokenUpdater)
if err != nil {
s.logger.Error("failed to update refresh token", "err", err)
s.logger.ErrorContext(ctx, "failed to update refresh token", "err", err)
return nil, ident, newInternalServerError()
}
rerr = s.updateOfflineSession(rCtx.storageToken, ident, lastUsed)
rerr = s.updateOfflineSession(ctx, rCtx.storageToken, ident, lastUsed)
if rerr != nil {
return nil, ident, rerr
}
@@ -337,7 +337,7 @@ func (s *Server) handleRefreshToken(w http.ResponseWriter, r *http.Request, clie
return
}
rCtx, rerr := s.getRefreshTokenFromStorage(&client.ID, token)
rCtx, rerr := s.getRefreshTokenFromStorage(r.Context(), &client.ID, token)
if rerr != nil {
s.refreshTokenErrHelper(w, rerr)
return
@@ -364,23 +364,23 @@ func (s *Server) handleRefreshToken(w http.ResponseWriter, r *http.Request, clie
Groups: ident.Groups,
}
accessToken, _, err := s.newAccessToken(client.ID, claims, rCtx.scopes, rCtx.storageToken.Nonce, rCtx.storageToken.ConnectorID)
accessToken, _, err := s.newAccessToken(r.Context(), client.ID, claims, rCtx.scopes, rCtx.storageToken.Nonce, rCtx.storageToken.ConnectorID)
if err != nil {
s.logger.Error("failed to create new access token", "err", err)
s.logger.ErrorContext(r.Context(), "failed to create new access token", "err", err)
s.refreshTokenErrHelper(w, newInternalServerError())
return
}
idToken, expiry, err := s.newIDToken(client.ID, claims, rCtx.scopes, rCtx.storageToken.Nonce, accessToken, "", rCtx.storageToken.ConnectorID)
idToken, expiry, err := s.newIDToken(r.Context(), client.ID, claims, rCtx.scopes, rCtx.storageToken.Nonce, accessToken, "", rCtx.storageToken.ConnectorID)
if err != nil {
s.logger.Error("failed to create ID token", "err", err)
s.logger.ErrorContext(r.Context(), "failed to create ID token", "err", err)
s.refreshTokenErrHelper(w, newInternalServerError())
return
}
rawNewToken, err := internal.Marshal(newToken)
if err != nil {
s.logger.Error("failed to marshal refresh token", "err", err)
s.logger.ErrorContext(r.Context(), "failed to marshal refresh token", "err", err)
s.refreshTokenErrHelper(w, newInternalServerError())
return
}
+63
View File
@@ -8,7 +8,9 @@ import (
"fmt"
"io/fs"
"log/slog"
"net"
"net/http"
"net/netip"
"net/url"
"os"
"path"
@@ -21,6 +23,7 @@ import (
gosundheit "github.com/AppsFlyer/go-sundheit"
"github.com/felixge/httpsnoop"
"github.com/google/uuid"
"github.com/gorilla/handlers"
"github.com/gorilla/mux"
"github.com/prometheus/client_golang/prometheus"
@@ -85,6 +88,10 @@ type Config struct {
// Headers is a map of headers to be added to the all responses.
Headers http.Header
// Header to extract real ip from.
RealIPHeader string
TrustedRealIPCIDRs []netip.Prefix
// List of allowed origins for CORS requests on discovery, token and keys endpoint.
// If none are indicated, CORS requests are disabled. Passing in "*" will allow any
// domain.
@@ -358,11 +365,52 @@ func newServer(ctx context.Context, c Config, rotationStrategy rotationStrategy)
}
}
parseRealIP := func(r *http.Request) (string, error) {
remoteAddr, _, err := net.SplitHostPort(r.RemoteAddr)
if err != nil {
return "", err
}
remoteIP, err := netip.ParseAddr(remoteAddr)
if err != nil {
return "", err
}
for _, n := range c.TrustedRealIPCIDRs {
if !n.Contains(remoteIP) {
return remoteAddr, nil // Fallback to the address from the request if the header is provided
}
}
ipVal := r.Header.Get(c.RealIPHeader)
if ipVal != "" {
ip, err := netip.ParseAddr(ipVal)
if err == nil {
return ip.String(), nil
}
}
return remoteAddr, nil
}
handlerWithHeaders := func(handlerName string, handler http.Handler) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
for k, v := range c.Headers {
w.Header()[k] = v
}
// Context values are used for logging purposes with the log/slog logger.
rCtx := r.Context()
rCtx = WithRequestID(rCtx)
if c.RealIPHeader != "" {
realIP, err := parseRealIP(r)
if err == nil {
rCtx = WithRemoteIP(rCtx, realIP)
}
}
r = r.WithContext(rCtx)
instrumentHandlerCounter(handlerName, handler)(w, r)
}
}
@@ -682,3 +730,18 @@ func (s *Server) getConnector(id string) (Connector, error) {
return conn, nil
}
type logRequestKey string
const (
RequestKeyRequestID logRequestKey = "request_id"
RequestKeyRemoteIP logRequestKey = "client_remote_addr"
)
func WithRequestID(ctx context.Context) context.Context {
return context.WithValue(ctx, RequestKeyRequestID, uuid.NewString())
}
func WithRemoteIP(ctx context.Context, ip string) context.Context {
return context.WithValue(ctx, RequestKeyRemoteIP, ip)
}