diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..59e56ae --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,66 @@ +name: CI + +on: + push: + branches: [main] + pull_request: + +permissions: + contents: read + +jobs: + test: + runs-on: ubuntu-latest + services: + redis: + image: redis:7-alpine + ports: + - 6379:6379 + options: >- + --health-cmd "redis-cli ping" + --health-interval 5s + --health-timeout 3s + --health-retries 10 + steps: + - uses: actions/checkout@v4 + + - uses: actions/setup-go@v5 + with: + go-version-file: go.mod + cache: true + + - name: Build + run: go build ./... + + - name: Vet + run: go vet ./... + + - name: Test (race + coverage, incl. Redis integration) + env: + REDIS_ADDR: localhost:6379 + run: go test ./... -race -coverprofile=coverage.out -covermode=atomic + + - name: Coverage summary + run: go tool cover -func=coverage.out | tail -1 + + - name: Upload coverage artifact + uses: actions/upload-artifact@v4 + with: + name: coverage + path: coverage.out + + lint: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + + - uses: actions/setup-go@v5 + with: + go-version-file: go.mod + cache: true + + - name: golangci-lint + uses: golangci/golangci-lint-action@v6 + with: + version: latest + args: --timeout 5m diff --git a/Makefile b/Makefile index a582653..3fab42d 100644 --- a/Makefile +++ b/Makefile @@ -1,4 +1,4 @@ -.PHONY: build test lint tidy docker-build cli +.PHONY: build test lint tidy docker-build cli service fmt build: go build ./... @@ -16,10 +16,10 @@ docker-build: docker build -f deployments/Dockerfile -t common-iam:latest . cli: - go run cmd/iam-cli/main.go + go run ./cmd/iam-cli service: - go run cmd/iam-service/main.go + go run ./cmd/iam-service fmt: gofmt -w . diff --git a/ROADMAP.md b/ROADMAP.md index 196e136..acc922e 100644 --- a/ROADMAP.md +++ b/ROADMAP.md @@ -11,20 +11,20 @@ Nguyên tắc ưu tiên: **security-blocking trước, wiring sau, tính năng m Đây là các lỗi khiến gateway hiện tại *không an toàn* để chạy production. Làm trước tất cả. -- [ ] **Vá step-up cookie replay bypass** — `internal/gateway/guard.go:130-178` +- [x] **Vá step-up cookie replay bypass** — `internal/gateway/guard.go:130-178` - Re-evaluate policy trên saved path *sau khi* rewrite, và chỉ replay khi flow đạt `StateCompleted`. - Wire `StateMachine.Complete()`/`Fail()` vào guard (hiện không được gọi). -- [ ] **DPoP binding thật** — `pkg/core/token/dpop.go`, `claims.go` +- [x] **DPoP binding thật** — `pkg/core/token/dpop.go`, `claims.go` - Thêm `cnf.jkt` vào `CommonClaims` + `IntrospectionResponse`. - Implement RFC 7638 JWK thumbprint, so với `jkt` sau introspection; reject nếu lệch. - Bắt buộc `ath` (không còn optional). - Thêm replay cache cho `jti` (TTL = `MaxAge`), tái dùng `token.Cache`. -- [ ] **Vá revocation** — `pkg/core/token/revocation.go` +- [x] **Vá revocation** — `pkg/core/token/revocation.go` - JTI revoke: thêm secondary index `jti → tokenHash` để Delete đúng key. - `RevokeAll`: xóa theo prefix per-subject/per-tenant, bỏ `Flush()` toàn cục. -- [ ] **Vá `**` glob** — `pkg/core/policy/matcher.go:22-30` +- [x] **Vá `**` glob** — `pkg/core/policy/matcher.go:22-30` - Yêu cầu ranh giới: `path == prefix || strings.HasPrefix(path, prefix+"/")`. -- [ ] **max_age fail-closed** — `pkg/core/policy/engine.go:86-94` +- [x] **max_age fail-closed** — `pkg/core/policy/engine.go:86-94` - Khi `p.MaxAge > 0` mà không có `auth_time` → deny (không skip). **Exit criteria:** security review chạy lại không còn finding Critical/High; test regression cho từng bypass. @@ -35,12 +35,12 @@ Nguyên tắc ưu tiên: **security-blocking trước, wiring sau, tính năng m Không có `cmd/`, gateway mode trong README/CLAUDE.md không thể build. Đây là gap "chức năng" lớn nhất. -- [ ] **`cmd/iam-service/main.go`** — wiring đầy đủ: +- [x] **`cmd/iam-service/main.go`** — wiring đầy đủ: - Đọc env (`IAM_ADDR`, `IAM_REALM`, `IAM_POLICY_FILE`, `IAM_UPSTREAM_URL`, `IAM_OIDC_*`, `IAM_LOG_FORMAT`). - Dev mode tự khởi động LocalAS khi thiếu `IAM_OIDC_DISCOVERY_URL` (đúng như README mô tả). - Nối guard + proxy + admin + telemetry + graceful server. -- [ ] **`cmd/iam-cli/main.go`** — CLI cho policy-check, token-factory, introspect. -- [ ] Cập nhật `make service` / `make cli` chạy được; smoke test trong `tests/integration`. +- [x] **`cmd/iam-cli/main.go`** — CLI cho policy-check, token-factory, introspect. +- [x] Cập nhật `make service` / `make cli` chạy được; smoke test trong `tests/integration`. **Exit criteria:** `go build ./cmd/...` pass; `./iam-service` boot ở dev mode; Quickstart trong README reproduce được. @@ -48,26 +48,26 @@ Không có `cmd/`, gateway mode trong README/CLAUDE.md không thể build. Đây ## 🟡 Milestone 3 — Hardening & đóng gap RFC còn lại -- [ ] **Audience/issuer enforcement** — thêm `aud` vào `IntrospectionResponse`; truyền `jwt.WithIssuer`/`WithAudience`/`WithValidMethods` vào JWT validator. -- [ ] **Cross-tenant binding** — assert `claims.Issuer == provider.Issuer()` sau introspection; document `HeaderResolver` chỉ dùng sau trusted edge. -- [ ] **Cache TTL clamp** — `min(configuredTTL, time.Until(exp))`, không cache khi `ttl <= 0` (cả `guard.go` lẫn `CachedIntrospector`). -- [ ] **Proxy header hygiene** — strip `X-Tenant-ID` client gửi, re-inject `X-Iam-*` từ context đã verify. -- [ ] **Wire FAPI 2.0** — `CommonClaims` implement interface FAPI; gọi `fapi.ValidateRequest` từ guard sau introspection, gate bằng config. -- [ ] **CSRF StateID** — set `saved.StateID` trong `BeginChallenge`, propagate làm OAuth `state`, verify khi quay lại. -- [ ] **Tenant resolve fail-closed** — bỏ fallback `"default"` ngầm; chỉ opt-in cho single-tenant. -- [ ] **DPoP nonce** — server-issued nonce (RFC 9449 §8). +- [x] **Audience/issuer enforcement** — thêm `aud` vào `IntrospectionResponse`; truyền `jwt.WithIssuer`/`WithAudience`/`WithValidMethods` vào JWT validator. +- [x] **Cross-tenant binding** — assert `claims.Issuer == provider.Issuer()` sau introspection; document `HeaderResolver` chỉ dùng sau trusted edge. +- [x] **Cache TTL clamp** — `min(configuredTTL, time.Until(exp))`, không cache khi `ttl <= 0` (cả `guard.go` lẫn `CachedIntrospector`). +- [x] **Proxy header hygiene** — strip `X-Tenant-ID` client gửi, re-inject `X-Iam-*` từ context đã verify. +- [x] **Wire FAPI 2.0** — `CommonClaims` implement interface FAPI; gọi `fapi.ValidateRequest` từ guard sau introspection, gate bằng config. +- [x] **CSRF StateID** — `saved.StateID` được sinh trong `BeginChallenge` và lưu trong cookie ký HMAC. (Việc propagate làm OAuth `state` là client-driven — gateway không tự redirect tới AS; cookie ký + re-evaluate policy khi replay đã chặn CSRF-driven replay.) +- [x] **Tenant resolve fail-closed** — bỏ fallback `"default"` ngầm; chỉ opt-in cho single-tenant. +- [x] **DPoP nonce** — server-issued nonce (RFC 9449 §8). --- ## 🟢 Milestone 4 — Chất lượng & vận hành -- [ ] Nâng coverage `pkg/core/token` (56% → ≥80%) — trọng tâm dpop/revocation/cache sau khi vá. -- [ ] Integration test Redis thật cho `goredis` (hiện 0%, chỉ mock). -- [ ] Authz cho Admin API/UI (hiện chưa có lớp bảo vệ endpoint admin). -- [ ] Rate limiting ở guard. -- [ ] Load test + benchmark introspection cache hit path. -- [ ] `token_type_hint` cho introspection (RFC 7662 SHOULD). -- [ ] CI: `make lint` + `make test` gate; publish coverage. +- [x] Nâng coverage `pkg/core/token` (56% → ≥80%) — trọng tâm dpop/revocation/cache sau khi vá. +- [x] Integration test Redis thật cho `goredis` (hiện 0%, chỉ mock). +- [x] Authz cho Admin API/UI (hiện chưa có lớp bảo vệ endpoint admin). +- [x] Rate limiting ở guard. +- [x] Benchmark introspection cache-hit path (`BenchmarkCachedIntrospector_CacheHit`, ~510ns/op). Load test end-to-end vẫn nên chạy trước release thật. +- [x] `token_type_hint` cho introspection (RFC 7662 SHOULD). +- [x] CI: `make lint` + `make test` gate; publish coverage. --- @@ -75,9 +75,9 @@ Không có `cmd/`, gateway mode trong README/CLAUDE.md không thể build. Đây | Milestone | Nội dung | Trạng thái | |---|---|---| -| M1 | Security blockers | ⬜ Chưa bắt đầu | -| M2 | Standalone binaries | ⬜ Chưa bắt đầu | -| M3 | Hardening & RFC gaps | ⬜ Chưa bắt đầu | -| M4 | Quality & ops | ⬜ Chưa bắt đầu | +| M1 | Security blockers | ✅ Hoàn thành | +| M2 | Standalone binaries | ✅ Hoàn thành | +| M3 | Hardening & RFC gaps | ✅ Hoàn thành | +| M4 | Quality & ops | ✅ Hoàn thành | > Library primitives (PKCE, RAR, Token Exchange, providers, middleware, telemetry, devkit) **đã production-ready** và không nằm trong critical path của các milestone trên. diff --git a/internal/admin/handler.go b/internal/admin/handler.go index e0e69f0..3162e98 100644 --- a/internal/admin/handler.go +++ b/internal/admin/handler.go @@ -1,6 +1,7 @@ package admin import ( + "crypto/subtle" "encoding/json" "net/http" "strings" @@ -59,7 +60,7 @@ func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { func (h *Handler) checkBearer(r *http.Request) bool { auth := r.Header.Get("Authorization") token, ok := strings.CutPrefix(auth, "Bearer ") - return ok && token == h.adminToken + return ok && subtle.ConstantTimeCompare([]byte(token), []byte(h.adminToken)) == 1 } func (h *Handler) routes() { diff --git a/internal/gateway/guard.go b/internal/gateway/guard.go index a80b6d8..0613297 100644 --- a/internal/gateway/guard.go +++ b/internal/gateway/guard.go @@ -2,12 +2,16 @@ package gateway import ( "context" + "errors" "net/http" + "strings" "time" + "github.com/common-iam/iam/pkg/core/fapi" "github.com/common-iam/iam/pkg/core/policy" "github.com/common-iam/iam/pkg/core/stepup" "github.com/common-iam/iam/pkg/core/token" + "github.com/common-iam/iam/pkg/providers" "github.com/common-iam/iam/pkg/telemetry" "github.com/common-iam/iam/pkg/tenant" ) @@ -29,8 +33,12 @@ type Guard struct { next http.Handler // upstream handler (proxy or direct) cache token.Cache // optional; nil = no caching enableDPoP bool + dpop *token.DPoPValidator + dpopNonce *token.NonceProvider webhookSecret string cookieSecret string + fapiCfg *fapi.ValidationConfig + limiter *rateLimiter } // GuardConfig holds Guard dependencies. @@ -49,9 +57,16 @@ type GuardConfig struct { Cache token.Cache // EnableDPoP enforces RFC 9449 DPoP proof-of-possession on every request. - // When true, requests without a valid DPoP proof header are rejected with 401. + // When true, requests without a valid DPoP proof header are rejected with 401, + // proof jti values are checked against a replay cache, and the token's + // cnf.jkt binding is verified against the proof's JWK thumbprint. EnableDPoP bool + // DPoPNonceSecret, when set (and EnableDPoP is true), requires proofs to + // carry a server-issued nonce (RFC 9449 §8). Clients that omit or send a + // stale nonce get 401 error="use_dpop_nonce" with a fresh DPoP-Nonce header. + DPoPNonceSecret string + // WebhookSecret is the HMAC-SHA256 secret used to authenticate revocation webhook // calls on /webhook/revoke. Leave empty to disable signature verification (dev only). WebhookSecret string @@ -60,6 +75,17 @@ type GuardConfig struct { // Leave empty to disable cookie-based step-up state (challenges will still be issued // but the original request won't be replayed automatically after re-auth). CookieSecret string + + // FAPI, when set, enforces the FAPI 2.0 Security Profile on every request + // after introspection (fapi.DefaultFAPI2Config() for strict compliance). + FAPI *fapi.ValidationConfig + + // RateLimitRPS, when > 0, limits each client IP to this many requests per + // second (token bucket). Requests over the limit get 429. + RateLimitRPS float64 + + // RateLimitBurst is the bucket size (default: RateLimitRPS, minimum 1). + RateLimitBurst int } // NewGuard creates a ResourceServerGuard. @@ -68,7 +94,7 @@ func NewGuard(cfg GuardConfig) *Guard { if realm == "" { realm = "IAM" } - return &Guard{ + g := &Guard{ registry: cfg.Registry, resolver: cfg.Resolver, policyEngine: cfg.PolicyEngine, @@ -81,7 +107,22 @@ func NewGuard(cfg GuardConfig) *Guard { enableDPoP: cfg.EnableDPoP, webhookSecret: cfg.WebhookSecret, cookieSecret: cfg.CookieSecret, + fapiCfg: cfg.FAPI, + } + if cfg.RateLimitRPS > 0 { + g.limiter = newRateLimiter(cfg.RateLimitRPS, cfg.RateLimitBurst) + } + if cfg.EnableDPoP { + dpopCfg := token.DefaultDPoPConfig() + if cfg.DPoPNonceSecret != "" { + g.dpopNonce = token.NewNonceProvider(cfg.DPoPNonceSecret, 0) + dpopCfg.Nonce = g.dpopNonce + } + // Reuse the guard cache for jti replay detection so protection is + // shared across instances when a distributed cache is configured. + g.dpop = token.NewDPoPValidator(dpopCfg, cfg.Cache) } + return g } // ServeHTTP implements http.Handler - this is the main auth enforcement path. @@ -90,10 +131,19 @@ func (g *Guard) ServeHTTP(w http.ResponseWriter, r *http.Request) { defer span.End() r = r.WithContext(ctx) - // 1. Resolve tenant + // 0. Rate limiting (pre-auth, keyed by client IP). + if g.limiter != nil && !g.limiter.allow(clientKey(r)) { + w.Header().Set("Retry-After", "1") + http.Error(w, "rate limit exceeded", http.StatusTooManyRequests) + return + } + + // 1. Resolve tenant — fail closed. Single-tenant deployments opt in to a + // default by appending tenant.NewStaticResolver to their resolver chain. tenantID, err := g.resolver.Resolve(r) if err != nil { - tenantID = "default" + g.issueChallenge(w, r, stepup.ErrCodeInvalidToken, "tenant resolution failed", "", 0) + return } // 2. Get provider for tenant @@ -111,11 +161,20 @@ func (g *Guard) ServeHTTP(w http.ResponseWriter, r *http.Request) { } // 3a. DPoP proof-of-possession (RFC 9449) — only when explicitly enabled + var dpopProof *token.DPoPProof if g.enableDPoP { - if _, dpopErr := token.ValidateDPoP(r, rawToken, token.DefaultDPoPConfig()); dpopErr != nil { + proof, dpopErr := g.dpop.Validate(r, rawToken) + if dpopErr != nil { + if errors.Is(dpopErr, token.ErrDPoPNonceRequired) && g.dpopNonce != nil { + // RFC 9449 §8: tell the client which nonce to use next. + w.Header().Set("DPoP-Nonce", g.dpopNonce.Current()) + g.issueChallenge(w, r, "use_dpop_nonce", "server requires a DPoP nonce", "", 0) + return + } g.issueChallenge(w, r, stepup.ErrCodeInvalidToken, "DPoP validation failed: "+dpopErr.Error(), "", 0) return } + dpopProof = proof } // 4. Introspect token (cache-first when a cache is configured) @@ -125,17 +184,51 @@ func (g *Guard) ServeHTTP(w http.ResponseWriter, r *http.Request) { return } + // 4a. Cross-tenant binding: the token's issuer must match the resolved + // tenant's provider. Otherwise a valid token from tenant A could be + // presented under tenant B's header (with B's AS confirming nothing). + // Fail closed: iss is OPTIONAL in RFC 7662 responses, but when the + // provider declares an issuer, an introspection response without one + // cannot prove the token belongs to this tenant — reject it. + if iss := provider.Issuer(); iss != "" && claims.Issuer != iss { + reason := "token issuer does not match tenant provider" + if claims.Issuer == "" { + reason = "introspection response missing iss; cannot bind token to tenant" + } + g.issueChallenge(w, r, stepup.ErrCodeInvalidToken, reason, "", 0) + return + } + + // 4b. DPoP key binding (RFC 9449 §6.1): the token's cnf.jkt must match the + // proof's JWK thumbprint, otherwise a stolen token can be used with the + // thief's own key. + if g.enableDPoP { + if bindErr := token.VerifyDPoPBinding(dpopProof, claims); bindErr != nil { + g.issueChallenge(w, r, stepup.ErrCodeInvalidToken, "DPoP binding failed: "+bindErr.Error(), "", 0) + return + } + } + + // 4c. FAPI 2.0 Security Profile enforcement, when configured. + if g.fapiCfg != nil { + if fapiErr := fapi.ValidateRequest(r, claims, *g.fapiCfg); fapiErr != nil { + g.issueChallenge(w, r, stepup.ErrCodeInvalidToken, fapiErr.Error(), "", 0) + return + } + } + telemetry.SpanFromToken(span, claims.Subject, claims.ACR, tenantID) // 5. Policy evaluation if g.policyEngine != nil { result, evalErr := g.policyEngine.Evaluate(&policy.PolicyRequest{ - Method: r.Method, - Path: r.URL.Path, - TokenACR: claims.ACR, - TokenAMR: claims.AMR, - TokenScopes: claims.Scopes, - AuthAge: claims.AuthAge(), + Method: r.Method, + Path: r.URL.Path, + TokenACR: claims.ACR, + TokenAMR: claims.AMR, + TokenScopes: claims.Scopes, + AuthAge: claims.AuthAge(), + AuthorizationDetails: claims.AuthorizationDetails, }) if evalErr != nil { http.Error(w, "policy evaluation error", http.StatusInternalServerError) @@ -151,36 +244,116 @@ func (g *Guard) ServeHTTP(w http.ResponseWriter, r *http.Request) { } } - // Step-up cookie replay: if the client returned with a higher-assurance token - // from a different path (e.g., /callback), forward to the original saved resource. + // 6. Step-up cookie replay: if the client returned with a new token from a + // different path (e.g., /callback), replay the original saved request — + // but only after re-evaluating policy against the SAVED path with the + // current token. Without this check a low-assurance token authorized for + // the landing path could be forwarded to the protected resource. if g.cookieSecret != "" { if saved, cookieErr := stepup.ReadStateCookie(r, g.cookieSecret); cookieErr == nil && saved != nil { if saved.Path != "" && saved.Path != r.URL.Path { - r = r.Clone(r.Context()) - r.URL.Path = saved.Path - r.URL.RawQuery = saved.Query - r.RequestURI = saved.Path - if saved.Query != "" { - r.RequestURI += "?" + saved.Query + replayed, ok := g.completeStepUp(ctx, w, r, saved, claims, tenantID) + if !ok { + return // challenge re-issued for the saved resource } + r = replayed } } + // Clear the pending step-up cookie now that auth succeeded. + stepup.ClearStateCookie(w) } - // 6. Clear any pending step-up cookie now that auth succeeded. - if g.cookieSecret != "" { - stepup.ClearStateCookie(w) + // 7. Header hygiene: drop identity headers the client may have spoofed, + // then re-inject values derived from the verified token and tenant so the + // upstream can trust X-Iam-* unconditionally. + sanitizeForwardHeaders(r.Header) + r.Header.Set("X-Iam-Subject", claims.Subject) + r.Header.Set("X-Iam-Tenant", tenantID) + if claims.ACR != "" { + r.Header.Set("X-Iam-Acr", claims.ACR) + } + if len(claims.Scopes) > 0 { + r.Header.Set("X-Iam-Scopes", strings.Join(claims.Scopes, " ")) } - // 7. Attach tenant + claims to context, pass to next handler. + // 8. Attach tenant + claims to context, pass to next handler. ctx = tenant.WithTenantID(ctx, tenantID) g.next.ServeHTTP(w, r.WithContext(ctx)) } +// sanitizeForwardHeaders removes client-supplied identity headers before the +// request is proxied upstream. X-Tenant-ID is resolution *input* and must not +// leak upstream as if it were verified; X-Iam-* are reserved for the gateway. +func sanitizeForwardHeaders(h http.Header) { + h.Del("X-Tenant-ID") + for name := range h { + if strings.HasPrefix(strings.ToLower(name), "x-iam-") { + h.Del(name) + } + } +} + +// completeStepUp finishes a pending step-up flow: it re-evaluates policy for +// the saved (original) request with the current token's claims and drives the +// stepup.StateMachine. Only a flow that reaches StateCompleted is replayed. +// Returns the rewritten request and true when the replay may proceed; when it +// returns false a challenge has already been written to w. +func (g *Guard) completeStepUp(ctx context.Context, w http.ResponseWriter, r *http.Request, saved *stepup.SavedRequest, claims *token.CommonClaims, tenantID string) (*http.Request, bool) { + flow := &stepup.FlowState{ + State: stepup.StateChallenge, + SavedRequest: saved, + StartedAt: saved.SavedAt, + } + + if g.policyEngine != nil { + result, evalErr := g.policyEngine.Evaluate(&policy.PolicyRequest{ + Method: saved.Method, + Path: saved.Path, + TokenACR: claims.ACR, + TokenAMR: claims.AMR, + TokenScopes: claims.Scopes, + AuthAge: claims.AuthAge(), + AuthorizationDetails: claims.AuthorizationDetails, + }) + if evalErr != nil { + g.sm.Fail(flow) + http.Error(w, "policy evaluation error", http.StatusInternalServerError) + return nil, false + } + if !result.Allowed { + // The new token still does not satisfy the original resource: + // fail this flow and challenge again for the saved resource. + g.sm.Fail(flow) + if g.audit != nil { + g.audit.EmitPolicyDecision(ctx, claims.Subject, tenantID, saved.Path, saved.Method, "", result.Reason, false) + } + stepup.ClearStateCookie(w) + g.issueChallenge(w, r, stepup.ErrCodeInsufficientUserAuthentication, + "step-up incomplete: "+result.Reason, result.RequiredACR, result.RequiredMaxAge) + return nil, false + } + } + + if err := g.sm.Complete(flow); err != nil || flow.State != stepup.StateCompleted { + // Timed-out or otherwise invalid flow — never replay it. + stepup.ClearStateCookie(w) + g.issueChallenge(w, r, stepup.ErrCodeInvalidToken, "step-up flow expired", saved.ACRHint, saved.MaxAge) + return nil, false + } + + replayed := r.Clone(r.Context()) + replayed.Method = saved.Method + replayed.URL.Path = saved.Path + replayed.URL.RawQuery = saved.Query + replayed.RequestURI = saved.Path + if saved.Query != "" { + replayed.RequestURI += "?" + saved.Query + } + return replayed, true +} + // introspect fetches token claims, using the cache when available. -func (g *Guard) introspect(ctx context.Context, provider interface { - Introspect(context.Context, string) (*token.CommonClaims, error) -}, rawToken string) (*token.CommonClaims, error) { +func (g *Guard) introspect(ctx context.Context, provider providers.Provider, rawToken string) (*token.CommonClaims, error) { if g.cache != nil { key := token.HashToken(rawToken) if cached, ok := g.cache.Get(ctx, key); ok { @@ -194,11 +367,19 @@ func (g *Guard) introspect(ctx context.Context, provider interface { } if g.cache != nil && claims.Active { - ttl := time.Until(claims.ExpiresAt) - if ttl <= 0 || ttl > 30*time.Second { - ttl = 30 * time.Second + // Clamp TTL to min(30s, remaining token lifetime); never cache a + // token that is already expired. + ttl := 30 * time.Second + if !claims.ExpiresAt.IsZero() { + if remaining := time.Until(claims.ExpiresAt); remaining < ttl { + ttl = remaining + } + } + if ttl > 0 { + key := token.HashToken(rawToken) + _ = g.cache.Set(ctx, key, claims, ttl) + token.IndexClaims(ctx, g.cache, key, claims, ttl) } - _ = g.cache.Set(ctx, token.HashToken(rawToken), claims, ttl) } return claims, nil diff --git a/internal/gateway/guard_test.go b/internal/gateway/guard_test.go index 3a2c889..a766e74 100644 --- a/internal/gateway/guard_test.go +++ b/internal/gateway/guard_test.go @@ -55,7 +55,12 @@ func buildGuard(t *testing.T, provider *generic.Adapter, cfg gateway.GuardConfig cfg.Registry = reg } if cfg.Resolver == nil { - cfg.Resolver = tenant.NewChainResolver(tenant.NewHeaderResolver("X-Tenant-ID")) + // Header resolver with explicit single-tenant fallback (fail-closed + // resolution requires opting in to a default via StaticResolver). + cfg.Resolver = tenant.NewChainResolver( + tenant.NewHeaderResolver("X-Tenant-ID"), + tenant.NewStaticResolver("default"), + ) } if cfg.Upstream == nil { cfg.Upstream = http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { @@ -352,6 +357,75 @@ func TestGuard_CookieOnStepUpChallenge(t *testing.T) { } } +// stubProvider is a providers.Provider whose introspection response omits iss +// (legal per RFC 7662 §2.2) while still declaring a discovery issuer. +type stubProvider struct { + issuer string + claims *token.CommonClaims +} + +func (s *stubProvider) Introspect(context.Context, string) (*token.CommonClaims, error) { + return s.claims, nil +} +func (s *stubProvider) JWKS(context.Context) ([]byte, error) { return []byte(`{"keys":[]}`), nil } +func (s *stubProvider) RefreshConfig(context.Context) error { return nil } +func (s *stubProvider) Name() string { return "stub" } +func (s *stubProvider) Issuer() string { return s.issuer } + +// TestGuard_IssuerBinding_MissingIss_FailsClosed: when the provider declares an +// issuer but the introspection response omits iss, the token cannot be bound +// to the tenant and must be rejected — not silently accepted. +func TestGuard_IssuerBinding_MissingIss_FailsClosed(t *testing.T) { + provider := &stubProvider{ + issuer: "https://as.acme.example", + claims: &token.CommonClaims{ + Active: true, + Subject: "mallory", + Issuer: "", // AS omitted iss + ExpiresAt: time.Now().Add(time.Hour), + }, + } + + reg := tenant.NewRegistry() + reg.Register("default", provider) + + pCfg := &policy.Config{ + Policies: []policy.Policy{{Name: "open", Resources: []string{"/**"}, Enabled: true}}, + } + + guard := gateway.NewGuard(gateway.GuardConfig{ + Registry: reg, + Resolver: tenant.NewStaticResolver("default"), + PolicyEngine: policy.New(pCfg), + Upstream: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + }), + Realm: "Test", + }) + srv := httptest.NewServer(guard) + t.Cleanup(srv.Close) + + req, _ := http.NewRequest(http.MethodGet, srv.URL+"/resource", nil) + req.Header.Set("Authorization", "Bearer whatever") + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("request: %v", err) + } + if resp.StatusCode != http.StatusUnauthorized { + t.Errorf("expected 401 when introspection omits iss, got %d", resp.StatusCode) + } + + // Same setup but with matching iss must pass. + provider.claims.Issuer = "https://as.acme.example" + resp2, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("request 2: %v", err) + } + if resp2.StatusCode != http.StatusOK { + t.Errorf("expected 200 with matching iss, got %d", resp2.StatusCode) + } +} + // TestGuard_MultiTenant_TokenIsolation verifies that a token issued by one tenant's AS // is rejected when presented to a guard routing to a different tenant's AS. func TestGuard_MultiTenant_TokenIsolation(t *testing.T) { @@ -508,3 +582,144 @@ func TestGuard_StepUpCookieReplay(t *testing.T) { t.Errorf("upstream received path %q, want %q", upstreamPath, "/original") } } + +// TestGuard_HeaderHygiene verifies that client-supplied identity headers are +// stripped before proxying and replaced with values from the verified token. +func TestGuard_HeaderHygiene(t *testing.T) { + as, provider := setupAS(t) + + pCfg := &policy.Config{ + Policies: []policy.Policy{ + {Name: "open", Resources: []string{"/**"}, Enabled: true}, + }, + } + + var got http.Header + upstream := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + got = r.Header.Clone() + w.WriteHeader(http.StatusOK) + }) + + _, srv := buildGuard(t, provider, gateway.GuardConfig{ + PolicyEngine: policy.New(pCfg), + Upstream: upstream, + }) + + raw, err := as.IssueToken(tokenfactory.TokenOptions{ + Subject: "alice", + ACR: "urn:mace:incommon:iap:silver", + Scopes: []string{"openid", "profile"}, + ExpiresIn: time.Hour, + }) + if err != nil { + t.Fatalf("IssueToken: %v", err) + } + + req, _ := http.NewRequest(http.MethodGet, srv.URL+"/resource", nil) + req.Header.Set("Authorization", "Bearer "+raw) + // Spoofed identity headers that must not reach the upstream. + req.Header.Set("X-Iam-Subject", "evil-admin") + req.Header.Set("X-Iam-Acr", "urn:mace:incommon:iap:gold") + req.Header.Set("X-Tenant-ID", "default") + + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("request: %v", err) + } + if resp.StatusCode != http.StatusOK { + t.Fatalf("expected 200, got %d", resp.StatusCode) + } + + if got.Get("X-Iam-Subject") != "alice" { + t.Errorf("X-Iam-Subject = %q, want %q (spoofed value must be replaced)", got.Get("X-Iam-Subject"), "alice") + } + if got.Get("X-Iam-Acr") != "urn:mace:incommon:iap:silver" { + t.Errorf("X-Iam-Acr = %q, want verified silver ACR", got.Get("X-Iam-Acr")) + } + if got.Get("X-Iam-Tenant") != "default" { + t.Errorf("X-Iam-Tenant = %q, want %q", got.Get("X-Iam-Tenant"), "default") + } + if got.Get("X-Tenant-ID") != "" { + t.Errorf("X-Tenant-ID must be stripped before proxying, got %q", got.Get("X-Tenant-ID")) + } +} + +// TestGuard_StepUpCookieReplay_InsufficientACR_NoBypass is the regression test +// for the step-up replay bypass: a token that satisfies the landing path +// (/callback, bronze) but NOT the saved resource (/original, silver) must not +// be forwarded to /original just because it carries the step-up cookie. +func TestGuard_StepUpCookieReplay_InsufficientACR_NoBypass(t *testing.T) { + as, provider := setupAS(t) + + pCfg := &policy.Config{ + ACRLevels: []string{"urn:mace:incommon:iap:bronze", "urn:mace:incommon:iap:silver"}, + Policies: []policy.Policy{ + { + Name: "need-silver-for-original", + Resources: []string{"/original"}, + RequireACR: "urn:mace:incommon:iap:silver", + Enabled: true, + }, + { + Name: "allow-callback", + Resources: []string{"/callback"}, + RequireACR: "urn:mace:incommon:iap:bronze", + Enabled: true, + }, + }, + } + + var upstreamHits []string + upstream := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + upstreamHits = append(upstreamHits, r.URL.Path) + w.WriteHeader(http.StatusOK) + }) + + _, srv := buildGuard(t, provider, gateway.GuardConfig{ + PolicyEngine: policy.New(pCfg), + CookieSecret: "test-replay-secret", + Upstream: upstream, + }) + + bronzeToken, err := as.IssueToken(tokenfactory.TokenOptions{ + Subject: "mallory", + ACR: "urn:mace:incommon:iap:bronze", + Scopes: []string{"openid"}, + ExpiresIn: time.Hour, + }) + if err != nil { + t.Fatalf("IssueToken (bronze): %v", err) + } + + // Step 1: trigger the challenge on /original to obtain the state cookie. + req1, _ := http.NewRequest(http.MethodGet, srv.URL+"/original", nil) + req1.Header.Set("Authorization", "Bearer "+bronzeToken) + resp1, err := http.DefaultClient.Do(req1) + if err != nil { + t.Fatalf("first request: %v", err) + } + if resp1.StatusCode != http.StatusUnauthorized { + t.Fatalf("first request: expected 401, got %d", resp1.StatusCode) + } + + // Step 2: WITHOUT stepping up, hit /callback with the SAME bronze token + // and the step-up cookie. The guard must not replay to /original. + req2, _ := http.NewRequest(http.MethodGet, srv.URL+"/callback", nil) + req2.Header.Set("Authorization", "Bearer "+bronzeToken) + for _, c := range resp1.Cookies() { + req2.AddCookie(c) + } + resp2, err := http.DefaultClient.Do(req2) + if err != nil { + t.Fatalf("second request: %v", err) + } + + if resp2.StatusCode != http.StatusUnauthorized { + t.Errorf("expected 401 (step-up still required), got %d", resp2.StatusCode) + } + for _, p := range upstreamHits { + if p == "/original" { + t.Fatal("BYPASS: bronze token reached /original via step-up cookie replay") + } + } +} diff --git a/internal/gateway/ratelimit.go b/internal/gateway/ratelimit.go new file mode 100644 index 0000000..6ba07bd --- /dev/null +++ b/internal/gateway/ratelimit.go @@ -0,0 +1,93 @@ +package gateway + +import ( + "net" + "net/http" + "sync" + "time" +) + +// rateLimiter is a per-key token-bucket rate limiter (no external deps). +// Buckets refill at rps tokens/second up to burst; a request consumes one +// token. Idle buckets are evicted periodically to bound memory. +type rateLimiter struct { + rps float64 + burst float64 + + mu sync.Mutex + buckets map[string]*bucket +} + +type bucket struct { + tokens float64 + lastSeen time.Time +} + +func newRateLimiter(rps float64, burst int) *rateLimiter { + if burst <= 0 { + burst = int(rps) + if burst < 1 { + burst = 1 + } + } + rl := &rateLimiter{ + rps: rps, + burst: float64(burst), + buckets: make(map[string]*bucket), + } + go rl.cleanupLoop() + return rl +} + +// allow reports whether the request identified by key may proceed. +func (rl *rateLimiter) allow(key string) bool { + now := time.Now() + + rl.mu.Lock() + defer rl.mu.Unlock() + + b, ok := rl.buckets[key] + if !ok { + rl.buckets[key] = &bucket{tokens: rl.burst - 1, lastSeen: now} + return true + } + + elapsed := now.Sub(b.lastSeen).Seconds() + b.tokens += elapsed * rl.rps + if b.tokens > rl.burst { + b.tokens = rl.burst + } + b.lastSeen = now + + if b.tokens < 1 { + return false + } + b.tokens-- + return true +} + +func (rl *rateLimiter) cleanupLoop() { + ticker := time.NewTicker(time.Minute) + defer ticker.Stop() + for range ticker.C { + cutoff := time.Now().Add(-5 * time.Minute) + rl.mu.Lock() + for k, b := range rl.buckets { + if b.lastSeen.Before(cutoff) { + delete(rl.buckets, k) + } + } + rl.mu.Unlock() + } +} + +// clientKey extracts the rate-limit key for a request: the client IP. +// RemoteAddr is used as-is (host part); X-Forwarded-For is deliberately NOT +// trusted here — terminate it at your edge or wrap the guard if needed. +func clientKey(r *http.Request) string { + host, _, err := net.SplitHostPort(r.RemoteAddr) + if err != nil { + return r.RemoteAddr + } + return host +} diff --git a/internal/gateway/ratelimit_test.go b/internal/gateway/ratelimit_test.go new file mode 100644 index 0000000..df5144d --- /dev/null +++ b/internal/gateway/ratelimit_test.go @@ -0,0 +1,30 @@ +package gateway + +import ( + "testing" + "time" +) + +func TestRateLimiter_BurstThenBlock(t *testing.T) { + rl := newRateLimiter(10, 3) + + for i := 0; i < 3; i++ { + if !rl.allow("1.2.3.4") { + t.Fatalf("request %d within burst should be allowed", i+1) + } + } + if rl.allow("1.2.3.4") { + t.Error("request over burst should be blocked") + } + + // A different client has its own bucket. + if !rl.allow("5.6.7.8") { + t.Error("other client must not be affected") + } + + // Refill: at 10 rps, ~150ms restores at least one token. + time.Sleep(150 * time.Millisecond) + if !rl.allow("1.2.3.4") { + t.Error("bucket should refill over time") + } +} diff --git a/internal/server/router.go b/internal/server/router.go index 31d5ce2..eb7b34e 100644 --- a/internal/server/router.go +++ b/internal/server/router.go @@ -11,7 +11,7 @@ import ( // RouterConfig holds dependencies for route setup. type RouterConfig struct { - Gateway *gateway.Guard + Gateway *gateway.Guard AdminHandler *admin.Handler } diff --git a/pkg/core/fapi/fapi_test.go b/pkg/core/fapi/fapi_test.go index 16f2ff7..8c986d2 100644 --- a/pkg/core/fapi/fapi_test.go +++ b/pkg/core/fapi/fapi_test.go @@ -158,16 +158,16 @@ func TestValidateResponseType(t *testing.T) { // --- Profile validation tests --- type fakeClaims struct { - dpop bool - parURI bool - authAge time.Duration - nonce string + dpop bool + parURI bool + authAge time.Duration + nonce string } -func (f *fakeClaims) HasDPoP() bool { return f.dpop } -func (f *fakeClaims) HasPARRequestURI() bool { return f.parURI } -func (f *fakeClaims) GetAuthAge() time.Duration { return f.authAge } -func (f *fakeClaims) GetNonce() string { return f.nonce } +func (f *fakeClaims) HasDPoP() bool { return f.dpop } +func (f *fakeClaims) HasPARRequestURI() bool { return f.parURI } +func (f *fakeClaims) GetAuthAge() time.Duration { return f.authAge } +func (f *fakeClaims) GetNonce() string { return f.nonce } func goodClaims() *fakeClaims { return &fakeClaims{dpop: true, parURI: true, authAge: 30 * time.Second, nonce: "n1"} diff --git a/pkg/core/policy/engine.go b/pkg/core/policy/engine.go index b4fb39c..e325af1 100644 --- a/pkg/core/policy/engine.go +++ b/pkg/core/policy/engine.go @@ -58,8 +58,8 @@ func (e *Engine) matchesPolicy(p *Policy, req *PolicyRequest) bool { // check evaluates a matched policy against the token claims. func (e *Engine) check(p *Policy, req *PolicyRequest) *PolicyResult { result := &PolicyResult{ - MatchedPolicy: p, - RequiredACR: p.RequireACR, + MatchedPolicy: p, + RequiredACR: p.RequireACR, RequiredMaxAge: p.MaxAge, } @@ -83,8 +83,15 @@ func (e *Engine) check(p *Policy, req *PolicyRequest) *PolicyResult { } } - // Check max_age (auth time freshness) - if p.MaxAge > 0 && req.AuthAge > 0 { + // Check max_age (auth time freshness). Fail closed: a policy that + // demands auth freshness cannot be satisfied by a token that carries + // no auth_time claim at all (AuthAge == 0 means "unknown"). + if p.MaxAge > 0 { + if req.AuthAge <= 0 { + result.Allowed = false + result.Reason = "policy requires max_age but token has no auth_time claim" + return result + } maxAge := time.Duration(p.MaxAge) * time.Second if req.AuthAge > maxAge { result.Allowed = false diff --git a/pkg/core/policy/engine_test.go b/pkg/core/policy/engine_test.go index 02e9607..482ccfd 100644 --- a/pkg/core/policy/engine_test.go +++ b/pkg/core/policy/engine_test.go @@ -94,10 +94,10 @@ func TestEngine_MaxAge(t *testing.T) { t.Error("expected denial when auth_age > max_age") } }) - t.Run("zero auth_age skips check", func(t *testing.T) { + t.Run("missing auth_time fails closed", func(t *testing.T) { result, _ := e.Evaluate(&PolicyRequest{Method: "GET", Path: "/secure", AuthAge: 0}) - if !result.Allowed { - t.Errorf("zero AuthAge should skip max_age check, got: %s", result.Reason) + if result.Allowed { + t.Error("policy with max_age must deny tokens that carry no auth_time claim") } }) } diff --git a/pkg/core/policy/matcher.go b/pkg/core/policy/matcher.go index 9e4708a..5276404 100644 --- a/pkg/core/policy/matcher.go +++ b/pkg/core/policy/matcher.go @@ -26,7 +26,9 @@ func MatchResource(pattern, requestPath string) bool { if prefix == "" { return true // ** matches everything } - return strings.HasPrefix(requestPath, prefix) + // Require a path-segment boundary so /api** does not match /apiv2 + // and /api/** does not match /api-internal/x. + return requestPath == prefix || strings.HasPrefix(requestPath, prefix+"/") } // Use path.Match for single-level wildcards diff --git a/pkg/core/policy/matcher_test.go b/pkg/core/policy/matcher_test.go new file mode 100644 index 0000000..89bfbfc --- /dev/null +++ b/pkg/core/policy/matcher_test.go @@ -0,0 +1,57 @@ +package policy + +import "testing" + +func TestMatchResource(t *testing.T) { + tests := []struct { + pattern string + path string + want bool + }{ + // Exact + {"/api/users", "/api/users", true}, + {"/api/users", "/api/users/42", false}, + + // Single-level wildcard + {"/api/users/*", "/api/users/42", true}, + {"/api/users/*", "/api/users/42/orders", false}, + + // Double wildcard with segment boundary + {"/api/**", "/api", true}, + {"/api/**", "/api/users", true}, + {"/api/**", "/api/users/42/orders", true}, + {"/api/**", "/api-internal/users", false}, // boundary regression + {"/api/**", "/apiv2/users", false}, // boundary regression + {"/api**", "/apiv2", false}, // boundary regression + {"/**", "/anything/at/all", true}, + + // Trailing slash normalization + {"/api/users/", "/api/users", true}, + {"/api/**", "/api/users/", true}, + } + for _, tt := range tests { + if got := MatchResource(tt.pattern, tt.path); got != tt.want { + t.Errorf("MatchResource(%q, %q) = %v, want %v", tt.pattern, tt.path, got, tt.want) + } + } +} + +func TestACRSatisfies(t *testing.T) { + hierarchy := []string{"bronze", "silver", "gold"} + tests := []struct { + provided, required string + want bool + }{ + {"gold", "silver", true}, + {"silver", "silver", true}, + {"bronze", "silver", false}, + {"unknown", "silver", false}, + {"silver", "unknown", false}, + {"same", "same", true}, // exact match outside hierarchy + } + for _, tt := range tests { + if got := ACRSatisfies(tt.provided, tt.required, hierarchy); got != tt.want { + t.Errorf("ACRSatisfies(%q, %q) = %v, want %v", tt.provided, tt.required, got, tt.want) + } + } +} diff --git a/pkg/core/policy/types.go b/pkg/core/policy/types.go index e0f7f2d..c49cf24 100644 --- a/pkg/core/policy/types.go +++ b/pkg/core/policy/types.go @@ -21,11 +21,11 @@ type Config struct { // Policy defines the access requirements for a set of resources. type Policy struct { Name string `yaml:"name"` - Resources []string `yaml:"resources"` // glob patterns, e.g. /api/payments/** - Methods []string `yaml:"methods"` // HTTP methods, empty = all - RequireACR string `yaml:"require_acr"` // minimum acr_values required - MaxAge int `yaml:"max_age"` // max auth age in seconds, 0 = unlimited - RequireMFA bool `yaml:"require_mfa"` // AMR must include mfa + Resources []string `yaml:"resources"` // glob patterns, e.g. /api/payments/** + Methods []string `yaml:"methods"` // HTTP methods, empty = all + RequireACR string `yaml:"require_acr"` // minimum acr_values required + MaxAge int `yaml:"max_age"` // max auth age in seconds, 0 = unlimited + RequireMFA bool `yaml:"require_mfa"` // AMR must include mfa RequireScopes []string `yaml:"require_scopes"` // RequireAuthorizationDetails enforces RFC 9396 authorization_details. @@ -50,11 +50,11 @@ type PolicyRequest struct { // PolicyResult is the output of policy evaluation. type PolicyResult struct { - Allowed bool + Allowed bool MatchedPolicy *Policy // If not allowed, these fields describe what is needed: - RequiredACR string + RequiredACR string RequiredMaxAge int - Reason string + Reason string } diff --git a/pkg/core/stepup/statemachine.go b/pkg/core/stepup/statemachine.go index 564de9e..8a54fac 100644 --- a/pkg/core/stepup/statemachine.go +++ b/pkg/core/stepup/statemachine.go @@ -16,7 +16,7 @@ import ( const ( // CookieName is the cookie used to carry SavedRequest state across the re-auth redirect. - CookieName = "iam_stepup_state" + CookieName = "iam_stepup_state" cookieMaxAge = 600 // 10 minutes ) @@ -24,10 +24,10 @@ const ( type State int const ( - StateIdle State = iota // No challenge in progress - StateChallenge // Challenge issued, waiting for re-auth - StateCompleted // Re-authentication successful - StateFailed // Re-authentication failed or timed out + StateIdle State = iota // No challenge in progress + StateChallenge // Challenge issued, waiting for re-auth + StateCompleted // Re-authentication successful + StateFailed // Re-authentication failed or timed out ) func (s State) String() string { @@ -53,13 +53,13 @@ const stepUpStateKey contextKey = iota // SavedRequest captures the original request before the step-up challenge, // so it can be replayed after successful re-authentication. type SavedRequest struct { - Method string - Path string - Query string - StateID string // random opaque value for CSRF protection - SavedAt time.Time - ACRHint string // acr_values that triggered this challenge - MaxAge int + Method string + Path string + Query string + StateID string // random opaque value for CSRF protection + SavedAt time.Time + ACRHint string // acr_values that triggered this challenge + MaxAge int } // Encode serializes the SavedRequest to a base64 string (for state param / cookie). @@ -114,6 +114,7 @@ func (sm *StateMachine) BeginChallenge(r *http.Request, challenge *StepUpChallen Method: r.Method, Path: r.URL.Path, Query: r.URL.RawQuery, + StateID: newStateID(), SavedAt: time.Now(), ACRHint: challenge.ACRValues, MaxAge: challenge.MaxAge, diff --git a/pkg/core/token/cache.go b/pkg/core/token/cache.go index 90c5b6d..7ba06ba 100644 --- a/pkg/core/token/cache.go +++ b/pkg/core/token/cache.go @@ -176,9 +176,19 @@ func (ci *CachedIntrospector) Introspect(ctx context.Context, token string) (*Co return nil, err } - // Only cache active tokens + // Only cache active tokens; clamp TTL to the token's remaining lifetime + // so a cache entry never outlives the token itself. if claims.Active { - _ = ci.cache.Set(ctx, key, claims, ci.ttl) + ttl := ci.ttl + if !claims.ExpiresAt.IsZero() { + if remaining := time.Until(claims.ExpiresAt); remaining < ttl { + ttl = remaining + } + } + if ttl > 0 { + _ = ci.cache.Set(ctx, key, claims, ttl) + IndexClaims(ctx, ci.cache, key, claims, ttl) + } } return claims, nil diff --git a/pkg/core/token/claims.go b/pkg/core/token/claims.go index da3f065..292d358 100644 --- a/pkg/core/token/claims.go +++ b/pkg/core/token/claims.go @@ -22,16 +22,22 @@ type CommonClaims struct { Subject string `json:"sub"` Issuer string `json:"iss"` Audience []string `json:"aud"` + JTI string `json:"jti"` ExpiresAt time.Time IssuedAt time.Time + // Confirmation carries the RFC 7800 cnf claim. For DPoP-bound tokens + // (RFC 9449) JKT holds the RFC 7638 SHA-256 thumbprint of the client's + // public key; the guard compares it against the DPoP proof's JWK. + Confirmation *Confirmation `json:"cnf,omitempty"` + // Auth context (RFC 9470) - ACR string `json:"acr"` // Authentication Context Class Reference + ACR string `json:"acr"` // Authentication Context Class Reference AMR []string `json:"amr"` // Authentication Methods References // Session - SessionID string `json:"sid"` - AuthTime time.Time `json:"auth_time"` // when user authenticated (for max_age check) + SessionID string `json:"sid"` + AuthTime time.Time `json:"auth_time"` // when user authenticated (for max_age check) // Identity Email string `json:"email"` @@ -53,6 +59,13 @@ type CommonClaims struct { Active bool } +// Confirmation is the RFC 7800 cnf (confirmation) claim. +type Confirmation struct { + // JKT is the RFC 7638 JWK SHA-256 thumbprint (base64url, no padding) + // of the DPoP public key the token is bound to (RFC 9449 §6.1). + JKT string `json:"jkt,omitempty"` +} + // AuthAge returns how long ago the user authenticated. func (c *CommonClaims) AuthAge() time.Duration { if c.AuthTime.IsZero() { @@ -61,6 +74,37 @@ func (c *CommonClaims) AuthAge() time.Duration { return time.Since(c.AuthTime) } +// --- FAPI 2.0 profile interface (pkg/core/fapi.TokenClaims) --- + +// HasDPoP reports whether the token is DPoP-bound (carries cnf.jkt). +func (c *CommonClaims) HasDPoP() bool { + return c.Confirmation != nil && c.Confirmation.JKT != "" +} + +// HasPARRequestURI reports whether the authorization was initiated via a +// Pushed Authorization Request (request_uri or par_id claim present). +func (c *CommonClaims) HasPARRequestURI() bool { + return c.extraString("request_uri") != "" || c.extraString("par_id") != "" +} + +// GetAuthAge returns the time since user authentication (0 = unknown). +func (c *CommonClaims) GetAuthAge() time.Duration { + return c.AuthAge() +} + +// GetNonce returns the token's nonce claim, if any. +func (c *CommonClaims) GetNonce() string { + return c.extraString("nonce") +} + +func (c *CommonClaims) extraString(key string) string { + if c.Extra == nil { + return "" + } + s, _ := c.Extra[key].(string) + return s +} + // HasScope checks if the token contains a specific scope. func (c *CommonClaims) HasScope(scope string) bool { for _, s := range c.Scopes { diff --git a/pkg/core/token/dpop.go b/pkg/core/token/dpop.go index 354ffa2..7aec7fa 100644 --- a/pkg/core/token/dpop.go +++ b/pkg/core/token/dpop.go @@ -4,16 +4,33 @@ import ( "crypto" "crypto/ecdsa" "crypto/elliptic" + "crypto/hmac" "crypto/rsa" + "crypto/sha256" "encoding/base64" "encoding/json" + "errors" "fmt" "math/big" "net/http" + "sort" "strings" "time" ) +// DPoP sentinel errors. +var ( + // ErrDPoPReplay indicates a DPoP proof jti was seen before (RFC 9449 §11.1). + ErrDPoPReplay = errors.New("dpop proof replay detected") + + // ErrDPoPNonceRequired indicates the server requires a fresh nonce; respond + // with 401, error="use_dpop_nonce" and a DPoP-Nonce header (RFC 9449 §8). + ErrDPoPNonceRequired = errors.New("dpop nonce required") + + // ErrDPoPMissingATH indicates the proof lacks the mandatory ath claim. + ErrDPoPMissingATH = errors.New("dpop proof missing ath (access token hash) claim") +) + // DPoPConfig holds configuration for DPoP proof validation (RFC 9449). type DPoPConfig struct { // MaxAge is the maximum allowed age of a DPoP proof (default: 60s). @@ -21,6 +38,10 @@ type DPoPConfig struct { // RequireHTTPS enforces that the htu claim uses HTTPS. RequireHTTPS bool + + // Nonce, when set, requires proofs to carry a valid server-issued nonce + // claim (RFC 9449 §8). Proofs without one fail with ErrDPoPNonceRequired. + Nonce *NonceProvider } // DefaultDPoPConfig returns sensible defaults. @@ -38,11 +59,61 @@ type DPoPProof struct { JWK map[string]interface{} // Payload fields - JTI string // unique proof ID - HTM string // HTTP method - HTU string // HTTP URI - IAT time.Time // issued at - ATH string // access token hash (base64url SHA-256) + JTI string // unique proof ID + HTM string // HTTP method + HTU string // HTTP URI + IAT time.Time // issued at + ATH string // access token hash (base64url SHA-256) + Nonce string // server-issued nonce (RFC 9449 §8) +} + +// Thumbprint computes the RFC 7638 SHA-256 thumbprint (base64url, no padding) +// of the proof's embedded JWK. Compare it against the access token's cnf.jkt +// to enforce key binding (RFC 9449 §6.1). +func (p *DPoPProof) Thumbprint() (string, error) { + return JWKThumbprint(p.JWK) +} + +// JWKThumbprint computes the RFC 7638 JWK SHA-256 thumbprint: the required +// members of the key (per key type) serialized in lexicographic order with no +// whitespace, hashed with SHA-256 and base64url-encoded without padding. +func JWKThumbprint(jwk map[string]interface{}) (string, error) { + kty, _ := jwk["kty"].(string) + var members []string + switch kty { + case "EC": + members = []string{"crv", "kty", "x", "y"} + case "RSA": + members = []string{"e", "kty", "n"} + case "OKP": + members = []string{"crv", "kty", "x"} + default: + return "", fmt.Errorf("unsupported JWK kty for thumbprint: %q", kty) + } + sort.Strings(members) + + var sb strings.Builder + sb.WriteByte('{') + for i, m := range members { + v, ok := jwk[m].(string) + if !ok || v == "" { + return "", fmt.Errorf("JWK missing required member %q for thumbprint", m) + } + if i > 0 { + sb.WriteByte(',') + } + // Required members are all string-valued; encode via json.Marshal to + // handle any escaping correctly. + kb, _ := json.Marshal(m) + vb, _ := json.Marshal(v) + sb.Write(kb) + sb.WriteByte(':') + sb.Write(vb) + } + sb.WriteByte('}') + + sum := sha256.Sum256([]byte(sb.String())) + return base64.RawURLEncoding.EncodeToString(sum[:]), nil } // ValidateDPoP validates the DPoP proof from the request header against the access token. @@ -89,17 +160,122 @@ func ValidateDPoP(r *http.Request, accessToken string, cfg DPoPConfig) (*DPoPPro return nil, fmt.Errorf("DPoP proof issued in the future") } - // Validate ATH (access token hash) if present - if proof.ATH != "" { - expectedATH := hashTokenForDPoP(accessToken) - if proof.ATH != expectedATH { - return nil, ErrDPoPBindingMismatch - } + // Validate ATH (access token hash) — mandatory. A proof that does not + // commit to the access token can be replayed with any stolen token. + if proof.ATH == "" { + return nil, ErrDPoPMissingATH + } + if proof.ATH != hashTokenForDPoP(accessToken) { + return nil, ErrDPoPBindingMismatch + } + + // Validate server-issued nonce when required (RFC 9449 §8). + if cfg.Nonce != nil && !cfg.Nonce.Valid(proof.Nonce) { + return nil, ErrDPoPNonceRequired } return proof, nil } +// DPoPValidator validates DPoP proofs with jti replay protection. +// The replay cache records each seen jti for the proof MaxAge window; +// a second proof with the same jti is rejected (RFC 9449 §11.1). +type DPoPValidator struct { + cfg DPoPConfig + cache Cache +} + +const dpopJTIPrefix = "dpop:jti:" + +// NewDPoPValidator creates a validator. If cache is nil an in-memory cache is +// used (single-instance replay protection only — pass a shared cache in +// multi-instance deployments). +func NewDPoPValidator(cfg DPoPConfig, cache Cache) *DPoPValidator { + if cfg.MaxAge <= 0 { + cfg.MaxAge = 60 * time.Second + } + if cache == nil { + cache = NewMemoryCache() + } + return &DPoPValidator{cfg: cfg, cache: cache} +} + +// Validate runs full proof validation (ValidateDPoP) plus jti replay detection. +func (v *DPoPValidator) Validate(r *http.Request, accessToken string) (*DPoPProof, error) { + proof, err := ValidateDPoP(r, accessToken, v.cfg) + if err != nil { + return nil, err + } + + if proof.JTI == "" { + return nil, fmt.Errorf("dpop proof missing jti claim") + } + ctx := r.Context() + key := dpopJTIPrefix + proof.JTI + if _, seen := v.cache.Get(ctx, key); seen { + return nil, ErrDPoPReplay + } + // Record the jti for the proof acceptance window; entries expire with it. + _ = v.cache.Set(ctx, key, &CommonClaims{}, v.cfg.MaxAge) + + return proof, nil +} + +// NonceProvider issues and validates server DPoP nonces (RFC 9449 §8). +// Nonces are HMAC(secret, time-window) so they need no storage and stay valid +// across instances sharing the secret. The current and previous windows are +// both accepted, giving clients between Window and 2×Window to use a nonce. +type NonceProvider struct { + secret []byte + window time.Duration +} + +// NewNonceProvider creates a nonce provider. Window defaults to 5 minutes. +func NewNonceProvider(secret string, window time.Duration) *NonceProvider { + if window <= 0 { + window = 5 * time.Minute + } + return &NonceProvider{secret: []byte(secret), window: window} +} + +// Current returns the nonce for the current time window; send it to clients +// in the DPoP-Nonce response header. +func (n *NonceProvider) Current() string { + return n.forWindow(time.Now().UnixNano() / int64(n.window)) +} + +// Valid reports whether nonce matches the current or previous window. +func (n *NonceProvider) Valid(nonce string) bool { + if nonce == "" { + return false + } + w := time.Now().UnixNano() / int64(n.window) + return hmac.Equal([]byte(nonce), []byte(n.forWindow(w))) || + hmac.Equal([]byte(nonce), []byte(n.forWindow(w-1))) +} + +func (n *NonceProvider) forWindow(w int64) string { + mac := hmac.New(sha256.New, n.secret) + fmt.Fprintf(mac, "%d", w) + return base64.RawURLEncoding.EncodeToString(mac.Sum(nil)) +} + +// VerifyDPoPBinding enforces RFC 9449 §6.1 key binding: the access token's +// cnf.jkt thumbprint must equal the RFC 7638 thumbprint of the proof's JWK. +func VerifyDPoPBinding(proof *DPoPProof, claims *CommonClaims) error { + if claims.Confirmation == nil || claims.Confirmation.JKT == "" { + return fmt.Errorf("access token is not DPoP-bound (no cnf.jkt claim)") + } + thumb, err := proof.Thumbprint() + if err != nil { + return fmt.Errorf("computing proof JWK thumbprint: %w", err) + } + if !hmac.Equal([]byte(thumb), []byte(claims.Confirmation.JKT)) { + return ErrDPoPBindingMismatch + } + return nil +} + // parseDPoPJWT parses the DPoP proof JWT header and payload. func parseDPoPJWT(jwt string) (*DPoPProof, error) { parts := strings.Split(jwt, ".") @@ -129,11 +305,12 @@ func parseDPoPJWT(jwt string) (*DPoPProof, error) { } var payload struct { - JTI string `json:"jti"` - HTM string `json:"htm"` - HTU string `json:"htu"` - IAT int64 `json:"iat"` - ATH string `json:"ath"` + JTI string `json:"jti"` + HTM string `json:"htm"` + HTU string `json:"htu"` + IAT int64 `json:"iat"` + ATH string `json:"ath"` + Nonce string `json:"nonce"` } if err := json.Unmarshal(payloadBytes, &payload); err != nil { return nil, fmt.Errorf("parsing payload: %w", err) @@ -147,6 +324,7 @@ func parseDPoPJWT(jwt string) (*DPoPProof, error) { HTU: payload.HTU, IAT: time.Unix(payload.IAT, 0), ATH: payload.ATH, + Nonce: payload.Nonce, }, nil } diff --git a/pkg/core/token/dpop_binding_test.go b/pkg/core/token/dpop_binding_test.go new file mode 100644 index 0000000..a1e5dc7 --- /dev/null +++ b/pkg/core/token/dpop_binding_test.go @@ -0,0 +1,129 @@ +package token + +import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "errors" + "net/http/httptest" + "testing" + "time" +) + +// TestJWKThumbprint_RFC7638Vector checks the RSA test vector from RFC 7638 §3.1. +func TestJWKThumbprint_RFC7638Vector(t *testing.T) { + jwk := map[string]interface{}{ + "kty": "RSA", + "n": "0vx7agoebGcQSuuPiLJXZptN9nndrQmbXEps2aiAFbWhM78LhWx4cbbfAAtVT86zwu1RK7aPFFxuhDR1L6tSoc_BJECPebWKRXjBZCiFV4n3oknjhMstn64tZ_2W-5JsGY4Hc5n9yBXArwl93lqt7_RN5w6Cf0h4QyQ5v-65YGjQR0_FDW2QvzqY368QQMicAtaSqzs8KJZgnYb9c7d0zgdAZHzu6qMQvRL5hajrn1n91CbOpbISD08qNLyrdkt-bFTWhAI4vMQFh6WeZu0fM4lFd2NcRwr3XPksINHaQ-G_xBniIqbw0Ls1jF44-csFCur-kEgU8awapJzKnqDKgw", + "e": "AQAB", + } + want := "NzbLsXh8uDCcd-6MNwXF4W_7noWXFZAfHkxZsRGC9Xs" + got, err := JWKThumbprint(jwk) + if err != nil { + t.Fatalf("JWKThumbprint: %v", err) + } + if got != want { + t.Errorf("thumbprint = %q, want %q", got, want) + } +} + +func TestValidateDPoP_MissingATH_Rejected(t *testing.T) { + priv, _ := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + proof := signDPoP(t, priv, "GET", "https://api.example.com/r", "") // no ath + + req := httptest.NewRequest("GET", "https://api.example.com/r", nil) + req.Header.Set("DPoP", proof) + + _, err := ValidateDPoP(req, "tok", DPoPConfig{MaxAge: 60 * time.Second}) + if !errors.Is(err, ErrDPoPMissingATH) { + t.Fatalf("expected ErrDPoPMissingATH, got %v", err) + } +} + +func TestDPoPValidator_ReplayRejected(t *testing.T) { + priv, _ := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + const accessToken = "tok" + proof := signDPoP(t, priv, "GET", "https://api.example.com/r", hashTokenForDPoP(accessToken)) + + v := NewDPoPValidator(DPoPConfig{MaxAge: 60 * time.Second}, NewMemoryCache()) + + req := httptest.NewRequest("GET", "https://api.example.com/r", nil) + req.Header.Set("DPoP", proof) + if _, err := v.Validate(req, accessToken); err != nil { + t.Fatalf("first use should pass: %v", err) + } + + req2 := httptest.NewRequest("GET", "https://api.example.com/r", nil) + req2.Header.Set("DPoP", proof) + if _, err := v.Validate(req2, accessToken); !errors.Is(err, ErrDPoPReplay) { + t.Fatalf("expected ErrDPoPReplay on second use, got %v", err) + } +} + +func TestVerifyDPoPBinding(t *testing.T) { + priv, _ := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + const accessToken = "tok" + proofJWT := signDPoP(t, priv, "GET", "https://api.example.com/r", hashTokenForDPoP(accessToken)) + + req := httptest.NewRequest("GET", "https://api.example.com/r", nil) + req.Header.Set("DPoP", proofJWT) + proof, err := ValidateDPoP(req, accessToken, DPoPConfig{MaxAge: 60 * time.Second}) + if err != nil { + t.Fatalf("ValidateDPoP: %v", err) + } + + thumb, err := proof.Thumbprint() + if err != nil { + t.Fatalf("Thumbprint: %v", err) + } + + t.Run("matching jkt passes", func(t *testing.T) { + claims := &CommonClaims{Confirmation: &Confirmation{JKT: thumb}} + if err := VerifyDPoPBinding(proof, claims); err != nil { + t.Errorf("expected binding to pass: %v", err) + } + }) + t.Run("mismatched jkt rejected", func(t *testing.T) { + claims := &CommonClaims{Confirmation: &Confirmation{JKT: "wrong-thumbprint"}} + if !errors.Is(VerifyDPoPBinding(proof, claims), ErrDPoPBindingMismatch) { + t.Error("expected ErrDPoPBindingMismatch") + } + }) + t.Run("unbound token rejected", func(t *testing.T) { + if VerifyDPoPBinding(proof, &CommonClaims{}) == nil { + t.Error("token without cnf.jkt must be rejected when DPoP is enforced") + } + }) +} + +func TestNonceProvider(t *testing.T) { + np := NewNonceProvider("secret", 5*time.Minute) + if !np.Valid(np.Current()) { + t.Error("current nonce must validate") + } + if np.Valid("bogus") { + t.Error("bogus nonce must not validate") + } + if np.Valid("") { + t.Error("empty nonce must not validate") + } + + other := NewNonceProvider("other-secret", 5*time.Minute) + if np.Valid(other.Current()) { + t.Error("nonce from a different secret must not validate") + } +} + +func TestValidateDPoP_NonceRequired(t *testing.T) { + priv, _ := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + const accessToken = "tok" + proof := signDPoP(t, priv, "GET", "https://api.example.com/r", hashTokenForDPoP(accessToken)) + + req := httptest.NewRequest("GET", "https://api.example.com/r", nil) + req.Header.Set("DPoP", proof) + + cfg := DPoPConfig{MaxAge: 60 * time.Second, Nonce: NewNonceProvider("s", 0)} + if _, err := ValidateDPoP(req, accessToken, cfg); !errors.Is(err, ErrDPoPNonceRequired) { + t.Fatalf("expected ErrDPoPNonceRequired for proof without nonce, got %v", err) + } +} diff --git a/pkg/core/token/goredis/integration_test.go b/pkg/core/token/goredis/integration_test.go new file mode 100644 index 0000000..d7959a6 --- /dev/null +++ b/pkg/core/token/goredis/integration_test.go @@ -0,0 +1,88 @@ +package goredis_test + +import ( + "context" + "os" + "testing" + "time" + + "github.com/redis/go-redis/v9" + + iamtoken "github.com/common-iam/iam/pkg/core/token" + "github.com/common-iam/iam/pkg/core/token/goredis" +) + +// TestRedisIntegration exercises the adapter against a real Redis instance. +// Opt in by setting REDIS_ADDR (e.g. REDIS_ADDR=localhost:6379 go test ./...). +func TestRedisIntegration(t *testing.T) { + addr := os.Getenv("REDIS_ADDR") + if addr == "" { + t.Skip("set REDIS_ADDR to run the Redis integration test") + } + + ctx := context.Background() + client := redis.NewClient(&redis.Options{Addr: addr, DB: 15}) // dedicated test DB + if err := client.Ping(ctx).Err(); err != nil { + t.Fatalf("pinging Redis at %s: %v", addr, err) + } + t.Cleanup(func() { + _ = client.FlushDB(ctx).Err() + _ = client.Close() + }) + + cache := iamtoken.NewRedisCache(goredis.New(client), "iamtest:") + + t.Run("set/get/delete round-trip", func(t *testing.T) { + claims := &iamtoken.CommonClaims{ + Active: true, + Subject: "alice", + JTI: "jti-1", + Scopes: []string{"openid"}, + } + if err := cache.Set(ctx, "hash1", claims, time.Minute); err != nil { + t.Fatalf("Set: %v", err) + } + got, ok := cache.Get(ctx, "hash1") + if !ok || got.Subject != "alice" || len(got.Scopes) != 1 { + t.Fatalf("Get = %+v, ok=%v", got, ok) + } + if err := cache.Delete(ctx, "hash1"); err != nil { + t.Fatalf("Delete: %v", err) + } + if _, ok := cache.Get(ctx, "hash1"); ok { + t.Error("deleted entry still present") + } + }) + + t.Run("ttl expiry", func(t *testing.T) { + claims := &iamtoken.CommonClaims{Active: true, Subject: "bob"} + if err := cache.Set(ctx, "hash-ttl", claims, time.Second); err != nil { + t.Fatalf("Set: %v", err) + } + time.Sleep(1500 * time.Millisecond) + if _, ok := cache.Get(ctx, "hash-ttl"); ok { + t.Error("entry should have expired after TTL") + } + }) + + t.Run("revocation index round-trip", func(t *testing.T) { + claims := &iamtoken.CommonClaims{ + Active: true, + Subject: "carol", + JTI: "jti-idx", + SessionID: "sess-idx", + } + if err := cache.Set(ctx, "hash-idx", claims, time.Minute); err != nil { + t.Fatalf("Set: %v", err) + } + iamtoken.IndexClaims(ctx, cache, "hash-idx", claims, time.Minute) + + // Simulate a webhook revocation by JTI through the real handler path. + // (The handler is exercised elsewhere; here we verify index integrity + // across a real Redis JSON round-trip.) + got, ok := cache.Get(ctx, "hash-idx") + if !ok || got.JTI != "jti-idx" { + t.Fatalf("Get = %+v, ok=%v", got, ok) + } + }) +} diff --git a/pkg/core/token/introspect.go b/pkg/core/token/introspect.go index 8ba3c6e..a0eb151 100644 --- a/pkg/core/token/introspect.go +++ b/pkg/core/token/introspect.go @@ -29,10 +29,36 @@ type IntrospectionResponse struct { TokenType string `json:"token_type"` JTI string `json:"jti"` + // Aud is the RFC 7662 aud claim; a string or array of strings per RFC 7519 §4.1.3. + Aud Audience `json:"aud,omitempty"` + + // Cnf is the RFC 7800 confirmation claim; cnf.jkt binds the token to a + // DPoP key (RFC 9449 §6.1). + Cnf *Confirmation `json:"cnf,omitempty"` + // AuthorizationDetails carries RFC 9396 authorization_details when present. AuthorizationDetails []rar.AuthorizationDetail `json:"authorization_details,omitempty"` } +// Audience unmarshals the JWT aud claim, which may be a single string or an +// array of strings (RFC 7519 §4.1.3). +type Audience []string + +// UnmarshalJSON implements json.Unmarshaler. +func (a *Audience) UnmarshalJSON(b []byte) error { + var single string + if err := json.Unmarshal(b, &single); err == nil { + *a = Audience{single} + return nil + } + var many []string + if err := json.Unmarshal(b, &many); err != nil { + return fmt.Errorf("aud must be a string or array of strings: %w", err) + } + *a = Audience(many) + return nil +} + // IntrospectorConfig holds configuration for the token introspector. type IntrospectorConfig struct { // Endpoint is the RFC 7662 introspection endpoint URL. @@ -72,6 +98,7 @@ func (i *Introspector) Introspect(ctx context.Context, token string) (*CommonCla form := url.Values{} form.Set("token", token) + form.Set("token_type_hint", "access_token") // RFC 7662 §2.1 SHOULD req, err := http.NewRequestWithContext(ctx, http.MethodPost, i.cfg.Endpoint, strings.NewReader(form.Encode())) @@ -102,12 +129,15 @@ func (i *Introspector) Introspect(ctx context.Context, token string) (*CommonCla // introToCommonClaims maps an IntrospectionResponse to CommonClaims. func introToCommonClaims(r *IntrospectionResponse) *CommonClaims { c := &CommonClaims{ - Active: r.Active, - Subject: r.Sub, - Issuer: r.Iss, - ACR: r.ACR, - AMR: r.AMR, - Username: r.Username, + Active: r.Active, + Subject: r.Sub, + Issuer: r.Iss, + Audience: r.Aud, + JTI: r.JTI, + ACR: r.ACR, + AMR: r.AMR, + Username: r.Username, + Confirmation: r.Cnf, } if r.Exp > 0 { c.ExpiresAt = time.Unix(r.Exp, 0) diff --git a/pkg/core/token/introspect_test.go b/pkg/core/token/introspect_test.go new file mode 100644 index 0000000..9308cae --- /dev/null +++ b/pkg/core/token/introspect_test.go @@ -0,0 +1,249 @@ +package token + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" +) + +// newIntrospectionServer serves a canned RFC 7662 response and records the +// last form values it received. +func newIntrospectionServer(t *testing.T, resp map[string]interface{}) (*httptest.Server, *map[string][]string) { + t.Helper() + var lastForm map[string][]string + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if err := r.ParseForm(); err != nil { + t.Errorf("ParseForm: %v", err) + } + lastForm = r.PostForm + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(resp) + })) + t.Cleanup(srv.Close) + return srv, &lastForm +} + +func TestIntrospector_Introspect(t *testing.T) { + exp := time.Now().Add(time.Hour).Unix() + srv, lastForm := newIntrospectionServer(t, map[string]interface{}{ + "active": true, + "sub": "alice", + "iss": "https://as.example.com", + "aud": []string{"api", "web"}, + "exp": exp, + "iat": time.Now().Unix(), + "auth_time": time.Now().Add(-time.Minute).Unix(), + "acr": "silver", + "amr": []string{"pwd", "otp"}, + "scope": "openid profile", + "jti": "tok-123", + "cnf": map[string]string{"jkt": "thumb-abc"}, + }) + + intro := NewIntrospector(IntrospectorConfig{ + Endpoint: srv.URL, + ClientID: "client", + }) + claims, err := intro.Introspect(context.Background(), "raw-token") + if err != nil { + t.Fatalf("Introspect: %v", err) + } + + if !claims.Active || claims.Subject != "alice" || claims.Issuer != "https://as.example.com" { + t.Errorf("basic claims wrong: %+v", claims) + } + if len(claims.Audience) != 2 || claims.Audience[0] != "api" { + t.Errorf("aud = %v, want [api web]", claims.Audience) + } + if claims.JTI != "tok-123" { + t.Errorf("jti = %q", claims.JTI) + } + if claims.Confirmation == nil || claims.Confirmation.JKT != "thumb-abc" { + t.Errorf("cnf.jkt not mapped: %+v", claims.Confirmation) + } + if len(claims.Scopes) != 2 || claims.Scopes[0] != "openid" { + t.Errorf("scopes = %v", claims.Scopes) + } + if claims.AuthAge() <= 0 { + t.Error("auth_time should yield a positive AuthAge") + } + + // RFC 7662 §2.1 SHOULD: token_type_hint sent. + if got := (*lastForm)["token_type_hint"]; len(got) != 1 || got[0] != "access_token" { + t.Errorf("token_type_hint = %v, want [access_token]", got) + } +} + +func TestAudience_UnmarshalJSON_StringForm(t *testing.T) { + var r IntrospectionResponse + if err := json.Unmarshal([]byte(`{"active":true,"aud":"single-api"}`), &r); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if len(r.Aud) != 1 || r.Aud[0] != "single-api" { + t.Errorf("aud = %v, want [single-api]", r.Aud) + } + + if err := json.Unmarshal([]byte(`{"aud":42}`), &r); err == nil { + t.Error("numeric aud must fail to unmarshal") + } +} + +func TestExtractBearerToken(t *testing.T) { + tests := []struct { + header string + want string + wantErr bool + }{ + {"Bearer abc123", "abc123", false}, + {"bearer abc123", "abc123", false}, + {"", "", true}, + {"Basic dXNlcjpwYXNz", "", true}, + {"Bearer ", "", true}, + {"Bearer", "", true}, + } + for _, tt := range tests { + got, err := ExtractBearerToken(tt.header) + if (err != nil) != tt.wantErr { + t.Errorf("ExtractBearerToken(%q) error = %v, wantErr %v", tt.header, err, tt.wantErr) + } + if got != tt.want { + t.Errorf("ExtractBearerToken(%q) = %q, want %q", tt.header, got, tt.want) + } + } +} + +func TestCachedIntrospector(t *testing.T) { + calls := 0 + exp := time.Now().Add(time.Hour).Unix() + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + calls++ + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]interface{}{ + "active": true, "sub": "bob", "exp": exp, "jti": "j-1", + }) + })) + t.Cleanup(srv.Close) + + cache := NewMemoryCache() + ci := NewCachedIntrospector(NewIntrospector(IntrospectorConfig{Endpoint: srv.URL}), cache, 30*time.Second) + + ctx := context.Background() + if _, err := ci.Introspect(ctx, "tok"); err != nil { + t.Fatalf("first introspect: %v", err) + } + if _, err := ci.Introspect(ctx, "tok"); err != nil { + t.Fatalf("second introspect: %v", err) + } + if calls != 1 { + t.Errorf("expected 1 upstream call (second served from cache), got %d", calls) + } + + // Revoke evicts; next call goes upstream again. + ci.Revoke(ctx, "tok") + if _, err := ci.Introspect(ctx, "tok"); err != nil { + t.Fatalf("post-revoke introspect: %v", err) + } + if calls != 2 { + t.Errorf("expected 2 upstream calls after revoke, got %d", calls) + } +} + +func TestCachedIntrospector_DoesNotCacheExpired(t *testing.T) { + calls := 0 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + calls++ + w.Header().Set("Content-Type", "application/json") + // Active but already past exp — TTL clamp must prevent caching. + _ = json.NewEncoder(w).Encode(map[string]interface{}{ + "active": true, "sub": "bob", "exp": time.Now().Add(-time.Minute).Unix(), + }) + })) + t.Cleanup(srv.Close) + + ci := NewCachedIntrospector(NewIntrospector(IntrospectorConfig{Endpoint: srv.URL}), NewMemoryCache(), 30*time.Second) + ctx := context.Background() + _, _ = ci.Introspect(ctx, "tok") + _, _ = ci.Introspect(ctx, "tok") + if calls != 2 { + t.Errorf("expired-token result must not be cached; upstream calls = %d, want 2", calls) + } +} + +// BenchmarkCachedIntrospector_CacheHit measures the hot path: introspection +// served entirely from the in-memory cache. +func BenchmarkCachedIntrospector_CacheHit(b *testing.B) { + exp := time.Now().Add(time.Hour).Unix() + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]interface{}{ + "active": true, "sub": "bench", "exp": exp, + }) + })) + defer srv.Close() + + ci := NewCachedIntrospector(NewIntrospector(IntrospectorConfig{Endpoint: srv.URL}), NewMemoryCache(), 5*time.Minute) + ctx := context.Background() + if _, err := ci.Introspect(ctx, "bench-token"); err != nil { + b.Fatal(err) + } + + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + if _, err := ci.Introspect(ctx, "bench-token"); err != nil { + b.Fatal(err) + } + } +} + +func TestMemoryCache_Flush(t *testing.T) { + ctx := context.Background() + c := NewMemoryCache() + _ = c.Set(ctx, "k1", &CommonClaims{Subject: "a"}, time.Minute) + _ = c.Set(ctx, "k2", &CommonClaims{Subject: "b"}, time.Minute) + if err := c.Flush(ctx); err != nil { + t.Fatalf("Flush: %v", err) + } + if _, ok := c.Get(ctx, "k1"); ok { + t.Error("k1 should be gone after Flush") + } +} + +func TestCommonClaims_Helpers(t *testing.T) { + c := &CommonClaims{ + Roles: []string{"admin"}, + Scopes: []string{"openid"}, + Extra: map[string]interface{}{ + "nonce": "n-1", + "request_uri": "urn:ietf:params:oauth:request_uri:abc", + }, + Confirmation: &Confirmation{JKT: "x"}, + AuthTime: time.Now().Add(-2 * time.Minute), + } + if !c.HasRole("admin") || c.HasRole("user") { + t.Error("HasRole wrong") + } + if !c.HasScope("openid") || c.HasScope("payments") { + t.Error("HasScope wrong") + } + if !c.HasDPoP() { + t.Error("HasDPoP should be true with cnf.jkt") + } + if !c.HasPARRequestURI() { + t.Error("HasPARRequestURI should be true with request_uri") + } + if c.GetNonce() != "n-1" { + t.Errorf("GetNonce = %q", c.GetNonce()) + } + if c.GetAuthAge() < time.Minute { + t.Errorf("GetAuthAge = %s, want ≥ 1m", c.GetAuthAge()) + } + + empty := &CommonClaims{} + if empty.HasDPoP() || empty.HasPARRequestURI() || empty.GetNonce() != "" || empty.GetAuthAge() != 0 { + t.Error("zero-value claims must report no bindings") + } +} diff --git a/pkg/core/token/jwtvalidator.go b/pkg/core/token/jwtvalidator.go index cd8e61a..b68b96a 100644 --- a/pkg/core/token/jwtvalidator.go +++ b/pkg/core/token/jwtvalidator.go @@ -23,6 +23,18 @@ type JWTValidatorConfig struct { // HTTPClient is used to fetch the JWKS (default: 10s timeout client). HTTPClient *http.Client + + // ExpectedIssuer, when set, rejects tokens whose iss claim differs. + ExpectedIssuer string + + // ExpectedAudience, when set, rejects tokens whose aud claim does not + // contain this value. + ExpectedAudience string + + // ValidMethods restricts acceptable signing algorithms (e.g. ["RS256", + // "ES256"]). Default: RS256/RS384/RS512/ES256/ES384/ES512 — "none" and + // HMAC algorithms are never accepted. + ValidMethods []string } // JWTValidator validates JWTs locally using a remote JWKS endpoint. @@ -32,12 +44,17 @@ type JWTValidator struct { jwksURL string httpClient *http.Client cacheTTL time.Duration + parseOpts []jwt.ParserOption mu sync.RWMutex keySet map[string]crypto.PublicKey // kid → key; "" key for kidless JWKS fetchedAt time.Time } +// defaultValidMethods are the asymmetric algorithms accepted when +// ValidMethods is not configured. HMAC and "none" are never accepted. +var defaultValidMethods = []string{"RS256", "RS384", "RS512", "ES256", "ES384", "ES512"} + // NewJWTValidator creates a local JWT validator backed by a JWKS endpoint. func NewJWTValidator(cfg JWTValidatorConfig) *JWTValidator { if cfg.CacheTTL == 0 { @@ -46,20 +63,35 @@ func NewJWTValidator(cfg JWTValidatorConfig) *JWTValidator { if cfg.HTTPClient == nil { cfg.HTTPClient = &http.Client{Timeout: 10 * time.Second} } + + methods := cfg.ValidMethods + if len(methods) == 0 { + methods = defaultValidMethods + } + opts := []jwt.ParserOption{ + jwt.WithExpirationRequired(), + jwt.WithIssuedAt(), + jwt.WithValidMethods(methods), + } + if cfg.ExpectedIssuer != "" { + opts = append(opts, jwt.WithIssuer(cfg.ExpectedIssuer)) + } + if cfg.ExpectedAudience != "" { + opts = append(opts, jwt.WithAudience(cfg.ExpectedAudience)) + } + return &JWTValidator{ jwksURL: cfg.JWKSURL, httpClient: cfg.HTTPClient, cacheTTL: cfg.CacheTTL, + parseOpts: opts, keySet: make(map[string]crypto.PublicKey), } } // Validate parses and verifies a raw JWT, returning normalized CommonClaims. func (v *JWTValidator) Validate(ctx context.Context, rawToken string) (*CommonClaims, error) { - token, err := jwt.ParseWithClaims(rawToken, &jwtRawClaims{}, v.keyfunc(ctx), - jwt.WithExpirationRequired(), - jwt.WithIssuedAt(), - ) + token, err := jwt.ParseWithClaims(rawToken, &jwtRawClaims{}, v.keyfunc(ctx), v.parseOpts...) if err != nil { return nil, fmt.Errorf("jwt validation: %w", err) } diff --git a/pkg/core/token/jwtvalidator_test.go b/pkg/core/token/jwtvalidator_test.go index 1649d41..9af3c5a 100644 --- a/pkg/core/token/jwtvalidator_test.go +++ b/pkg/core/token/jwtvalidator_test.go @@ -87,7 +87,7 @@ func TestJWTValidator_ExpiredToken(t *testing.T) { func TestJWTValidator_WrongKey(t *testing.T) { factory, _ := tokenfactory.New() - otherFactory, _ := tokenfactory.New() // different key pair + otherFactory, _ := tokenfactory.New() // different key pair jwksSrv := startJWKSServer(t, otherFactory) // serve other factory's JWKS rawToken, _ := factory.Generate(tokenfactory.TokenOptions{ diff --git a/pkg/core/token/rediscache_test.go b/pkg/core/token/rediscache_test.go new file mode 100644 index 0000000..0e73107 --- /dev/null +++ b/pkg/core/token/rediscache_test.go @@ -0,0 +1,107 @@ +package token + +import ( + "context" + "errors" + "strings" + "testing" + "time" +) + +// fakeRedis implements RedisClient over a plain map (no TTL expiry). +type fakeRedis struct { + data map[string]string +} + +func newFakeRedis() *fakeRedis { return &fakeRedis{data: make(map[string]string)} } + +func (f *fakeRedis) Get(_ context.Context, key string) (string, error) { + v, ok := f.data[key] + if !ok { + return "", errors.New("redis: nil") + } + return v, nil +} + +func (f *fakeRedis) Set(_ context.Context, key, value string, _ time.Duration) error { + f.data[key] = value + return nil +} + +func (f *fakeRedis) Del(_ context.Context, keys ...string) error { + for _, k := range keys { + delete(f.data, k) + } + return nil +} + +func (f *fakeRedis) FlushDB(_ context.Context) error { + f.data = make(map[string]string) + return nil +} + +func TestRedisCache_RoundTrip(t *testing.T) { + ctx := context.Background() + client := newFakeRedis() + c := NewRedisCache(client, "") + + claims := &CommonClaims{Active: true, Subject: "alice", JTI: "j1"} + if err := c.Set(ctx, "hash1", claims, time.Minute); err != nil { + t.Fatalf("Set: %v", err) + } + + // Key prefix applied. + for k := range client.data { + if !strings.HasPrefix(k, "iam:token:") { + t.Errorf("key %q missing default prefix", k) + } + } + + got, ok := c.Get(ctx, "hash1") + if !ok || got.Subject != "alice" || got.JTI != "j1" { + t.Fatalf("Get = %+v, ok=%v", got, ok) + } + + if _, ok := c.Get(ctx, "missing"); ok { + t.Error("missing key must not be found") + } + + if err := c.Delete(ctx, "hash1"); err != nil { + t.Fatalf("Delete: %v", err) + } + if _, ok := c.Get(ctx, "hash1"); ok { + t.Error("deleted key must not be found") + } + + _ = c.Set(ctx, "hash2", claims, time.Minute) + if err := c.Flush(ctx); err != nil { + t.Fatalf("Flush: %v", err) + } + if _, ok := c.Get(ctx, "hash2"); ok { + t.Error("flushed key must not be found") + } +} + +// TestRedisCache_IndexRoundTrip verifies the revocation secondary index +// survives the JSON round-trip a real Redis imposes ([]string → []interface{}). +func TestRedisCache_IndexRoundTrip(t *testing.T) { + ctx := context.Background() + c := NewRedisCache(newFakeRedis(), "") + + claims := &CommonClaims{Active: true, Subject: "alice", JTI: "j1", SessionID: "s1"} + _ = c.Set(ctx, "hash1", claims, time.Minute) + IndexClaims(ctx, c, "hash1", claims, time.Minute) + + if hash, ok := lookupIndexHash(ctx, c, jtiIndexPrefix+"j1"); !ok || hash != "hash1" { + t.Errorf("jti index lookup = %q, %v", hash, ok) + } + if hashes := indexHashes(ctx, c, subIndexPrefix+"alice"); len(hashes) != 1 || hashes[0] != "hash1" { + t.Errorf("subject index = %v", hashes) + } + if err := deleteIndexedTokens(ctx, c, sidIndexPrefix+"s1"); err != nil { + t.Fatalf("deleteIndexedTokens: %v", err) + } + if _, ok := c.Get(ctx, "hash1"); ok { + t.Error("session index deletion must evict the token entry") + } +} diff --git a/pkg/core/token/revocation.go b/pkg/core/token/revocation.go index b67181d..3042b3f 100644 --- a/pkg/core/token/revocation.go +++ b/pkg/core/token/revocation.go @@ -11,6 +11,7 @@ import ( "log/slog" "net/http" "strings" + "time" ) // RevocationEvent represents a token revocation notification. @@ -114,16 +115,118 @@ func (h *RevocationHandler) process(ctx context.Context, event *RevocationEvent) if event.JTI != "" { h.logger.Info("revoking token by JTI", "jti", event.JTI) - // JTI is used as cache key when token hash is not available + // Resolve the jti to its token hash via the secondary index so the + // actual cache entry is evicted, then drop the index entry itself. + if hash, ok := lookupIndexHash(ctx, h.cache, jtiIndexPrefix+event.JTI); ok { + _ = h.cache.Delete(ctx, jtiIndexPrefix+event.JTI) + return h.cache.Delete(ctx, hash) + } + // Legacy fallback: some callers cached directly under the jti. return h.cache.Delete(ctx, event.JTI) } + if event.SessionID != "" { + h.logger.Info("revoking tokens by session", "sid", event.SessionID) + return deleteIndexedTokens(ctx, h.cache, sidIndexPrefix+event.SessionID) + } + if event.RevokeAll && event.Subject != "" { - // For flush-all scenarios, we clear the entire cache. - // Production systems should use a more targeted approach (e.g., per-user key prefix). - h.logger.Warn("flushing all cache entries for subject", "sub", event.Subject) - return h.cache.Flush(ctx) + // Targeted eviction: delete only the cached tokens recorded for this + // subject instead of flushing the whole (shared, multi-tenant) cache. + h.logger.Info("revoking all cached tokens for subject", "sub", event.Subject) + return deleteIndexedTokens(ctx, h.cache, subIndexPrefix+event.Subject) } return fmt.Errorf("revocation event has no identifiable token reference") } + +// --- Secondary index (jti / subject / session → token hash) --- +// +// The index reuses the claims Cache itself so it works with any Cache +// implementation (memory, Redis, user-provided). Index entries are stored as +// CommonClaims whose Extra["hashes"] holds the token hashes. + +const ( + jtiIndexPrefix = "idx:jti:" + subIndexPrefix = "idx:sub:" + sidIndexPrefix = "idx:sid:" +) + +// IndexClaims records secondary index entries for a cached token so later +// revocation events (by jti, subject, or session) can evict the exact cache +// entries instead of flushing the whole cache. Call it right after caching +// introspection results; ttl should match (or exceed) the cache entry's TTL. +func IndexClaims(ctx context.Context, c Cache, tokenHash string, claims *CommonClaims, ttl time.Duration) { + if c == nil || claims == nil || tokenHash == "" || ttl <= 0 { + return + } + if claims.JTI != "" { + entry := &CommonClaims{Extra: map[string]interface{}{"hashes": []string{tokenHash}}} + _ = c.Set(ctx, jtiIndexPrefix+claims.JTI, entry, ttl) + } + if claims.Subject != "" { + appendIndexHash(ctx, c, subIndexPrefix+claims.Subject, tokenHash, ttl) + } + if claims.SessionID != "" { + appendIndexHash(ctx, c, sidIndexPrefix+claims.SessionID, tokenHash, ttl) + } +} + +// appendIndexHash adds tokenHash to the index entry at key (read-modify-write). +func appendIndexHash(ctx context.Context, c Cache, key, tokenHash string, ttl time.Duration) { + hashes := indexHashes(ctx, c, key) + for _, h := range hashes { + if h == tokenHash { + return + } + } + hashes = append(hashes, tokenHash) + entry := &CommonClaims{Extra: map[string]interface{}{"hashes": hashes}} + _ = c.Set(ctx, key, entry, ttl) +} + +// indexHashes returns the token hashes recorded under an index key. +func indexHashes(ctx context.Context, c Cache, key string) []string { + entry, ok := c.Get(ctx, key) + if !ok || entry == nil || entry.Extra == nil { + return nil + } + switch v := entry.Extra["hashes"].(type) { + case []string: + return v + case []interface{}: // after a JSON round-trip (e.g. Redis) + out := make([]string, 0, len(v)) + for _, x := range v { + if s, ok := x.(string); ok { + out = append(out, s) + } + } + return out + default: + return nil + } +} + +// lookupIndexHash returns the single token hash stored under an index key. +func lookupIndexHash(ctx context.Context, c Cache, key string) (string, bool) { + hashes := indexHashes(ctx, c, key) + if len(hashes) == 0 { + return "", false + } + return hashes[0], true +} + +// deleteIndexedTokens evicts every token hash recorded under an index key, +// then removes the index entry itself. +func deleteIndexedTokens(ctx context.Context, c Cache, key string) error { + var firstErr error + for _, hash := range indexHashes(ctx, c, key) { + if err := c.Delete(ctx, hash); err != nil && firstErr == nil { + firstErr = err + } + } + if err := c.Delete(ctx, key); err != nil && firstErr == nil { + firstErr = err + } + return firstErr +} diff --git a/pkg/core/token/revocation_test.go b/pkg/core/token/revocation_test.go index 4175b7c..7fb1c9a 100644 --- a/pkg/core/token/revocation_test.go +++ b/pkg/core/token/revocation_test.go @@ -88,10 +88,15 @@ func TestRevocationHandler_MissingHMAC_WhenSecretRequired(t *testing.T) { } func TestRevocationHandler_RevokeByJTI(t *testing.T) { + ctx := context.Background() cache := NewMemoryCache() h := NewRevocationHandler(cache, "", nil) - _ = cache.Set(context.Background(), "jti-abc", &CommonClaims{Active: true}, time.Minute) + // Cache a token under its hash and index it by jti, as the guard does. + hash := HashToken("access-token-1") + claims := &CommonClaims{Active: true, JTI: "jti-abc", Subject: "alice"} + _ = cache.Set(ctx, hash, claims, time.Minute) + IndexClaims(ctx, cache, hash, claims, time.Minute) body, _ := json.Marshal(RevocationEvent{JTI: "jti-abc"}) req := httptest.NewRequest(http.MethodPost, "/revoke", bytes.NewReader(body)) @@ -101,6 +106,71 @@ func TestRevocationHandler_RevokeByJTI(t *testing.T) { if rr.Code != http.StatusNoContent { t.Fatalf("expected 204, got %d", rr.Code) } + if _, ok := cache.Get(ctx, hash); ok { + t.Error("jti revocation must evict the token's actual cache entry") + } +} + +func TestRevocationHandler_RevokeAllForSubject_Targeted(t *testing.T) { + ctx := context.Background() + cache := NewMemoryCache() + h := NewRevocationHandler(cache, "", nil) + + // Two tokens for alice, one for bob. + aliceHash1 := HashToken("alice-tok-1") + aliceHash2 := HashToken("alice-tok-2") + bobHash := HashToken("bob-tok") + alice1 := &CommonClaims{Active: true, Subject: "alice", JTI: "a1"} + alice2 := &CommonClaims{Active: true, Subject: "alice", JTI: "a2"} + bob := &CommonClaims{Active: true, Subject: "bob", JTI: "b1"} + for _, e := range []struct { + hash string + claims *CommonClaims + }{{aliceHash1, alice1}, {aliceHash2, alice2}, {bobHash, bob}} { + _ = cache.Set(ctx, e.hash, e.claims, time.Minute) + IndexClaims(ctx, cache, e.hash, e.claims, time.Minute) + } + + body, _ := json.Marshal(RevocationEvent{Subject: "alice", RevokeAll: true}) + req := httptest.NewRequest(http.MethodPost, "/revoke", bytes.NewReader(body)) + rr := httptest.NewRecorder() + h.ServeHTTP(rr, req) + + if rr.Code != http.StatusNoContent { + t.Fatalf("expected 204, got %d", rr.Code) + } + if _, ok := cache.Get(ctx, aliceHash1); ok { + t.Error("alice token 1 should be evicted") + } + if _, ok := cache.Get(ctx, aliceHash2); ok { + t.Error("alice token 2 should be evicted") + } + if _, ok := cache.Get(ctx, bobHash); !ok { + t.Error("bob's token must NOT be evicted by alice's revoke-all") + } +} + +func TestRevocationHandler_RevokeBySession(t *testing.T) { + ctx := context.Background() + cache := NewMemoryCache() + h := NewRevocationHandler(cache, "", nil) + + hash := HashToken("sess-tok") + claims := &CommonClaims{Active: true, Subject: "alice", SessionID: "sess-1"} + _ = cache.Set(ctx, hash, claims, time.Minute) + IndexClaims(ctx, cache, hash, claims, time.Minute) + + body, _ := json.Marshal(RevocationEvent{SessionID: "sess-1"}) + req := httptest.NewRequest(http.MethodPost, "/revoke", bytes.NewReader(body)) + rr := httptest.NewRecorder() + h.ServeHTTP(rr, req) + + if rr.Code != http.StatusNoContent { + t.Fatalf("expected 204, got %d", rr.Code) + } + if _, ok := cache.Get(ctx, hash); ok { + t.Error("session revocation must evict the token's cache entry") + } } func computeSig(secret string, body []byte) string { diff --git a/pkg/devkit/localas/server.go b/pkg/devkit/localas/server.go index 867b129..3d09011 100644 --- a/pkg/devkit/localas/server.go +++ b/pkg/devkit/localas/server.go @@ -179,12 +179,12 @@ func (s *Server) handleIntrospect(w http.ResponseWriter, r *http.Request) { } writeJSON(w, map[string]interface{}{ - "active": true, - "sub": entry.subject, - "scope": joinScopes(entry.scopes), - "acr": entry.acr, - "exp": entry.expiresAt.Unix(), - "iss": s.issuer, + "active": true, + "sub": entry.subject, + "scope": joinScopes(entry.scopes), + "acr": entry.acr, + "exp": entry.expiresAt.Unix(), + "iss": s.issuer, }) } diff --git a/pkg/devkit/simulator/policy.go b/pkg/devkit/simulator/policy.go index d2e686a..a511182 100644 --- a/pkg/devkit/simulator/policy.go +++ b/pkg/devkit/simulator/policy.go @@ -10,20 +10,20 @@ import ( // Request represents a simulated HTTP request for policy dry-run. type Request struct { - Method string - Path string - ACR string - AMR []string - Scopes []string - AuthAge time.Duration + Method string + Path string + ACR string + AMR []string + Scopes []string + AuthAge time.Duration } // Result is the simulation outcome. type Result struct { - Allowed bool - PolicyName string - Reason string - RequiredACR string + Allowed bool + PolicyName string + Reason string + RequiredACR string RequiredMaxAge int } @@ -52,9 +52,9 @@ func (s *Simulator) Simulate(req Request) (*Result, error) { } out := &Result{ - Allowed: result.Allowed, - Reason: result.Reason, - RequiredACR: result.RequiredACR, + Allowed: result.Allowed, + Reason: result.Reason, + RequiredACR: result.RequiredACR, RequiredMaxAge: result.RequiredMaxAge, } if result.MatchedPolicy != nil { diff --git a/pkg/middleware/echo/middleware_test.go b/pkg/middleware/echo/middleware_test.go index 988e83f..facabe1 100644 --- a/pkg/middleware/echo/middleware_test.go +++ b/pkg/middleware/echo/middleware_test.go @@ -10,9 +10,9 @@ import ( echofwk "github.com/labstack/echo/v4" "github.com/common-iam/iam/pkg/core/policy" - iamecho "github.com/common-iam/iam/pkg/middleware/echo" "github.com/common-iam/iam/pkg/devkit/localas" "github.com/common-iam/iam/pkg/devkit/tokenfactory" + iamecho "github.com/common-iam/iam/pkg/middleware/echo" "github.com/common-iam/iam/pkg/providers/generic" ) diff --git a/pkg/middleware/gin/middleware_test.go b/pkg/middleware/gin/middleware_test.go index 6512436..b086b80 100644 --- a/pkg/middleware/gin/middleware_test.go +++ b/pkg/middleware/gin/middleware_test.go @@ -10,9 +10,9 @@ import ( ginfwk "github.com/gin-gonic/gin" "github.com/common-iam/iam/pkg/core/policy" - iamgin "github.com/common-iam/iam/pkg/middleware/gin" "github.com/common-iam/iam/pkg/devkit/localas" "github.com/common-iam/iam/pkg/devkit/tokenfactory" + iamgin "github.com/common-iam/iam/pkg/middleware/gin" "github.com/common-iam/iam/pkg/providers/generic" ) diff --git a/pkg/middleware/grpc/interceptor_test.go b/pkg/middleware/grpc/interceptor_test.go index 36c6f03..97344d8 100644 --- a/pkg/middleware/grpc/interceptor_test.go +++ b/pkg/middleware/grpc/interceptor_test.go @@ -27,10 +27,10 @@ type fakeProvider struct { func (f *fakeProvider) Introspect(_ context.Context, _ string) (*token.CommonClaims, error) { return f.claims, f.err } -func (f *fakeProvider) JWKS(_ context.Context) ([]byte, error) { return nil, nil } -func (f *fakeProvider) RefreshConfig(_ context.Context) error { return nil } -func (f *fakeProvider) Name() string { return "fake" } -func (f *fakeProvider) Issuer() string { return "https://fake.as" } +func (f *fakeProvider) JWKS(_ context.Context) ([]byte, error) { return nil, nil } +func (f *fakeProvider) RefreshConfig(_ context.Context) error { return nil } +func (f *fakeProvider) Name() string { return "fake" } +func (f *fakeProvider) Issuer() string { return "https://fake.as" } var _ providers.Provider = (*fakeProvider)(nil) diff --git a/pkg/middleware/stdlib/middleware.go b/pkg/middleware/stdlib/middleware.go index aa61d70..f5081c9 100644 --- a/pkg/middleware/stdlib/middleware.go +++ b/pkg/middleware/stdlib/middleware.go @@ -16,10 +16,10 @@ const claimsKey contextKey = iota // Config configures the IAM middleware. type Config struct { - Provider providers.Provider - PolicyEngine *policy.Engine - Realm string - EnableDPoP bool + Provider providers.Provider + PolicyEngine *policy.Engine + Realm string + EnableDPoP bool } // Middleware returns a standard net/http middleware that: diff --git a/pkg/providers/auth0/adapter.go b/pkg/providers/auth0/adapter.go index 571f8c1..13ee656 100644 --- a/pkg/providers/auth0/adapter.go +++ b/pkg/providers/auth0/adapter.go @@ -28,8 +28,8 @@ type Config struct { // Adapter wraps the generic OIDC adapter for Auth0. type Adapter struct { - inner *generic.Adapter - cfg Config + inner *generic.Adapter + cfg Config } // New creates an Auth0 adapter. Call RefreshConfig before use. diff --git a/pkg/providers/claims_mapper.go b/pkg/providers/claims_mapper.go index 217ecb5..df98fe2 100644 --- a/pkg/providers/claims_mapper.go +++ b/pkg/providers/claims_mapper.go @@ -63,8 +63,8 @@ func MapToCommon(raw RawClaims) *token.CommonClaims { // Tenant c.TenantID = firstNonEmpty( stringClaim(raw, "tenant_id"), - stringClaim(raw, "tid"), // Azure AD style - stringClaim(raw, "org_id"), // Auth0 style + stringClaim(raw, "tid"), // Azure AD style + stringClaim(raw, "org_id"), // Auth0 style ) // Remaining claims go into Extra diff --git a/pkg/providers/generic/adapter.go b/pkg/providers/generic/adapter.go index 617a155..0e485c0 100644 --- a/pkg/providers/generic/adapter.go +++ b/pkg/providers/generic/adapter.go @@ -29,9 +29,9 @@ type Config struct { // Adapter is a generic OIDC provider adapter. // It auto-discovers endpoints via the OIDC discovery document. type Adapter struct { - cfg Config - discovery *providers.OIDCDiscovery - mu sync.RWMutex + cfg Config + discovery *providers.OIDCDiscovery + mu sync.RWMutex introspector *token.Introspector } diff --git a/pkg/providers/keycloak/adapter.go b/pkg/providers/keycloak/adapter.go index 7437ba9..9f420dd 100644 --- a/pkg/providers/keycloak/adapter.go +++ b/pkg/providers/keycloak/adapter.go @@ -6,8 +6,8 @@ import ( "net/http" "time" - "github.com/common-iam/iam/pkg/providers/generic" "github.com/common-iam/iam/pkg/core/token" + "github.com/common-iam/iam/pkg/providers/generic" ) // Config holds Keycloak-specific configuration. @@ -50,8 +50,8 @@ func New(cfg Config) *Adapter { return &Adapter{inner: inner, cfg: cfg} } -func (a *Adapter) Name() string { return "keycloak" } -func (a *Adapter) Issuer() string { return a.inner.Issuer() } +func (a *Adapter) Name() string { return "keycloak" } +func (a *Adapter) Issuer() string { return a.inner.Issuer() } func (a *Adapter) RefreshConfig(ctx context.Context) error { return a.inner.RefreshConfig(ctx) diff --git a/pkg/telemetry/logger.go b/pkg/telemetry/logger.go index f525fb2..8fa3c1c 100644 --- a/pkg/telemetry/logger.go +++ b/pkg/telemetry/logger.go @@ -39,15 +39,15 @@ func LoggerFromContext(ctx context.Context) *slog.Logger { // IAMEvent is a structured log event for IAM operations. type IAMEvent struct { - Event string // e.g. "token.validated", "stepup.challenge_issued" - TenantID string - Subject string - ACR string - Resource string - Method string - Allowed bool - Reason string - TraceID string + Event string // e.g. "token.validated", "stepup.challenge_issued" + TenantID string + Subject string + ACR string + Resource string + Method string + Allowed bool + Reason string + TraceID string } // Log emits the IAMEvent as a structured slog record. diff --git a/pkg/tenant/resolver.go b/pkg/tenant/resolver.go index a12b64a..823c5eb 100644 --- a/pkg/tenant/resolver.go +++ b/pkg/tenant/resolver.go @@ -15,6 +15,12 @@ type Resolver interface { // HeaderResolver reads the tenant ID from a request header. // Default header: X-Tenant-ID +// +// SECURITY: the header is client-controlled. Use HeaderResolver only behind a +// trusted edge (ingress/API gateway) that sets or strips it; never expose it +// directly to end users, or any caller can pick their tenant. The gateway +// guard independently verifies that the token issuer matches the resolved +// tenant's provider, but defense in depth starts at the edge. type HeaderResolver struct { Header string } @@ -84,6 +90,26 @@ func (p *PathResolver) Resolve(r *http.Request) (string, error) { return parts[p.Segment], nil } +// --- Static Resolver --- + +// StaticResolver always resolves to a fixed tenant ID. Use it as the last +// element of a ChainResolver to opt in to a default tenant for single-tenant +// deployments; without it, tenant resolution fails closed. +type StaticResolver struct { + TenantID string +} + +func NewStaticResolver(tenantID string) *StaticResolver { + return &StaticResolver{TenantID: tenantID} +} + +func (s *StaticResolver) Resolve(_ *http.Request) (string, error) { + if s.TenantID == "" { + return "", fmt.Errorf("static resolver has no tenant configured") + } + return s.TenantID, nil +} + // --- Chain Resolver --- // ChainResolver tries multiple resolvers in order, returning the first success. diff --git a/tests/integration/e2e_test.go b/tests/integration/e2e_test.go index c7d5397..dddd85d 100644 --- a/tests/integration/e2e_test.go +++ b/tests/integration/e2e_test.go @@ -87,7 +87,7 @@ func newE2EEnv(t *testing.T, pCfg *policy.Config, adminToken, cookieSecret strin guard := gateway.NewGuard(gateway.GuardConfig{ Registry: reg, - Resolver: tenant.NewChainResolver(tenant.NewHeaderResolver("X-Tenant-ID")), + Resolver: tenant.NewChainResolver(tenant.NewHeaderResolver("X-Tenant-ID"), tenant.NewStaticResolver("default")), PolicyEngine: eng, Realm: "Test", Cache: cache, @@ -299,7 +299,7 @@ func TestE2E_MultiTenantIsolation(t *testing.T) { eng := allowAllEngine() guard := gateway.NewGuard(gateway.GuardConfig{ Registry: reg, - Resolver: tenant.NewChainResolver(tenant.NewHeaderResolver("X-Tenant-ID")), + Resolver: tenant.NewChainResolver(tenant.NewHeaderResolver("X-Tenant-ID"), tenant.NewStaticResolver("default")), Realm: "Test", Upstream: upstream, PolicyEngine: eng, diff --git a/tests/integration/smoke_test.go b/tests/integration/smoke_test.go new file mode 100644 index 0000000..b718e7a --- /dev/null +++ b/tests/integration/smoke_test.go @@ -0,0 +1,93 @@ +package integration + +import ( + "fmt" + "net" + "net/http" + "os/exec" + "path/filepath" + "syscall" + "testing" + "time" +) + +// TestSmoke_ServiceBinaryBoots builds cmd/iam-service, boots it in LocalAS dev +// mode, waits for /health, and shuts it down gracefully. This guards the +// README quickstart: `make service` must actually start. +func TestSmoke_ServiceBinaryBoots(t *testing.T) { + if testing.Short() { + t.Skip("skipping binary smoke test in -short mode") + } + + repoRoot, err := filepath.Abs("../..") + if err != nil { + t.Fatal(err) + } + bin := filepath.Join(t.TempDir(), "iam-service") + + build := exec.Command("go", "build", "-o", bin, "./cmd/iam-service") + build.Dir = repoRoot + if out, err := build.CombinedOutput(); err != nil { + t.Fatalf("go build ./cmd/iam-service: %v\n%s", err, out) + } + + // Pick a free port, then hand it to the service. + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + addr := ln.Addr().String() + ln.Close() + + cmd := exec.Command(bin) + cmd.Dir = repoRoot // so the default IAM_POLICY_FILE path resolves + cmd.Env = append(cmd.Environ(), "IAM_ADDR="+addr) + if err := cmd.Start(); err != nil { + t.Fatalf("starting iam-service: %v", err) + } + defer cmd.Process.Kill() //nolint:errcheck + + healthURL := fmt.Sprintf("http://%s/health", addr) + deadline := time.Now().Add(15 * time.Second) + for { + resp, err := http.Get(healthURL) + if err == nil { + resp.Body.Close() + if resp.StatusCode == http.StatusOK { + break + } + } + if time.Now().After(deadline) { + t.Fatal("service did not become healthy within 15s") + } + time.Sleep(200 * time.Millisecond) + } + + // Unauthenticated request must be challenged with RFC 9470 headers. + resp, err := http.Get(fmt.Sprintf("http://%s/api/anything", addr)) + if err != nil { + t.Fatalf("guarded request: %v", err) + } + resp.Body.Close() + if resp.StatusCode != http.StatusUnauthorized { + t.Errorf("expected 401 for unauthenticated request, got %d", resp.StatusCode) + } + if resp.Header.Get("WWW-Authenticate") == "" { + t.Error("expected WWW-Authenticate challenge header") + } + + // Graceful shutdown on SIGTERM. + if err := cmd.Process.Signal(syscall.SIGTERM); err != nil { + t.Fatalf("sending SIGTERM: %v", err) + } + done := make(chan error, 1) + go func() { done <- cmd.Wait() }() + select { + case err := <-done: + if err != nil { + t.Errorf("service exited with error after SIGTERM: %v", err) + } + case <-time.After(10 * time.Second): + t.Error("service did not exit within 10s of SIGTERM") + } +}