From 20e683ce543fb3ca04617b751f78f799de8849b4 Mon Sep 17 00:00:00 2001 From: Mads Jon Nielsen Date: Tue, 4 Aug 2026 12:53:43 +0200 Subject: [PATCH 1/2] Add oidcauth package Authorization code flow with PKCE against an OIDC provider, with the verified id token kept in an httpOnly cookie as the session so there is no session store to run, and role gating read from the client_roles claim. Services that need a browser login hand-roll this today: bearer-guard spends 450 lines on it in package main, so none of it is importable, and mysql-admin another 300. --- go.mod | 5 +- go.sum | 6 + oidcauth/cookie.go | 49 ++++ oidcauth/fakeidp_test.go | 262 +++++++++++++++++ oidcauth/oidcauth.go | 503 ++++++++++++++++++++++++++++++++ oidcauth/oidcauth_test.go | 595 ++++++++++++++++++++++++++++++++++++++ readme.md | 48 +++ 7 files changed, 1467 insertions(+), 1 deletion(-) create mode 100644 oidcauth/cookie.go create mode 100644 oidcauth/fakeidp_test.go create mode 100644 oidcauth/oidcauth.go create mode 100644 oidcauth/oidcauth_test.go diff --git a/go.mod b/go.mod index 696857e..30d190e 100644 --- a/go.mod +++ b/go.mod @@ -1,14 +1,17 @@ module github.com/cego/go-lib/v2 -go 1.25 +go 1.25.0 require ( + github.com/coreos/go-oidc/v3 v3.18.0 github.com/jarcoal/httpmock v1.4.1 github.com/stretchr/testify v1.11.1 + golang.org/x/oauth2 v0.36.0 ) require ( github.com/davecgh/go-spew v1.1.1 // indirect + github.com/go-jose/go-jose/v4 v4.1.4 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect github.com/stretchr/objx v0.5.2 // indirect gopkg.in/yaml.v3 v3.0.1 // indirect diff --git a/go.sum b/go.sum index 41314fc..9fa080a 100644 --- a/go.sum +++ b/go.sum @@ -1,5 +1,9 @@ +github.com/coreos/go-oidc/v3 v3.18.0 h1:V9orjXynvu5wiC9SemFTWnG4F45v403aIcjWo0d41+A= +github.com/coreos/go-oidc/v3 v3.18.0/go.mod h1:DYCf24+ncYi+XkIH97GY1+dqoRlbaSI26KVTCI9SrY4= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/go-jose/go-jose/v4 v4.1.4 h1:moDMcTHmvE6Groj34emNPLs/qtYXRVcd6S7NHbHz3kA= +github.com/go-jose/go-jose/v4 v4.1.4/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08= github.com/jarcoal/httpmock v1.4.1 h1:0Ju+VCFuARfFlhVXFc2HxlcQkfB+Xq12/EotHko+x2A= github.com/jarcoal/httpmock v1.4.1/go.mod h1:ftW1xULwo+j0R0JJkJIIi7UKigZUXCLLanykgjwBXL0= github.com/maxatome/go-testdeep v1.14.0 h1:rRlLv1+kI8eOI3OaBXZwb3O7xY3exRzdW5QyX48g9wI= @@ -10,6 +14,8 @@ github.com/stretchr/objx v0.5.2 h1:xuMeJ0Sdp5ZMRXx/aWO6RZxdr3beISkG5/G/aIRr3pY= github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +golang.org/x/oauth2 v0.36.0 h1:peZ/1z27fi9hUOFCAZaHyrpWG5lwe0RJEEEeH0ThlIs= +golang.org/x/oauth2 v0.36.0/go.mod h1:YDBUJMTkDnJS+A4BP4eZBjCqtokkg1hODuPjwiGPO7Q= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= diff --git a/oidcauth/cookie.go b/oidcauth/cookie.go new file mode 100644 index 0000000..fce8fd2 --- /dev/null +++ b/oidcauth/cookie.go @@ -0,0 +1,49 @@ +package oidcauth + +import ( + "context" + "net/http" + "time" +) + +const ( + sessionCookie = "session" + stateCookie = "state" + nonceCookie = "nonce" + verifierCookie = "verifier" + returnCookie = "return" +) + +func (o *OidcAuth) cookieName(name string) string { + return "__Host-" + o.cookiePrefix + "_" + name +} + +func (o *OidcAuth) setCookie(w http.ResponseWriter, name, value string, maxAge time.Duration) { + http.SetCookie(w, &http.Cookie{ + Name: o.cookieName(name), + Value: value, + Path: "/", + MaxAge: int(maxAge.Seconds()), + HttpOnly: true, + Secure: true, + SameSite: http.SameSiteLaxMode, + }) +} + +func (o *OidcAuth) clearCookie(w http.ResponseWriter, name string) { + o.setCookie(w, name, "", -time.Second) +} + +type contextKey struct{} + +func withUser(ctx context.Context, user User) context.Context { + return context.WithValue(ctx, contextKey{}, user) +} + +func UserFromContext(ctx context.Context) User { + user, ok := ctx.Value(contextKey{}).(User) + if !ok { + return User{} + } + return user +} diff --git a/oidcauth/fakeidp_test.go b/oidcauth/fakeidp_test.go new file mode 100644 index 0000000..1b8213a --- /dev/null +++ b/oidcauth/fakeidp_test.go @@ -0,0 +1,262 @@ +package oidcauth_test + +import ( + "crypto" + "crypto/rand" + "crypto/rsa" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "maps" + "math/big" + "net/http" + "net/http/httptest" + "net/url" + "sync" + "testing" + "time" +) + +const ( + testClientID = "test-client" + testClientSecret = "test-secret" +) + +type fakeIDP struct { + server *httptest.Server + priv *rsa.PrivateKey + kid string + + mu sync.Mutex + pending map[string]map[string]any + challenges map[string]string + nextUser map[string]any + redirects map[string]string + expiry time.Duration +} + +func newFakeIDP(t *testing.T) *fakeIDP { + t.Helper() + priv, err := rsa.GenerateKey(rand.Reader, 2048) + if err != nil { + t.Fatalf("rsa key: %v", err) + } + idp := &fakeIDP{ + priv: priv, + kid: "test-key", + pending: map[string]map[string]any{}, + challenges: map[string]string{}, + redirects: map[string]string{}, + expiry: 5 * time.Minute, + } + mux := http.NewServeMux() + mux.HandleFunc("/.well-known/openid-configuration", idp.handleDiscovery) + mux.HandleFunc("/jwks", idp.handleJWKS) + mux.HandleFunc("/auth", idp.handleAuthorize) + mux.HandleFunc("/token", idp.handleToken) + idp.server = httptest.NewTLSServer(mux) + t.Cleanup(idp.server.Close) + return idp +} + +func (idp *fakeIDP) IssuerURL() string { return idp.server.URL } + +func (idp *fakeIDP) Client() *http.Client { return idp.server.Client() } + +func (idp *fakeIDP) AllowRedirectURI(uri string) { + idp.mu.Lock() + defer idp.mu.Unlock() + idp.redirects[uri] = uri +} + +func (idp *fakeIDP) LoginAs(claims map[string]any) { + idp.mu.Lock() + defer idp.mu.Unlock() + idp.nextUser = claims +} + +func (idp *fakeIDP) MintAccessToken(t *testing.T, audience string, claims map[string]any) string { + t.Helper() + all := idp.idClaims(claims) + all["aud"] = audience + all["azp"] = "some-cli-client" + token, err := idp.signJWT(all) + if err != nil { + t.Fatalf("sign jwt: %v", err) + } + return token +} + +func (idp *fakeIDP) MintIDToken(t *testing.T, claims map[string]any) string { + t.Helper() + token, err := idp.signJWT(idp.idClaims(claims)) + if err != nil { + t.Fatalf("sign jwt: %v", err) + } + return token +} + +func (idp *fakeIDP) idClaims(userClaims map[string]any) map[string]any { + now := time.Now() + claims := map[string]any{ + "iss": idp.server.URL, + "aud": testClientID, + "sub": "test-subject", + "iat": now.Unix(), + "exp": now.Add(idp.expiry).Unix(), + } + maps.Copy(claims, userClaims) + return claims +} + +func (idp *fakeIDP) signJWT(claims map[string]any) (string, error) { + header, _ := json.Marshal(map[string]any{"alg": "RS256", "typ": "JWT", "kid": idp.kid}) + body, _ := json.Marshal(claims) + enc := base64.RawURLEncoding + signingInput := enc.EncodeToString(header) + "." + enc.EncodeToString(body) + sum := sha256.Sum256([]byte(signingInput)) + sig, err := rsa.SignPKCS1v15(rand.Reader, idp.priv, crypto.SHA256, sum[:]) + if err != nil { + return "", err + } + return signingInput + "." + enc.EncodeToString(sig), nil +} + +func (idp *fakeIDP) handleDiscovery(w http.ResponseWriter, _ *http.Request) { + u := idp.server.URL + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{ + "issuer": u, + "authorization_endpoint": u + "/auth", + "token_endpoint": u + "/token", + "end_session_endpoint": u + "/logout", + "jwks_uri": u + "/jwks", + "id_token_signing_alg_values_supported": []string{"RS256"}, + "response_types_supported": []string{"code"}, + "subject_types_supported": []string{"public"}, + }) +} + +func (idp *fakeIDP) handleJWKS(w http.ResponseWriter, _ *http.Request) { + pub := idp.priv.PublicKey + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{ + "keys": []map[string]any{{ + "kty": "RSA", + "use": "sig", + "alg": "RS256", + "kid": idp.kid, + "n": base64.RawURLEncoding.EncodeToString(pub.N.Bytes()), + "e": base64.RawURLEncoding.EncodeToString(big.NewInt(int64(pub.E)).Bytes()), + }}, + }) +} + +func (idp *fakeIDP) handleAuthorize(w http.ResponseWriter, r *http.Request) { + idp.mu.Lock() + redirectURI, registered := idp.redirects[r.URL.Query().Get("redirect_uri")] + idp.mu.Unlock() + if !registered { + http.Error(w, "redirect_uri is not registered for this client", http.StatusBadRequest) + return + } + challenge := r.URL.Query().Get("code_challenge") + if challenge == "" { + http.Error(w, "missing code_challenge", http.StatusBadRequest) + return + } + if method := r.URL.Query().Get("code_challenge_method"); method != "S256" { + http.Error(w, "code_challenge_method must be S256, got "+method, http.StatusBadRequest) + return + } + + idp.mu.Lock() + user := idp.nextUser + idp.nextUser = nil + idp.mu.Unlock() + if user == nil { + http.Error(w, "fakeIDP: /auth called with no LoginAs primed", http.StatusBadRequest) + return + } + + claims := map[string]any{} + maps.Copy(claims, user) + if nonce := r.URL.Query().Get("nonce"); nonce != "" { + claims["nonce"] = nonce + } + + code, err := randomCode() + if err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + idp.mu.Lock() + idp.pending[code] = claims + idp.challenges[code] = challenge + idp.mu.Unlock() + + target, err := url.Parse(redirectURI) + if err != nil { + http.Error(w, "bad redirect_uri", http.StatusBadRequest) + return + } + query := target.Query() + query.Set("code", code) + query.Set("state", r.URL.Query().Get("state")) + target.RawQuery = query.Encode() + http.Redirect(w, r, target.String(), http.StatusFound) +} + +func (idp *fakeIDP) handleToken(w http.ResponseWriter, r *http.Request) { + if err := r.ParseForm(); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + clientID, clientSecret, ok := r.BasicAuth() + if !ok { + clientID, clientSecret = r.PostFormValue("client_id"), r.PostFormValue("client_secret") + } + if clientID != testClientID || clientSecret != testClientSecret { + http.Error(w, "invalid client", http.StatusUnauthorized) + return + } + + code := r.PostFormValue("code") + idp.mu.Lock() + userClaims, found := idp.pending[code] + challenge := idp.challenges[code] + delete(idp.pending, code) + delete(idp.challenges, code) + idp.mu.Unlock() + if !found { + http.Error(w, "invalid code", http.StatusBadRequest) + return + } + + sum := sha256.Sum256([]byte(r.PostFormValue("code_verifier"))) + if base64.RawURLEncoding.EncodeToString(sum[:]) != challenge { + http.Error(w, "PKCE verifier does not match challenge", http.StatusBadRequest) + return + } + + idToken, err := idp.signJWT(idp.idClaims(userClaims)) + if err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{ + "access_token": "not-used-by-oidcauth", + "id_token": idToken, + "token_type": "Bearer", + "expires_in": int(idp.expiry.Seconds()), + }) +} + +func randomCode() (string, error) { + buf := make([]byte, 16) + if _, err := rand.Read(buf); err != nil { + return "", err + } + return base64.RawURLEncoding.EncodeToString(buf), nil +} diff --git a/oidcauth/oidcauth.go b/oidcauth/oidcauth.go new file mode 100644 index 0000000..354001e --- /dev/null +++ b/oidcauth/oidcauth.go @@ -0,0 +1,503 @@ +package oidcauth + +import ( + "context" + "crypto/rand" + "encoding/base64" + "errors" + "fmt" + "net/http" + "net/url" + "slices" + "strings" + "time" + + "github.com/cego/go-lib/v2/headers" + "github.com/cego/go-lib/v2/logger" + "github.com/cego/go-lib/v2/renderer" + coreoidc "github.com/coreos/go-oidc/v3/oidc" + "golang.org/x/oauth2" +) + +const ( + LoginPath = "/auth/login" + CallbackPath = "/auth/callback" + LogoutPath = "/auth/logout" + + DefaultRolesClaim = "client_roles" + DefaultCookiePrefix = "oidcauth" + + transientCookieMaxAge = 10 * time.Minute + + loginExpiredMessage = "login expired, please try again" + loginFailedMessage = "login failed" + + maxSessionCookieBytes = 3800 +) + +type Config struct { + Issuer string + ClientID string + ClientSecret string + BaseURL string +} + +type User struct { + Subject string + Email string + EmailVerified bool + Name string + GivenName string + FamilyName string + PreferredUsername string + Roles []string +} + +func (u User) HasRole(role string) bool { + return slices.Contains(u.Roles, role) +} + +func (u User) HasAnyRole(roles ...string) bool { + if len(roles) == 0 { + return true + } + + for _, role := range roles { + if u.HasRole(role) { + return true + } + } + + return false +} + +type OptionFunc func(o *OidcAuth) + +func WithHTTPClient(httpClient *http.Client) OptionFunc { + return func(o *OidcAuth) { + o.httpClient = httpClient + } +} + +func WithScopes(scopes ...string) OptionFunc { + return func(o *OidcAuth) { + o.scopes = scopes + if !slices.Contains(o.scopes, coreoidc.ScopeOpenID) { + o.scopes = append([]string{coreoidc.ScopeOpenID}, o.scopes...) + } + } +} + +func WithRolesClaim(claim string) OptionFunc { + return func(o *OidcAuth) { + o.rolesClaim = claim + } +} + +func WithBearerAudience(audience string) OptionFunc { + return func(o *OidcAuth) { + o.bearerAudience = audience + } +} + +func WithCookiePrefix(prefix string) OptionFunc { + return func(o *OidcAuth) { + o.cookiePrefix = prefix + } +} + +type OidcAuth struct { + logger logger.Logger + baseURL string + renderer *renderer.Renderer + oauth oauth2.Config + verifier *coreoidc.IDTokenVerifier + bearerVerifier *coreoidc.IDTokenVerifier + bearerAudience string + endSessionURL string + httpClient *http.Client + scopes []string + rolesClaim string + cookiePrefix string +} + +func New(ctx context.Context, l logger.Logger, cfg Config, opts ...OptionFunc) (*OidcAuth, error) { + if err := requireSecureURL("issuer", cfg.Issuer); err != nil { + return nil, err + } + if err := requireSecureURL("base url", cfg.BaseURL); err != nil { + return nil, err + } + if cfg.ClientID == "" { + return nil, errors.New("client id is required") + } + if cfg.ClientSecret == "" { + return nil, errors.New("client secret is required") + } + + o := &OidcAuth{ + logger: l, + renderer: renderer.New(l), + httpClient: &http.Client{Timeout: 10 * time.Second}, + scopes: []string{coreoidc.ScopeOpenID, "profile", "email"}, + rolesClaim: DefaultRolesClaim, + cookiePrefix: DefaultCookiePrefix, + } + for _, opt := range opts { + opt(o) + } + + provider, err := coreoidc.NewProvider(coreoidc.ClientContext(ctx, o.httpClient), cfg.Issuer) + if err != nil { + return nil, fmt.Errorf("discovering issuer %s: %w", cfg.Issuer, err) + } + + var discovery struct { + EndSessionEndpoint string `json:"end_session_endpoint"` + JwksURI string `json:"jwks_uri"` + } + if err := provider.Claims(&discovery); err != nil { + return nil, fmt.Errorf("reading discovery document: %w", err) + } + + for name, endpoint := range map[string]string{ + "authorization_endpoint": provider.Endpoint().AuthURL, + "token_endpoint": provider.Endpoint().TokenURL, + "jwks_uri": discovery.JwksURI, + "end_session_endpoint": discovery.EndSessionEndpoint, + } { + if endpoint == "" { + continue + } + if err := requireSecureURL("discovered "+name, endpoint); err != nil { + return nil, err + } + } + + o.baseURL = strings.TrimSuffix(cfg.BaseURL, "/") + o.oauth = oauth2.Config{ + ClientID: cfg.ClientID, + ClientSecret: cfg.ClientSecret, + Endpoint: provider.Endpoint(), + RedirectURL: o.baseURL + CallbackPath, + Scopes: o.scopes, + } + o.verifier = provider.Verifier(&coreoidc.Config{ClientID: cfg.ClientID}) + if o.bearerAudience == "" { + o.bearerAudience = cfg.ClientID + } + o.bearerVerifier = provider.Verifier(&coreoidc.Config{ClientID: o.bearerAudience}) + o.endSessionURL = discovery.EndSessionEndpoint + + return o, nil +} + +func (o *OidcAuth) Register(mux *http.ServeMux) { + mux.HandleFunc("GET "+LoginPath, o.Login) + mux.HandleFunc("GET "+CallbackPath, o.Callback) + mux.HandleFunc("POST "+LogoutPath, o.Logout) +} + +func (o *OidcAuth) Handler(handler http.Handler, roles ...string) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + user, err := o.userFromRequest(r) + if err != nil { + o.startAuthentication(w, r) + return + } + + if !user.HasAnyRole(roles...) { + o.logger.Info("denied request without required role", "user", user.Email, "required", roles, "roles", user.Roles, "path", r.URL.Path) + o.renderer.Text(w, http.StatusForbidden, "you do not hold any of the roles "+strings.Join(roles, ", ")) + return + } + + handler.ServeHTTP(w, r.WithContext(withUser(r.Context(), user))) + }) +} + +func (o *OidcAuth) HandlerFunc(handlerFunc http.HandlerFunc, roles ...string) http.Handler { + return o.Handler(handlerFunc, roles...) +} + +func (o *OidcAuth) Middleware(roles ...string) func(next http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return o.Handler(next, roles...) + } +} + +func (o *OidcAuth) startAuthentication(w http.ResponseWriter, r *http.Request) { + if r.Header.Get("HX-Request") == "true" { + w.Header().Set("HX-Redirect", LoginPath) + w.WriteHeader(http.StatusUnauthorized) + return + } + + if r.Header.Get(headers.Authorization) != "" || !strings.Contains(r.Header.Get("Accept"), "text/html") { + o.renderer.Text(w, http.StatusUnauthorized, "authentication required") + return + } + + o.Login(w, r) +} + +func (o *OidcAuth) Login(w http.ResponseWriter, r *http.Request) { + state, err := randomString() + if err != nil { + o.renderer.Text(w, http.StatusInternalServerError, "could not start login") + o.logger.Error("failed to generate state", "error", err) + return + } + nonce, err := randomString() + if err != nil { + o.renderer.Text(w, http.StatusInternalServerError, "could not start login") + o.logger.Error("failed to generate nonce", "error", err) + return + } + + verifier := oauth2.GenerateVerifier() + o.setCookie(w, stateCookie, state, transientCookieMaxAge) + o.setCookie(w, nonceCookie, nonce, transientCookieMaxAge) + o.setCookie(w, verifierCookie, verifier, transientCookieMaxAge) + o.setCookie(w, returnCookie, returnTarget(r), transientCookieMaxAge) + + http.Redirect(w, r, o.oauth.AuthCodeURL(state, coreoidc.Nonce(nonce), oauth2.S256ChallengeOption(verifier)), http.StatusFound) +} + +func (o *OidcAuth) Callback(w http.ResponseWriter, r *http.Request) { + state, err := r.Cookie(o.cookieName(stateCookie)) + if err != nil || state.Value == "" || state.Value != r.URL.Query().Get("state") { + o.renderer.Text(w, http.StatusBadRequest, loginExpiredMessage) + return + } + + verifier, err := r.Cookie(o.cookieName(verifierCookie)) + if err != nil { + o.renderer.Text(w, http.StatusBadRequest, loginExpiredMessage) + return + } + + token, err := o.oauth.Exchange(coreoidc.ClientContext(r.Context(), o.httpClient), r.URL.Query().Get("code"), oauth2.VerifierOption(verifier.Value)) + if err != nil { + o.renderer.Text(w, http.StatusUnauthorized, loginFailedMessage) + o.logger.Error("failed to exchange authorization code", "error", err) + return + } + + rawIDToken, ok := token.Extra("id_token").(string) + if !ok { + o.renderer.Text(w, http.StatusUnauthorized, loginFailedMessage) + o.logger.Error("token response had no id_token") + return + } + + idToken, err := o.verifier.Verify(coreoidc.ClientContext(r.Context(), o.httpClient), rawIDToken) + if err != nil { + o.renderer.Text(w, http.StatusUnauthorized, loginFailedMessage) + o.logger.Error("failed to verify id token", "error", err) + return + } + + nonce, err := r.Cookie(o.cookieName(nonceCookie)) + if err != nil || idToken.Nonce != nonce.Value { + o.renderer.Text(w, http.StatusBadRequest, loginExpiredMessage) + return + } + + user, err := o.claimsToUser(idToken, true) + if err != nil { + o.renderer.Text(w, http.StatusUnauthorized, loginFailedMessage) + o.logger.Error("failed to read id token claims", "error", err) + return + } + + if len(rawIDToken) > maxSessionCookieBytes { + o.renderer.Text(w, http.StatusInternalServerError, loginFailedMessage) + o.logger.Error("id token is too large to keep in a cookie", "bytes", len(rawIDToken), "limit", maxSessionCookieBytes) + return + } + + o.setCookie(w, sessionCookie, rawIDToken, time.Until(idToken.Expiry)) + o.clearCookie(w, stateCookie) + o.clearCookie(w, nonceCookie) + o.clearCookie(w, verifierCookie) + + target := "/" + if returnTo, err := r.Cookie(o.cookieName(returnCookie)); err == nil { + target = safeTarget(returnTo.Value) + } + o.clearCookie(w, returnCookie) + + o.logger.Info("logged in", "user", user.Email, "roles", user.Roles) + http.Redirect(w, r, target, http.StatusFound) +} + +func (o *OidcAuth) Logout(w http.ResponseWriter, r *http.Request) { + o.clearCookie(w, sessionCookie) + + endSession, err := url.Parse(o.endSessionURL) + if err != nil || o.endSessionURL == "" { + http.Redirect(w, r, "/", http.StatusSeeOther) + return + } + + query := endSession.Query() + query.Set("client_id", o.oauth.ClientID) + query.Set("post_logout_redirect_uri", o.baseURL+"/") + endSession.RawQuery = query.Encode() + + http.Redirect(w, r, endSession.String(), http.StatusSeeOther) +} + +func (o *OidcAuth) userFromRequest(r *http.Request) (User, error) { + ctx := coreoidc.ClientContext(r.Context(), o.httpClient) + + if authorization := r.Header.Get(headers.Authorization); authorization != "" { + bearer, found := strings.CutPrefix(authorization, "Bearer ") + if !found { + return User{}, errors.New("authorization header is not a bearer token") + } + + token, err := o.bearerVerifier.Verify(ctx, bearer) + if err != nil { + return User{}, fmt.Errorf("verifying bearer token: %w", err) + } + + return o.claimsToUser(token, false) + } + + cookie, err := r.Cookie(o.cookieName(sessionCookie)) + if err != nil { + return User{}, errors.New("no session") + } + + idToken, err := o.verifier.Verify(ctx, cookie.Value) + if err != nil { + return User{}, fmt.Errorf("verifying session: %w", err) + } + + return o.claimsToUser(idToken, true) +} + +func (o *OidcAuth) claimsToUser(idToken *coreoidc.IDToken, requireOurAuthorizedParty bool) (User, error) { + claims := map[string]any{} + if err := idToken.Claims(&claims); err != nil { + return User{}, err + } + + if requireOurAuthorizedParty { + if err := o.validateAuthorizedParty(idToken, claims); err != nil { + return User{}, err + } + } + if idToken.Subject == "" { + return User{}, errors.New("token has no sub claim") + } + + user := User{ + Subject: idToken.Subject, + Email: stringFrom(claims, "email"), + Name: stringFrom(claims, "name"), + GivenName: stringFrom(claims, "given_name"), + FamilyName: stringFrom(claims, "family_name"), + PreferredUsername: stringFrom(claims, "preferred_username"), + Roles: o.rolesFrom(claims), + } + user.EmailVerified, _ = claims["email_verified"].(bool) + + return user, nil +} + +func (o *OidcAuth) validateAuthorizedParty(idToken *coreoidc.IDToken, claims map[string]any) error { + authorizedParty := stringFrom(claims, "azp") + if authorizedParty != "" && authorizedParty != o.oauth.ClientID { + return fmt.Errorf("token authorized for %q, not %q", authorizedParty, o.oauth.ClientID) + } + if authorizedParty == "" && len(idToken.Audience) > 1 { + return errors.New("token has several audiences and no azp claim") + } + + return nil +} + +func (o *OidcAuth) rolesFrom(claims map[string]any) []string { + roles := stringsFrom(claims[o.rolesClaim]) + + if resourceAccess, ok := claims["resource_access"].(map[string]any); ok { + if ours, ok := resourceAccess[o.oauth.ClientID].(map[string]any); ok { + roles = append(roles, stringsFrom(ours["roles"])...) + } + } + + return roles +} + +func stringFrom(claims map[string]any, key string) string { + value, _ := claims[key].(string) + return value +} + +func requireSecureURL(name, value string) error { + parsed, err := url.Parse(value) + if err != nil || parsed.Host == "" { + return fmt.Errorf("%s must be an absolute url, got %q", name, value) + } + + if parsed.Scheme == "https" || (parsed.Scheme == "http" && isLoopback(parsed.Hostname())) { + return nil + } + + return fmt.Errorf("%s must be https, got %q", name, value) +} + +func isLoopback(host string) bool { + return host == "localhost" || host == "127.0.0.1" || host == "::1" +} + +func safeTarget(value string) string { + if !strings.HasPrefix(value, "/") || strings.HasPrefix(value, "//") || strings.HasPrefix(value, `/\`) { + return "/" + } + + parsed, err := url.Parse(value) + if err != nil || parsed.Scheme != "" || parsed.Host != "" || isAuthPath(parsed.Path) { + return "/" + } + + return value +} + +func stringsFrom(claim any) []string { + values, ok := claim.([]any) + if !ok { + return nil + } + + var strings []string + for _, value := range values { + if value, ok := value.(string); ok { + strings = append(strings, value) + } + } + + return strings +} + +func returnTarget(r *http.Request) string { + if r.Method != http.MethodGet || isAuthPath(r.URL.Path) { + return "/" + } + return r.URL.RequestURI() +} + +func isAuthPath(path string) bool { + return path == LoginPath || path == CallbackPath || path == LogoutPath +} + +func randomString() (string, error) { + buf := make([]byte, 32) + if _, err := rand.Read(buf); err != nil { + return "", err + } + return base64.RawURLEncoding.EncodeToString(buf), nil +} diff --git a/oidcauth/oidcauth_test.go b/oidcauth/oidcauth_test.go new file mode 100644 index 0000000..1e0e3a3 --- /dev/null +++ b/oidcauth/oidcauth_test.go @@ -0,0 +1,595 @@ +package oidcauth_test + +import ( + "context" + "net/http" + "net/http/cookiejar" + "net/http/httptest" + "net/url" + "testing" + + "github.com/cego/go-lib/v2/logger" + "github.com/cego/go-lib/v2/oidcauth" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func browserGet(t *testing.T, client *http.Client, url string) *http.Response { + t.Helper() + request, err := http.NewRequest(http.MethodGet, url, nil) + require.NoError(t, err) + request.Header.Set("Accept", "text/html") + + response, err := client.Do(request) + require.NoError(t, err) + + return response +} + +func newTestServer(t *testing.T, idp *fakeIDP, role string, opts ...oidcauth.OptionFunc) (*httptest.Server, *http.Client) { + t.Helper() + + mux := http.NewServeMux() + server := httptest.NewTLSServer(mux) + t.Cleanup(server.Close) + idp.AllowRedirectURI(server.URL + oidcauth.CallbackPath) + + auth, err := oidcauth.New(context.Background(), logger.NewMock(), oidcauth.Config{ + Issuer: idp.IssuerURL(), + ClientID: testClientID, + ClientSecret: testClientSecret, + BaseURL: server.URL, + }, append([]oidcauth.OptionFunc{oidcauth.WithHTTPClient(idp.Client())}, opts...)...) + require.NoError(t, err) + + auth.Register(mux) + handler := func(w http.ResponseWriter, r *http.Request) { + user := oidcauth.UserFromContext(r.Context()) + _, _ = w.Write([]byte(user.Email + " " + user.Subject)) + } + + var protected http.Handler + if role == "" { + protected = auth.HandlerFunc(handler) + } else { + protected = auth.HandlerFunc(handler, role) + } + mux.Handle("GET /protected", protected) + mux.Handle("POST /protected", protected) + mux.Handle("GET /names", auth.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + user := oidcauth.UserFromContext(r.Context()) + _, _ = w.Write([]byte(user.GivenName + " " + user.FamilyName + " " + user.PreferredUsername)) + }, "reader")) + + jar, err := cookiejar.New(nil) + require.NoError(t, err) + + client := server.Client() + client.Jar = jar + + return server, client +} + +func TestOidcAuth(t *testing.T) { + t.Run("a user with the role completes the flow and reaches the handler", func(t *testing.T) { + idp := newFakeIDP(t) + server, client := newTestServer(t, idp, "reader") + idp.LoginAs(map[string]any{"email": "mjn@cego.dk", "name": "Mads", "client_roles": []any{"reader"}}) + + response := browserGet(t, client, server.URL+"/protected") + defer func() { _ = response.Body.Close() }() + + assert.Equal(t, http.StatusOK, response.StatusCode) + body := make([]byte, 64) + n, _ := response.Body.Read(body) + assert.Contains(t, string(body[:n]), "mjn@cego.dk") + }) + + t.Run("a user without the role is refused", func(t *testing.T) { + idp := newFakeIDP(t) + server, client := newTestServer(t, idp, "process-admin") + idp.LoginAs(map[string]any{"email": "lejo@cego.dk", "client_roles": []any{"reader"}}) + + response := browserGet(t, client, server.URL+"/protected") + defer func() { _ = response.Body.Close() }() + + assert.Equal(t, http.StatusForbidden, response.StatusCode) + }) + + t.Run("no session redirects to the provider", func(t *testing.T) { + idp := newFakeIDP(t) + server, _ := newTestServer(t, idp, "reader") + client := server.Client() + client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse } + + response := browserGet(t, client, server.URL+"/protected") + defer func() { _ = response.Body.Close() }() + + assert.Equal(t, http.StatusFound, response.StatusCode) + location, err := url.Parse(response.Header.Get("Location")) + require.NoError(t, err) + assert.Equal(t, idp.IssuerURL()+"/auth", location.Scheme+"://"+location.Host+location.Path) + assert.Equal(t, testClientID, location.Query().Get("client_id")) + assert.Equal(t, "S256", location.Query().Get("code_challenge_method")) + assert.NotEmpty(t, location.Query().Get("state")) + assert.NotEmpty(t, location.Query().Get("nonce")) + }) + + t.Run("an htmx request is told to redirect instead of being sent to the provider", func(t *testing.T) { + idp := newFakeIDP(t) + server, _ := newTestServer(t, idp, "reader") + client := server.Client() + client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse } + + request, err := http.NewRequest(http.MethodGet, server.URL+"/protected", nil) + require.NoError(t, err) + request.Header.Set("HX-Request", "true") + + response, err := client.Do(request) + require.NoError(t, err) + defer func() { _ = response.Body.Close() }() + + assert.Equal(t, http.StatusUnauthorized, response.StatusCode) + assert.Equal(t, oidcauth.LoginPath, response.Header.Get("HX-Redirect")) + }) + + t.Run("a forged session cookie is not trusted", func(t *testing.T) { + idp := newFakeIDP(t) + server, _ := newTestServer(t, idp, "reader") + client := server.Client() + client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse } + + request, err := http.NewRequest(http.MethodGet, server.URL+"/protected", nil) + require.NoError(t, err) + request.Header.Set("Accept", "text/html") + request.AddCookie(&http.Cookie{Name: "__Host-oidcauth_session", Value: "not.a.jwt"}) + + response, err := client.Do(request) + require.NoError(t, err) + defer func() { _ = response.Body.Close() }() + + assert.Equal(t, http.StatusFound, response.StatusCode) + }) + + t.Run("a token signed for another client is not trusted", func(t *testing.T) { + idp := newFakeIDP(t) + server, _ := newTestServer(t, idp, "reader") + client := server.Client() + client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse } + + request, err := http.NewRequest(http.MethodGet, server.URL+"/protected", nil) + require.NoError(t, err) + request.Header.Set("Accept", "text/html") + request.AddCookie(&http.Cookie{ + Name: "__Host-oidcauth_session", + Value: idp.MintIDToken(t, map[string]any{"aud": "some-other-client", "client_roles": []any{"reader"}}), + }) + + response, err := client.Do(request) + require.NoError(t, err) + defer func() { _ = response.Body.Close() }() + + assert.Equal(t, http.StatusFound, response.StatusCode) + }) + + t.Run("callback without the state cookie is rejected", func(t *testing.T) { + idp := newFakeIDP(t) + server, _ := newTestServer(t, idp, "reader") + + response, err := server.Client().Get(server.URL + oidcauth.CallbackPath + "?code=whatever&state=whatever") + require.NoError(t, err) + defer func() { _ = response.Body.Close() }() + + assert.Equal(t, http.StatusBadRequest, response.StatusCode) + }) + + t.Run("roles can be read from another claim", func(t *testing.T) { + idp := newFakeIDP(t) + server, client := newTestServer(t, idp, "reader", oidcauth.WithRolesClaim("realm_roles")) + idp.LoginAs(map[string]any{"email": "mjn@cego.dk", "realm_roles": []any{"reader"}}) + + response := browserGet(t, client, server.URL+"/protected") + defer func() { _ = response.Body.Close() }() + + assert.Equal(t, http.StatusOK, response.StatusCode) + }) + + t.Run("logout clears the session and ends it at the provider", func(t *testing.T) { + idp := newFakeIDP(t) + server, _ := newTestServer(t, idp, "reader") + client := server.Client() + client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse } + + response, err := client.Post(server.URL+oidcauth.LogoutPath, "", nil) + require.NoError(t, err) + defer func() { _ = response.Body.Close() }() + + assert.Equal(t, http.StatusSeeOther, response.StatusCode) + assert.Contains(t, response.Header.Get("Location"), idp.IssuerURL()+"/logout") + assert.Contains(t, response.Header.Get("Set-Cookie"), "oidcauth_session=;") + }) + + t.Run("the options are applied", func(t *testing.T) { + idp := newFakeIDP(t) + server, _ := newTestServer(t, idp, "reader", + oidcauth.WithHTTPClient(idp.Client()), + oidcauth.WithScopes("openid", "email"), + oidcauth.WithCookiePrefix("myservice"), + ) + client := server.Client() + client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse } + + response := browserGet(t, client, server.URL+"/protected") + defer func() { _ = response.Body.Close() }() + + location, err := url.Parse(response.Header.Get("Location")) + require.NoError(t, err) + assert.Equal(t, "openid email", location.Query().Get("scope")) + assert.NotEmpty(t, cookieValue(t, response, "__Host-myservice_state")) + }) + + t.Run("the visited url is returned to after login", func(t *testing.T) { + idp := newFakeIDP(t) + server, client := newTestServer(t, idp, "reader") + idp.LoginAs(map[string]any{"email": "mjn@cego.dk", "client_roles": []any{"reader"}}) + + response := browserGet(t, client, server.URL+"/protected?sort=time") + defer func() { _ = response.Body.Close() }() + + assert.Equal(t, http.StatusOK, response.StatusCode) + assert.Equal(t, "/protected?sort=time", response.Request.URL.RequestURI()) + }) + + t.Run("a callback with an unusable code fails the login", func(t *testing.T) { + idp := newFakeIDP(t) + server, client := newTestServer(t, idp, "reader") + client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse } + + started := browserGet(t, client, server.URL+"/protected") + _ = started.Body.Close() + location, err := url.Parse(started.Header.Get("Location")) + require.NoError(t, err) + + response, err := client.Get(server.URL + oidcauth.CallbackPath + "?code=not-a-real-code&state=" + location.Query().Get("state")) + require.NoError(t, err) + defer func() { _ = response.Body.Close() }() + + assert.Equal(t, http.StatusUnauthorized, response.StatusCode) + }) + + t.Run("a post is returned to the root after login", func(t *testing.T) { + idp := newFakeIDP(t) + server, client := newTestServer(t, idp, "reader") + client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse } + + request, err := http.NewRequest(http.MethodPost, server.URL+"/protected", nil) + require.NoError(t, err) + request.Header.Set("Accept", "text/html") + + response, err := client.Do(request) + require.NoError(t, err) + defer func() { _ = response.Body.Close() }() + + assert.Equal(t, "/", cookieValue(t, response, "__Host-oidcauth_return")) + }) + + t.Run("a forged return cookie cannot redirect off this host", func(t *testing.T) { + for _, target := range []string{"//evil.com", "https://evil.com/x", "not-a-path"} { + idp := newFakeIDP(t) + server, client := newTestServer(t, idp, "reader") + client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse } + idp.LoginAs(map[string]any{"email": "mjn@cego.dk", "client_roles": []any{"reader"}}) + + started := browserGet(t, client, server.URL+"/protected") + _ = started.Body.Close() + serverURL, err := url.Parse(server.URL) + require.NoError(t, err) + client.Jar.SetCookies(serverURL, []*http.Cookie{{Name: "__Host-oidcauth_return", Value: target, Path: "/"}}) + + atProvider, err := client.Get(started.Header.Get("Location")) + require.NoError(t, err) + _ = atProvider.Body.Close() + atCallback, err := client.Get(atProvider.Header.Get("Location")) + require.NoError(t, err) + _ = atCallback.Body.Close() + + assert.Equal(t, http.StatusFound, atCallback.StatusCode, target) + assert.Equal(t, "/", atCallback.Header.Get("Location"), "login must not redirect to %s", target) + } + }) + + t.Run("http issuer and base url are refused", func(t *testing.T) { + idp := newFakeIDP(t) + _, err := oidcauth.New(context.Background(), logger.NewMock(), oidcauth.Config{ + Issuer: "http://keycloak.example.com/realms/cego", ClientID: testClientID, + ClientSecret: testClientSecret, BaseURL: "https://example.com", + }) + require.ErrorContains(t, err, "issuer must be https") + + _, err = oidcauth.New(context.Background(), logger.NewMock(), oidcauth.Config{ + Issuer: idp.IssuerURL(), ClientID: testClientID, + ClientSecret: testClientSecret, BaseURL: "http://example.com", + }, oidcauth.WithHTTPClient(idp.Client())) + require.ErrorContains(t, err, "base url must be https") + }) + + t.Run("a token authorized for another client is refused", func(t *testing.T) { + idp := newFakeIDP(t) + server, _ := newTestServer(t, idp, "reader") + client := server.Client() + client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse } + + for _, claims := range []map[string]any{ + {"azp": "another-client", "client_roles": []any{"reader"}}, + {"aud": []string{testClientID, "another-client"}, "client_roles": []any{"reader"}}, + {"sub": "", "client_roles": []any{"reader"}}, + } { + request, err := http.NewRequest(http.MethodGet, server.URL+"/protected", nil) + require.NoError(t, err) + request.Header.Set("Accept", "text/html") + request.AddCookie(&http.Cookie{Name: "__Host-oidcauth_session", Value: idp.MintIDToken(t, claims)}) + + response, err := client.Do(request) + require.NoError(t, err) + _ = response.Body.Close() + assert.Equal(t, http.StatusFound, response.StatusCode, claims) + } + }) + + t.Run("visiting the login route does not loop back to itself", func(t *testing.T) { + idp := newFakeIDP(t) + server, client := newTestServer(t, idp, "reader") + client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse } + + response := browserGet(t, client, server.URL+oidcauth.LoginPath) + defer func() { _ = response.Body.Close() }() + + assert.Equal(t, "/", cookieValue(t, response, "__Host-oidcauth_return")) + }) + + t.Run("logout never puts the token in the url", func(t *testing.T) { + idp := newFakeIDP(t) + server, _ := newTestServer(t, idp, "reader") + client := server.Client() + client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse } + + token := idp.MintIDToken(t, map[string]any{"client_roles": []any{"reader"}}) + request, err := http.NewRequest(http.MethodPost, server.URL+oidcauth.LogoutPath, nil) + require.NoError(t, err) + request.AddCookie(&http.Cookie{Name: "__Host-oidcauth_session", Value: token}) + + response, err := client.Do(request) + require.NoError(t, err) + defer func() { _ = response.Body.Close() }() + + location := response.Header.Get("Location") + assert.NotContains(t, location, token) + parsed, err := url.Parse(location) + require.NoError(t, err) + assert.Empty(t, parsed.Query().Get("id_token_hint")) + assert.Equal(t, testClientID, parsed.Query().Get("client_id")) + assert.Equal(t, server.URL+"/", parsed.Query().Get("post_logout_redirect_uri")) + }) + + t.Run("the openid scope cannot be dropped", func(t *testing.T) { + idp := newFakeIDP(t) + server, _ := newTestServer(t, idp, "reader", oidcauth.WithScopes("email")) + client := server.Client() + client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse } + + response := browserGet(t, client, server.URL+"/protected") + defer func() { _ = response.Body.Close() }() + + location, err := url.Parse(response.Header.Get("Location")) + require.NoError(t, err) + assert.Equal(t, "openid email", location.Query().Get("scope")) + }) + + t.Run("an api client authenticates with a bearer token", func(t *testing.T) { + idp := newFakeIDP(t) + server, _ := newTestServer(t, idp, "reader") + + request, err := http.NewRequest(http.MethodGet, server.URL+"/protected", nil) + require.NoError(t, err) + request.Header.Set("Authorization", "Bearer "+idp.MintAccessToken(t, testClientID, map[string]any{ + "email": "robot@cego.dk", "client_roles": []any{"reader"}, + })) + + response, err := server.Client().Do(request) + require.NoError(t, err) + defer func() { _ = response.Body.Close() }() + + assert.Equal(t, http.StatusOK, response.StatusCode) + }) + + t.Run("an api client is answered 401 rather than redirected", func(t *testing.T) { + idp := newFakeIDP(t) + server, _ := newTestServer(t, idp, "reader") + client := server.Client() + client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse } + + response, err := client.Get(server.URL + "/protected") + require.NoError(t, err) + _ = response.Body.Close() + assert.Equal(t, http.StatusUnauthorized, response.StatusCode) + + request, err := http.NewRequest(http.MethodGet, server.URL+"/protected", nil) + require.NoError(t, err) + request.Header.Set("Authorization", "Bearer not.a.jwt") + request.Header.Set("Accept", "text/html") + + response, err = client.Do(request) + require.NoError(t, err) + defer func() { _ = response.Body.Close() }() + assert.Equal(t, http.StatusUnauthorized, response.StatusCode) + }) + + t.Run("a bearer audience can differ from the client id", func(t *testing.T) { + idp := newFakeIDP(t) + server, _ := newTestServer(t, idp, "reader", oidcauth.WithBearerAudience("rule-engine-api")) + + request, err := http.NewRequest(http.MethodGet, server.URL+"/protected", nil) + require.NoError(t, err) + request.Header.Set("Authorization", "Bearer "+idp.MintAccessToken(t, "rule-engine-api", map[string]any{ + "client_roles": []any{"reader"}, + })) + + response, err := server.Client().Do(request) + require.NoError(t, err) + defer func() { _ = response.Body.Close() }() + + assert.Equal(t, http.StatusOK, response.StatusCode) + }) + + t.Run("a route can require authentication without any role", func(t *testing.T) { + idp := newFakeIDP(t) + server, client := newTestServer(t, idp, "") + idp.LoginAs(map[string]any{"email": "anyone@cego.dk"}) + + response := browserGet(t, client, server.URL+"/protected") + defer func() { _ = response.Body.Close() }() + + assert.Equal(t, http.StatusOK, response.StatusCode) + }) + + t.Run("roles are read from resource_access when there is no flat claim", func(t *testing.T) { + idp := newFakeIDP(t) + server, client := newTestServer(t, idp, "tool-admin") + idp.LoginAs(map[string]any{"email": "mjn@cego.dk", "resource_access": map[string]any{ + testClientID: map[string]any{"roles": []any{"tool-admin"}}, + }}) + + response := browserGet(t, client, server.URL+"/protected") + defer func() { _ = response.Body.Close() }() + + assert.Equal(t, http.StatusOK, response.StatusCode) + }) + + t.Run("http on loopback is allowed for local development", func(t *testing.T) { + idp := newFakeIDP(t) + _, err := oidcauth.New(context.Background(), logger.NewMock(), oidcauth.Config{ + Issuer: idp.IssuerURL(), ClientID: testClientID, + ClientSecret: testClientSecret, BaseURL: "http://localhost:8080", + }, oidcauth.WithHTTPClient(idp.Client())) + require.NoError(t, err) + }) + + t.Run("the gate also works as chi middleware", func(t *testing.T) { + idp := newFakeIDP(t) + + mux := http.NewServeMux() + server := httptest.NewTLSServer(mux) + t.Cleanup(server.Close) + idp.AllowRedirectURI(server.URL + oidcauth.CallbackPath) + + auth, err := oidcauth.New(context.Background(), logger.NewMock(), oidcauth.Config{ + Issuer: idp.IssuerURL(), ClientID: testClientID, + ClientSecret: testClientSecret, BaseURL: server.URL, + }, oidcauth.WithHTTPClient(idp.Client())) + require.NoError(t, err) + + auth.Register(mux) + gate := auth.Middleware("reader") + mux.Handle("GET /wrapped", gate(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + user := oidcauth.UserFromContext(r.Context()) + _, _ = w.Write([]byte(user.Email)) + assert.True(t, user.EmailVerified) + }))) + + jar, err := cookiejar.New(nil) + require.NoError(t, err) + client := server.Client() + client.Jar = jar + idp.LoginAs(map[string]any{"email": "mjn@cego.dk", "email_verified": true, "client_roles": []any{"reader"}}) + + response := browserGet(t, client, server.URL+"/wrapped") + defer func() { _ = response.Body.Close() }() + + assert.Equal(t, http.StatusOK, response.StatusCode) + }) + + t.Run("an authorization header that is not a bearer token is refused", func(t *testing.T) { + idp := newFakeIDP(t) + server, client := newTestServer(t, idp, "reader") + idp.LoginAs(map[string]any{"email": "mjn@cego.dk", "client_roles": []any{"reader"}}) + + loggedIn := browserGet(t, client, server.URL+"/protected") + _ = loggedIn.Body.Close() + require.Equal(t, http.StatusOK, loggedIn.StatusCode) + + client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse } + + for _, authorization := range []string{"Basic dXNlcjpwYXNz", "bearer lowercase", "Bearer not.a.jwt"} { + request, err := http.NewRequest(http.MethodGet, server.URL+"/protected", nil) + require.NoError(t, err) + request.Header.Set("Accept", "text/html") + request.Header.Set("Authorization", authorization) + + response, err := client.Do(request) + require.NoError(t, err) + _ = response.Body.Close() + + assert.Equal(t, http.StatusUnauthorized, response.StatusCode, authorization) + } + }) + + t.Run("the configuration is validated", func(t *testing.T) { + idp := newFakeIDP(t) + for _, tc := range []struct { + cfg oidcauth.Config + want string + }{ + {oidcauth.Config{Issuer: "not-a-url", ClientID: "c", ClientSecret: "s", BaseURL: "https://example.com"}, "issuer must be an absolute url"}, + {oidcauth.Config{Issuer: idp.IssuerURL(), ClientID: "c", ClientSecret: "s", BaseURL: "/relative"}, "base url must be an absolute url"}, + {oidcauth.Config{Issuer: idp.IssuerURL(), ClientID: "", ClientSecret: "s", BaseURL: "https://example.com"}, "client id is required"}, + {oidcauth.Config{Issuer: idp.IssuerURL(), ClientID: "c", ClientSecret: "", BaseURL: "https://example.com"}, "client secret is required"}, + } { + _, err := oidcauth.New(context.Background(), logger.NewMock(), tc.cfg, oidcauth.WithHTTPClient(idp.Client())) + require.ErrorContains(t, err, tc.want) + } + }) + + t.Run("the standard name claims are exposed", func(t *testing.T) { + idp := newFakeIDP(t) + server, client := newTestServer(t, idp, "reader") + idp.LoginAs(map[string]any{ + "email": "mjn@cego.dk", "given_name": "Mads", "family_name": "Nielsen", + "preferred_username": "mjn", "client_roles": []any{"reader"}, + }) + + response := browserGet(t, client, server.URL+"/names") + defer func() { _ = response.Body.Close() }() + + body := make([]byte, 128) + n, _ := response.Body.Read(body) + assert.Equal(t, "Mads Nielsen mjn", string(body[:n])) + }) + + t.Run("discovery failure is reported", func(t *testing.T) { + _, err := oidcauth.New(context.Background(), logger.NewMock(), oidcauth.Config{ + Issuer: "https://127.0.0.1:1/realms/nope", + ClientID: testClientID, + ClientSecret: testClientSecret, + BaseURL: "https://example.com", + }) + assert.ErrorContains(t, err, "discovering issuer") + }) +} + +func cookieValue(t *testing.T, response *http.Response, name string) string { + t.Helper() + for _, cookie := range response.Cookies() { + if cookie.Name == name { + return cookie.Value + } + } + t.Fatalf("response did not set cookie %s", name) + return "" +} + +func TestUserHasRole(t *testing.T) { + user := oidcauth.User{Roles: []string{"reader", "process-admin"}} + assert.True(t, user.HasRole("process-admin")) + assert.False(t, user.HasRole("nope")) + assert.False(t, oidcauth.User{}.HasRole("reader")) +} diff --git a/readme.md b/readme.md index 3a66097..c7ad377 100644 --- a/readme.md +++ b/readme.md @@ -17,6 +17,7 @@ import ( "github.com/cego/go-lib/v2/logger" "github.com/cego/go-lib/v2/renderer" "github.com/cego/go-lib/v2/forwardauth" + "github.com/cego/go-lib/v2/oidcauth" "github.com/cego/go-lib/v2/headers" "github.com/cego/go-lib/v2/serve" "github.com/cego/go-lib/v2/periodic" @@ -82,6 +83,53 @@ mux.Handle("/data", fa.HandlerFunc(func (w http.ResponseWriter, req *http.Reques })) ``` +## Using OidcAuth + +Authorization code flow with PKCE for browsers, bearer tokens for api clients. The verified id token +is the session cookie, so there is no session store. + +```go +mux := http.NewServeMux() +auth, err := oidcauth.New(context.Background(), l, oidcauth.Config{ + Issuer: "https://keycloak.example.com/realms/cego", + ClientID: "myservice", + ClientSecret: os.Getenv("MYSERVICE_OIDC_CLIENT_SECRET"), + BaseURL: "https://myservice.example.com", +}) + +auth.Register(mux) // Adds GET /auth/login, GET /auth/callback and POST /auth/logout + +mux.Handle("GET /{$}", auth.HandlerFunc(index)) // Authenticated +mux.Handle("GET /things", auth.HandlerFunc(things, "reader", "tool-admin")) // Either role +mux.Handle("POST /things/{id}/delete", auth.HandlerFunc(del, "tool-admin")) // That role + +r.Use(auth.Middleware("reader")) // The same gate as chi middleware + +// Inside a wrapped handler, or a template +user := oidcauth.UserFromContext(r.Context()) +user.HasRole("tool-admin") +user.HasAnyRole("reader", "tool-admin") +``` + +- Register `/auth/callback` on the client, or the provider rejects the login +- Roles are read from `client_roles` and `resource_access..roles` +- Issuer and base url must be https, except on loopback +- No session: browsers go to the provider, htmx gets `HX-Redirect`, api clients get `401` +- Logout is a `POST`, and hands the provider no token, so it asks the user to confirm +- Sessions cannot be revoked before the id token expires, so keep that lifetime short +- Cookies are `__Host-` prefixed, secure and `SameSite=Lax`, which is not a CSRF token + +### Options +```go +auth, err := oidcauth.New(ctx, l, cfg, + oidcauth.WithHTTPClient(httpClient), // default timeout 10s + oidcauth.WithScopes("openid", "email"), // default openid, profile, email + oidcauth.WithRolesClaim("realm_roles"), // default client_roles + oidcauth.WithCookiePrefix("myservice"), // default oidcauth + oidcauth.WithBearerAudience("myservice-api"), // default the client id +) +``` + ## Headers ```go req.Header.Get(headers.Authorization) From 3825c0ce0339fc6eb21bbc11df3e6329eb201fcf Mon Sep 17 00:00:00 2001 From: Mads Jon Nielsen Date: Wed, 5 Aug 2026 11:29:52 +0200 Subject: [PATCH 2/2] Log ecs shaped fields User implements slog.LogValuer so "user", user renders user.id, user.type, user.email, user.full_name and user.roles, matching what the php package logs, and the request path and errors go through url.path and the go-lib error helper. --- oidcauth/oidcauth.go | 30 +++++++++++++++++++++++------- oidcauth/oidcauth_test.go | 33 +++++++++++++++++++++++++++++++++ readme.md | 1 + 3 files changed, 57 insertions(+), 7 deletions(-) diff --git a/oidcauth/oidcauth.go b/oidcauth/oidcauth.go index 354001e..52525f1 100644 --- a/oidcauth/oidcauth.go +++ b/oidcauth/oidcauth.go @@ -6,6 +6,7 @@ import ( "encoding/base64" "errors" "fmt" + "log/slog" "net/http" "net/url" "slices" @@ -53,6 +54,21 @@ type User struct { Roles []string } +func (u User) LogValue() slog.Value { + id := u.PreferredUsername + if id == "" { + id = u.Subject + } + + return slog.GroupValue( + slog.String("id", id), + slog.String("type", "oidc"), + slog.String("email", u.Email), + slog.String("full_name", u.Name), + slog.Any("roles", u.Roles), + ) +} + func (u User) HasRole(role string) bool { return slices.Contains(u.Roles, role) } @@ -207,7 +223,7 @@ func (o *OidcAuth) Handler(handler http.Handler, roles ...string) http.Handler { } if !user.HasAnyRole(roles...) { - o.logger.Info("denied request without required role", "user", user.Email, "required", roles, "roles", user.Roles, "path", r.URL.Path) + o.logger.Info("denied request without required role", "user", user, slog.Group("url", slog.String("path", r.URL.Path)), "required_roles", roles) o.renderer.Text(w, http.StatusForbidden, "you do not hold any of the roles "+strings.Join(roles, ", ")) return } @@ -245,13 +261,13 @@ func (o *OidcAuth) Login(w http.ResponseWriter, r *http.Request) { state, err := randomString() if err != nil { o.renderer.Text(w, http.StatusInternalServerError, "could not start login") - o.logger.Error("failed to generate state", "error", err) + o.logger.Error("failed to generate state", logger.GetSlogAttrFromError(err)) return } nonce, err := randomString() if err != nil { o.renderer.Text(w, http.StatusInternalServerError, "could not start login") - o.logger.Error("failed to generate nonce", "error", err) + o.logger.Error("failed to generate nonce", logger.GetSlogAttrFromError(err)) return } @@ -280,7 +296,7 @@ func (o *OidcAuth) Callback(w http.ResponseWriter, r *http.Request) { token, err := o.oauth.Exchange(coreoidc.ClientContext(r.Context(), o.httpClient), r.URL.Query().Get("code"), oauth2.VerifierOption(verifier.Value)) if err != nil { o.renderer.Text(w, http.StatusUnauthorized, loginFailedMessage) - o.logger.Error("failed to exchange authorization code", "error", err) + o.logger.Error("failed to exchange authorization code", logger.GetSlogAttrFromError(err)) return } @@ -294,7 +310,7 @@ func (o *OidcAuth) Callback(w http.ResponseWriter, r *http.Request) { idToken, err := o.verifier.Verify(coreoidc.ClientContext(r.Context(), o.httpClient), rawIDToken) if err != nil { o.renderer.Text(w, http.StatusUnauthorized, loginFailedMessage) - o.logger.Error("failed to verify id token", "error", err) + o.logger.Error("failed to verify id token", logger.GetSlogAttrFromError(err)) return } @@ -307,7 +323,7 @@ func (o *OidcAuth) Callback(w http.ResponseWriter, r *http.Request) { user, err := o.claimsToUser(idToken, true) if err != nil { o.renderer.Text(w, http.StatusUnauthorized, loginFailedMessage) - o.logger.Error("failed to read id token claims", "error", err) + o.logger.Error("failed to read id token claims", logger.GetSlogAttrFromError(err)) return } @@ -328,7 +344,7 @@ func (o *OidcAuth) Callback(w http.ResponseWriter, r *http.Request) { } o.clearCookie(w, returnCookie) - o.logger.Info("logged in", "user", user.Email, "roles", user.Roles) + o.logger.Info("logged in", "user", user) http.Redirect(w, r, target, http.StatusFound) } diff --git a/oidcauth/oidcauth_test.go b/oidcauth/oidcauth_test.go index 1e0e3a3..3c5c2e0 100644 --- a/oidcauth/oidcauth_test.go +++ b/oidcauth/oidcauth_test.go @@ -2,10 +2,13 @@ package oidcauth_test import ( "context" + "encoding/json" + "log/slog" "net/http" "net/http/cookiejar" "net/http/httptest" "net/url" + "strings" "testing" "github.com/cego/go-lib/v2/logger" @@ -565,6 +568,36 @@ func TestOidcAuth(t *testing.T) { assert.Equal(t, "Mads Nielsen mjn", string(body[:n])) }) + t.Run("the user logs as an ecs shaped group", func(t *testing.T) { + user := oidcauth.User{ + Subject: "cb18311d", PreferredUsername: "mjn", Email: "mjn@cego.dk", + Name: "Mads Jon Nielsen", Roles: []string{"reader", "process-admin"}, + } + + var out strings.Builder + slog.New(slog.NewJSONHandler(&out, nil)).Info("logged in", "user", user) + + logged := map[string]any{} + require.NoError(t, json.Unmarshal([]byte(out.String()), &logged)) + fields, ok := logged["user"].(map[string]any) + require.True(t, ok, out.String()) + + assert.Equal(t, "mjn", fields["id"]) + assert.Equal(t, "oidc", fields["type"]) + assert.Equal(t, "mjn@cego.dk", fields["email"]) + assert.Equal(t, "Mads Jon Nielsen", fields["full_name"]) + assert.Equal(t, []any{"reader", "process-admin"}, fields["roles"]) + }) + + t.Run("the log id falls back to the subject", func(t *testing.T) { + var out strings.Builder + slog.New(slog.NewJSONHandler(&out, nil)).Info("logged in", "user", oidcauth.User{Subject: "cb18311d"}) + + logged := map[string]any{} + require.NoError(t, json.Unmarshal([]byte(out.String()), &logged)) + assert.Equal(t, "cb18311d", logged["user"].(map[string]any)["id"]) + }) + t.Run("discovery failure is reported", func(t *testing.T) { _, err := oidcauth.New(context.Background(), logger.NewMock(), oidcauth.Config{ Issuer: "https://127.0.0.1:1/realms/nope", diff --git a/readme.md b/readme.md index c7ad377..b764fee 100644 --- a/readme.md +++ b/readme.md @@ -118,6 +118,7 @@ user.HasAnyRole("reader", "tool-admin") - Logout is a `POST`, and hands the provider no token, so it asks the user to confirm - Sessions cannot be revoked before the id token expires, so keep that lifetime short - Cookies are `__Host-` prefixed, secure and `SameSite=Lax`, which is not a CSRF token +- `User` implements `slog.LogValuer`, so `"user", user` logs the ecs shaped group ### Options ```go