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..52525f1 --- /dev/null +++ b/oidcauth/oidcauth.go @@ -0,0 +1,519 @@ +package oidcauth + +import ( + "context" + "crypto/rand" + "encoding/base64" + "errors" + "fmt" + "log/slog" + "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) 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) +} + +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, 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 + } + + 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", 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", logger.GetSlogAttrFromError(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", logger.GetSlogAttrFromError(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", logger.GetSlogAttrFromError(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", logger.GetSlogAttrFromError(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) + 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..3c5c2e0 --- /dev/null +++ b/oidcauth/oidcauth_test.go @@ -0,0 +1,628 @@ +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" + "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("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", + 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..b764fee 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,54 @@ 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 +- `User` implements `slog.LogValuer`, so `"user", user` logs the ecs shaped group + +### 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)