mirror of
https://github.com/netbirdio/dex.git
synced 2026-05-22 18:43:53 -07:00
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:
co-authored by
Márk Sági-Kazár
parent
6ceb26509b
commit
225660785c
+33
-9
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user