diff --git a/docs/docs.go b/docs/docs.go index 2f8cea2..2b43bb9 100644 --- a/docs/docs.go +++ b/docs/docs.go @@ -326,6 +326,12 @@ const docTemplate = `{ "schema": { "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" } + }, + "503": { + "description": "Rate limiting temporarily unavailable", + "schema": { + "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" + } } } } @@ -383,6 +389,12 @@ const docTemplate = `{ "schema": { "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" } + }, + "503": { + "description": "Authentication temporarily unavailable", + "schema": { + "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" + } } } } @@ -420,6 +432,12 @@ const docTemplate = `{ "schema": { "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" } + }, + "503": { + "description": "Authentication temporarily unavailable", + "schema": { + "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" + } } } } @@ -477,6 +495,12 @@ const docTemplate = `{ "schema": { "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" } + }, + "503": { + "description": "Authentication temporarily unavailable", + "schema": { + "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" + } } } } @@ -529,6 +553,12 @@ const docTemplate = `{ "schema": { "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" } + }, + "503": { + "description": "Rate limiting temporarily unavailable", + "schema": { + "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" + } } } } @@ -569,6 +599,12 @@ const docTemplate = `{ "schema": { "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" } + }, + "503": { + "description": "Rate limiting temporarily unavailable", + "schema": { + "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" + } } } } @@ -621,6 +657,12 @@ const docTemplate = `{ "schema": { "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" } + }, + "503": { + "description": "Rate limiting temporarily unavailable", + "schema": { + "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" + } } } } @@ -673,6 +715,12 @@ const docTemplate = `{ "schema": { "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" } + }, + "503": { + "description": "Rate limiting temporarily unavailable", + "schema": { + "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" + } } } } diff --git a/docs/horizontal-scaling.md b/docs/horizontal-scaling.md index 0e7b2fd..a2adf0a 100644 --- a/docs/horizontal-scaling.md +++ b/docs/horizontal-scaling.md @@ -44,10 +44,11 @@ REDIS_URL=redis://:password@your-redis-host:6379/0 ``` Redis is used for rate limiting and access token blacklisting. If Redis is temporarily unavailable: -- Rate limiting is **skipped** (requests are allowed — graceful degradation) -- Blacklist checks are **skipped** (already-revoked tokens may be accepted until Redis recovers) +- Public auth routes under `/api/v1/auth/*` fail closed with `503 Service Unavailable` +- JWT-protected routes fail closed during blacklist checks with `503 Service Unavailable` +- Non-auth public routes such as `/health`, `/ready`, and `/swagger/*` remain available -This means Redis is **not** a hard availability dependency, but it should be highly available in production. +This means Redis is part of UniAuth's auth control plane. It should be highly available in production, and `/ready` should be used to remove unhealthy instances from rotation quickly. ### Shared JWT Secret @@ -74,11 +75,11 @@ When a user logs out or changes their password, UniAuth: 1. Revokes the session in PostgreSQL (invalidates the refresh token for all instances) 2. Blacklists the current access token JTI in Redis with the remaining TTL -The blacklist entry is checked by the `JWTAuth` middleware on every subsequent authenticated request across all instances. This ensures that a user who logs out cannot continue using their short-lived access token on a different instance. +The blacklist entry is checked by the `JWTAuth` middleware on every subsequent authenticated request across all instances. This ensures that a user who logs out cannot continue using their short-lived access token on a different instance. If Redis cannot be reached for that check, UniAuth now rejects the request with `503` instead of failing open. ### Access Token Expiry -Access tokens have a short TTL (default 15 minutes). Even without Redis, a blacklisted token expires naturally within that window. The Redis blacklist entry TTL matches the remaining token lifetime — it cleans itself up automatically. +Access tokens still have a short TTL (default 15 minutes), and the Redis blacklist entry TTL matches the remaining token lifetime so it cleans itself up automatically. During a Redis outage, protected requests fail closed until Redis recovers rather than accepting potentially revoked access tokens. --- diff --git a/docs/security-report.md b/docs/security-report.md index e20a1d3..f37268d 100644 --- a/docs/security-report.md +++ b/docs/security-report.md @@ -1,9 +1,9 @@ **1. Executive Summary** -This repo’s remaining main security problems are real and concentrated in operational hardening. The highest-risk remaining issue is now that Redis fail-open behavior can weaken revocation and throttling during outages, with weak JWT secret acceptance as the next tier of residual risk. +This repo’s remaining main security problems are now concentrated in configuration and deployment hardening. The highest-risk remaining issue is weak JWT secret acceptance, followed by the remaining lower-severity readiness and supply-chain hardening gaps. -Update (2026-03-28): findings 1, 2, 3, 4, 5, 6, and 7 have been fixed in the repo. The remaining findings below are still outstanding unless explicitly marked otherwise. +Update (2026-03-28): findings 1, 2, 3, 4, 5, 6, 7, and 8 have been fixed in the repo. The remaining findings below are still outstanding unless explicitly marked otherwise. -This report began as a read-only review of the codebase across auth/session logic, authorization/tenant boundaries, and config/deployment. Findings 1, 2, 3, 4, 5, 6, and 7 have since been remediated in the repo. `go test ./...` passed locally; `govulncheck` was not installed, so dependency-CVE verification remains open. +This report began as a read-only review of the codebase across auth/session logic, authorization/tenant boundaries, and config/deployment. Findings 1, 2, 3, 4, 5, 6, 7, and 8 have since been remediated in the repo. `go test ./...` passed locally; `govulncheck` was not installed, so dependency-CVE verification remains open. **2. Architecture and Attack Surface** - Entry points: public auth routes, `/health`, `/ready`, `/swagger/*`, and JWT-protected `/api/v1/*` in [router.go](/Users/osamamuhammed/uniauth/internal/api/router.go#L76). @@ -48,6 +48,8 @@ This report began as a read-only review of the codebase across auth/session logi Remediation implemented: `chi`'s unconditional `RealIP` middleware has been removed, a shared trusted-proxy-aware client IP resolver now populates request context before logging and rate limiting, auth handlers consume that resolved value for audit/session metadata, and `TRUSTED_PROXY_CIDRS` now explicitly defines which proxy CIDRs or IPs may supply forwarded client-IP headers. Regression coverage was added in [ratelimit_test.go](/Users/osamamuhammed/uniauth/internal/api/middleware/ratelimit_test.go#L1) and [config_test.go](/Users/osamamuhammed/uniauth/internal/config/config_test.go#L1). 8. **Revocation and throttling fail open when Redis is unavailable** — Medium, High confidence. Affected: [auth middleware](/Users/osamamuhammed/uniauth/internal/api/middleware/auth.go#L44), [auth service](/Users/osamamuhammed/uniauth/internal/service/auth.go#L313), [ratelimit middleware](/Users/osamamuhammed/uniauth/internal/api/middleware/ratelimit.go#L18). Evidence: blacklist lookup failures allow the request, blacklist writes are ignored, and rate limiting is skipped on Redis errors. Risk: during a Redis outage or partition, revoked access tokens remain usable until expiry and brute-force controls disappear. Preconditions: Redis outage/partition. Fix: decide explicitly which endpoints must fail closed, and at minimum fail closed for logout-sensitive auth decisions or degrade with alarms/feature flags. Verify with tests that simulate Redis failures. Category: session management / availability-hardening. + Status: Fixed. + Remediation implemented: Redis-backed auth controls now use a strict-auth fail-closed policy: `JWTAuth` returns `503` when blacklist lookups fail, `/api/v1/auth/*` requests return `503` when rate limiting cannot be enforced, and logout/password-change flows abort with `ErrServiceUnavailable` before mutating DB-backed auth state if the access-token blacklist write fails. Supporting docs were updated in [horizontal-scaling.md](/Users/osamamuhammed/uniauth/docs/horizontal-scaling.md#L1). Regression coverage was added in [auth_test.go](/Users/osamamuhammed/uniauth/internal/api/middleware/auth_test.go#L1), [ratelimit_test.go](/Users/osamamuhammed/uniauth/internal/api/middleware/ratelimit_test.go#L1), [auth_redis_test.go](/Users/osamamuhammed/uniauth/internal/service/auth_redis_test.go#L1), and [response_test.go](/Users/osamamuhammed/uniauth/internal/api/handlers/response_test.go#L1). 9. **JWT secret strength is documented but not enforced** — Medium, Medium confidence. Affected: [main.go](/Users/osamamuhammed/uniauth/cmd/server/main.go#L115), [config.go](/Users/osamamuhammed/uniauth/internal/config/config.go#L53), [docker-compose](/Users/osamamuhammed/uniauth/docker-compose.yml#L12). Evidence: startup only checks `JWT_SECRET != ""`; docs say “minimum 32 characters,” but code does not enforce it. Risk: weak or placeholder HMAC secrets make token forgery practical in misconfigured self-hosted deployments. Preconditions: operator sets a weak/default secret. Fix: reject short or known-placeholder secrets at startup. Verify with config validation tests. OWASP: A02 Cryptographic Failures. @@ -67,7 +69,7 @@ This report began as a read-only review of the codebase across auth/session logi - Completed hotfix 4: stop logging reset tokens and fail closed when reset-email delivery is unavailable. Complexity: Small. Regression risk: Low. Order: 4. - Completed hardening 1: add webhook URL validation and outbound SSRF guards. Complexity: Medium. Regression risk: Medium. Order: 5. - Completed hardening 2: replace naive forwarded-header trust with trusted-proxy-aware IP extraction and explicit proxy CIDR configuration. Complexity: Small-Medium. Regression risk: Medium. Order: 6. -- Short-term hardening 3: enforce JWT secret quality, sanitize `/ready`, and define Redis degraded-mode behavior. Complexity: Small. Regression risk: Low-Medium. Order: 7. +- Short-term hardening 3: enforce JWT secret quality and sanitize `/ready`. Complexity: Small. Regression risk: Low-Medium. Order: 7. - Completed structural improvement 1: introduce reusable authorization middleware/policy mapping instead of ad hoc handler checks. Complexity: High. Regression risk: Medium. - Completed structural improvement 2: add DB-level same-org enforcement for role assignments and auto-remove legacy cross-tenant links during migration. Complexity: Medium. Regression risk: Medium. - Structural improvement 3: pin CI actions/artifacts/images and add `govulncheck` to CI. Complexity: Small. Regression risk: Low. @@ -84,6 +86,7 @@ This report began as a read-only review of the codebase across auth/session logi - Add startup/config tests that weak or placeholder `JWT_SECRET` values fail fast. - Implemented: password reset email tests now confirm reset tokens and links are never written to logs in [email_test.go](/Users/osamamuhammed/uniauth/internal/service/email_test.go#L1). - Implemented: auth reset-flow tests now confirm failed delivery deletes reset-token rows while successful delivery preserves a usable token in [auth_reset_test.go](/Users/osamamuhammed/uniauth/internal/service/auth_reset_test.go#L1). +- Implemented: Redis outage regression tests now confirm blacklist lookup failures return `503`, auth-route rate limiting fails closed with `503`, non-auth public routes remain available, and logout/password-change flows do not mutate DB-backed auth state when blacklist writes fail in [auth_test.go](/Users/osamamuhammed/uniauth/internal/api/middleware/auth_test.go#L1), [ratelimit_test.go](/Users/osamamuhammed/uniauth/internal/api/middleware/ratelimit_test.go#L1), and [auth_redis_test.go](/Users/osamamuhammed/uniauth/internal/service/auth_redis_test.go#L1). - Add `govulncheck ./...` to CI and pin supply-chain artifacts. **7. Open Questions / Assumptions** @@ -91,4 +94,4 @@ This report began as a read-only review of the codebase across auth/session logi - I did not verify pod/network egress policy. Existing egress controls remain useful defense-in-depth on top of the new webhook validation and delivery guards. - I could not verify dependency-level vulns because `govulncheck` is not installed in this workspace. -Quick note: 4 security issues remain open in the report: findings 8 and 9, plus the 2 lower-severity hardening items. +Quick note: 3 security issues remain open in the report: finding 9 plus the 2 lower-severity hardening items. diff --git a/docs/swagger.json b/docs/swagger.json index ad5ec15..919a489 100644 --- a/docs/swagger.json +++ b/docs/swagger.json @@ -320,6 +320,12 @@ "schema": { "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" } + }, + "503": { + "description": "Rate limiting temporarily unavailable", + "schema": { + "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" + } } } } @@ -377,6 +383,12 @@ "schema": { "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" } + }, + "503": { + "description": "Authentication temporarily unavailable", + "schema": { + "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" + } } } } @@ -414,6 +426,12 @@ "schema": { "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" } + }, + "503": { + "description": "Authentication temporarily unavailable", + "schema": { + "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" + } } } } @@ -471,6 +489,12 @@ "schema": { "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" } + }, + "503": { + "description": "Authentication temporarily unavailable", + "schema": { + "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" + } } } } @@ -523,6 +547,12 @@ "schema": { "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" } + }, + "503": { + "description": "Rate limiting temporarily unavailable", + "schema": { + "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" + } } } } @@ -563,6 +593,12 @@ "schema": { "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" } + }, + "503": { + "description": "Rate limiting temporarily unavailable", + "schema": { + "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" + } } } } @@ -615,6 +651,12 @@ "schema": { "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" } + }, + "503": { + "description": "Rate limiting temporarily unavailable", + "schema": { + "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" + } } } } @@ -667,6 +709,12 @@ "schema": { "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" } + }, + "503": { + "description": "Rate limiting temporarily unavailable", + "schema": { + "$ref": "#/definitions/internal_api_handlers.SwaggerErrorResponse" + } } } } diff --git a/docs/swagger.yaml b/docs/swagger.yaml index 37c4a26..c7585cc 100644 --- a/docs/swagger.yaml +++ b/docs/swagger.yaml @@ -702,6 +702,10 @@ paths: description: Internal Server Error schema: $ref: '#/definitions/internal_api_handlers.SwaggerErrorResponse' + "503": + description: Rate limiting temporarily unavailable + schema: + $ref: '#/definitions/internal_api_handlers.SwaggerErrorResponse' summary: Authenticate user tags: - Auth @@ -738,6 +742,10 @@ paths: description: Internal Server Error schema: $ref: '#/definitions/internal_api_handlers.SwaggerErrorResponse' + "503": + description: Authentication temporarily unavailable + schema: + $ref: '#/definitions/internal_api_handlers.SwaggerErrorResponse' security: - BearerAuth: [] summary: Logout current session @@ -762,6 +770,10 @@ paths: description: Internal Server Error schema: $ref: '#/definitions/internal_api_handlers.SwaggerErrorResponse' + "503": + description: Authentication temporarily unavailable + schema: + $ref: '#/definitions/internal_api_handlers.SwaggerErrorResponse' security: - BearerAuth: [] summary: Logout all sessions @@ -799,6 +811,10 @@ paths: description: Internal Server Error schema: $ref: '#/definitions/internal_api_handlers.SwaggerErrorResponse' + "503": + description: Authentication temporarily unavailable + schema: + $ref: '#/definitions/internal_api_handlers.SwaggerErrorResponse' security: - BearerAuth: [] summary: Change password @@ -835,6 +851,10 @@ paths: description: Internal Server Error schema: $ref: '#/definitions/internal_api_handlers.SwaggerErrorResponse' + "503": + description: Rate limiting temporarily unavailable + schema: + $ref: '#/definitions/internal_api_handlers.SwaggerErrorResponse' summary: Confirm a password reset tags: - Auth @@ -862,6 +882,10 @@ paths: description: Bad Request schema: $ref: '#/definitions/internal_api_handlers.SwaggerErrorResponse' + "503": + description: Rate limiting temporarily unavailable + schema: + $ref: '#/definitions/internal_api_handlers.SwaggerErrorResponse' summary: Request a password reset email tags: - Auth @@ -897,6 +921,10 @@ paths: description: Internal Server Error schema: $ref: '#/definitions/internal_api_handlers.SwaggerErrorResponse' + "503": + description: Rate limiting temporarily unavailable + schema: + $ref: '#/definitions/internal_api_handlers.SwaggerErrorResponse' summary: Refresh access token tags: - Auth @@ -932,6 +960,10 @@ paths: description: Internal Server Error schema: $ref: '#/definitions/internal_api_handlers.SwaggerErrorResponse' + "503": + description: Rate limiting temporarily unavailable + schema: + $ref: '#/definitions/internal_api_handlers.SwaggerErrorResponse' summary: Register a new organization and admin user tags: - Auth diff --git a/internal/api/handlers/auth.go b/internal/api/handlers/auth.go index 9f0b2f2..417fe2a 100644 --- a/internal/api/handlers/auth.go +++ b/internal/api/handlers/auth.go @@ -27,6 +27,7 @@ func NewAuthHandler(authSvc *service.AuthService) *AuthHandler { // @Success 201 {object} RegisterResponse // @Failure 400 {object} SwaggerErrorResponse // @Failure 409 {object} SwaggerErrorResponse "Organization or email already exists" +// @Failure 503 {object} SwaggerErrorResponse "Rate limiting temporarily unavailable" // @Failure 500 {object} SwaggerErrorResponse // @Router /api/v1/auth/register [post] func (h *AuthHandler) Register(w http.ResponseWriter, r *http.Request) { @@ -76,6 +77,7 @@ func (h *AuthHandler) Register(w http.ResponseWriter, r *http.Request) { // @Failure 400 {object} SwaggerErrorResponse // @Failure 401 {object} SwaggerErrorResponse "Invalid credentials" // @Failure 403 {object} SwaggerErrorResponse "User or organization inactive" +// @Failure 503 {object} SwaggerErrorResponse "Rate limiting temporarily unavailable" // @Failure 500 {object} SwaggerErrorResponse // @Router /api/v1/auth/login [post] func (h *AuthHandler) Login(w http.ResponseWriter, r *http.Request) { @@ -127,6 +129,7 @@ func (h *AuthHandler) Login(w http.ResponseWriter, r *http.Request) { // @Success 200 {object} TokenPairResponse // @Failure 400 {object} SwaggerErrorResponse // @Failure 401 {object} SwaggerErrorResponse "Token invalid or expired" +// @Failure 503 {object} SwaggerErrorResponse "Rate limiting temporarily unavailable" // @Failure 500 {object} SwaggerErrorResponse // @Router /api/v1/auth/refresh [post] func (h *AuthHandler) Refresh(w http.ResponseWriter, r *http.Request) { @@ -165,6 +168,7 @@ func (h *AuthHandler) Refresh(w http.ResponseWriter, r *http.Request) { // @Success 200 {object} SwaggerMessageResponse // @Failure 400 {object} SwaggerErrorResponse // @Failure 401 {object} SwaggerErrorResponse +// @Failure 503 {object} SwaggerErrorResponse "Authentication temporarily unavailable" // @Failure 500 {object} SwaggerErrorResponse // @Security BearerAuth // @Router /api/v1/auth/logout [post] @@ -195,6 +199,7 @@ func (h *AuthHandler) Logout(w http.ResponseWriter, r *http.Request) { // @Produce json // @Success 200 {object} SwaggerMessageResponse // @Failure 401 {object} SwaggerErrorResponse +// @Failure 503 {object} SwaggerErrorResponse "Authentication temporarily unavailable" // @Failure 500 {object} SwaggerErrorResponse // @Security BearerAuth // @Router /api/v1/auth/logout-all [post] @@ -225,6 +230,7 @@ func (h *AuthHandler) LogoutAll(w http.ResponseWriter, r *http.Request) { // @Param body body ResetRequestBody true "Organization slug and email" // @Success 200 {object} SwaggerMessageResponse // @Failure 400 {object} SwaggerErrorResponse +// @Failure 503 {object} SwaggerErrorResponse "Rate limiting temporarily unavailable" // @Router /api/v1/auth/password/reset-request [post] func (h *AuthHandler) RequestPasswordReset(w http.ResponseWriter, r *http.Request) { var req struct { @@ -254,6 +260,7 @@ func (h *AuthHandler) RequestPasswordReset(w http.ResponseWriter, r *http.Reques // @Success 200 {object} SwaggerMessageResponse // @Failure 400 {object} SwaggerErrorResponse // @Failure 401 {object} SwaggerErrorResponse "Token invalid or expired" +// @Failure 503 {object} SwaggerErrorResponse "Rate limiting temporarily unavailable" // @Failure 500 {object} SwaggerErrorResponse // @Router /api/v1/auth/password/reset-confirm [post] func (h *AuthHandler) ConfirmPasswordReset(w http.ResponseWriter, r *http.Request) { @@ -288,6 +295,7 @@ func (h *AuthHandler) ConfirmPasswordReset(w http.ResponseWriter, r *http.Reques // @Success 200 {object} SwaggerMessageResponse // @Failure 400 {object} SwaggerErrorResponse // @Failure 401 {object} SwaggerErrorResponse "Wrong current password" +// @Failure 503 {object} SwaggerErrorResponse "Authentication temporarily unavailable" // @Failure 500 {object} SwaggerErrorResponse // @Security BearerAuth // @Router /api/v1/auth/password/change [put] diff --git a/internal/api/handlers/response.go b/internal/api/handlers/response.go index 7de9882..64c9f85 100644 --- a/internal/api/handlers/response.go +++ b/internal/api/handlers/response.go @@ -27,6 +27,8 @@ func handleServiceError(w http.ResponseWriter, err error) { writeError(w, http.StatusNotFound, "not found") case errors.Is(err, domain.ErrAlreadyExists): writeError(w, http.StatusConflict, err.Error()) + case errors.Is(err, domain.ErrServiceUnavailable): + writeError(w, http.StatusServiceUnavailable, "service temporarily unavailable") case errors.Is(err, domain.ErrInvalidCredentials): writeError(w, http.StatusUnauthorized, "invalid credentials") case errors.Is(err, domain.ErrUnauthorized): diff --git a/internal/api/handlers/response_test.go b/internal/api/handlers/response_test.go index 6d6309c..410657e 100644 --- a/internal/api/handlers/response_test.go +++ b/internal/api/handlers/response_test.go @@ -30,6 +30,11 @@ func TestHandleServiceError(t *testing.T) { err: domain.ErrAlreadyExists, wantStatus: http.StatusConflict, }, + { + name: "ErrServiceUnavailable → 503", + err: domain.ErrServiceUnavailable, + wantStatus: http.StatusServiceUnavailable, + }, { name: "ErrInvalidCredentials → 401", err: domain.ErrInvalidCredentials, diff --git a/internal/api/middleware/auth.go b/internal/api/middleware/auth.go index 5d094ba..1892862 100644 --- a/internal/api/middleware/auth.go +++ b/internal/api/middleware/auth.go @@ -2,18 +2,23 @@ package middleware import ( "context" + "log/slog" "net/http" + "reflect" "strings" "time" "github.com/google/uuid" "github.com/osama1998h/uniauth/internal/domain" - "github.com/osama1998h/uniauth/internal/repository/cache" "github.com/osama1998h/uniauth/internal/service" "github.com/osama1998h/uniauth/pkg/token" ) +type tokenBlacklistChecker interface { + IsTokenBlacklisted(ctx context.Context, tokenID string) (bool, error) +} + type contextKey string const ( @@ -27,7 +32,7 @@ const ( // JWTAuth extracts and validates a Bearer JWT from the Authorization header. // It also checks the token blacklist in Redis so that revoked tokens are rejected // across all instances (required for correct horizontal scaling behaviour). -func JWTAuth(maker *token.Maker, c *cache.Cache) func(next http.Handler) http.Handler { +func JWTAuth(maker *token.Maker, c tokenBlacklistChecker, logger *slog.Logger) func(next http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { tokenStr := extractBearer(r) @@ -42,12 +47,17 @@ func JWTAuth(maker *token.Maker, c *cache.Cache) func(next http.Handler) http.Ha return } - // Check Redis blacklist. If the cache is nil (e.g. in tests) or Redis - // is unavailable, allow the request (graceful degradation — same - // pattern used by the rate limiter). - if c != nil { + // Allow a nil checker in tests, but fail closed on runtime Redis errors. + if !isNilTokenBlacklistChecker(c) { blacklisted, err := c.IsTokenBlacklisted(r.Context(), claims.TokenID.String()) - if err == nil && blacklisted { + if err != nil { + if logger != nil { + logger.WarnContext(r.Context(), "authentication unavailable during token blacklist lookup", "path", r.URL.Path, "error", err) + } + writeServiceUnavailable(w, "authentication temporarily unavailable") + return + } + if blacklisted { writeUnauthorized(w, "token has been revoked") return } @@ -125,3 +135,23 @@ func writeUnauthorized(w http.ResponseWriter, msg string) { w.WriteHeader(http.StatusUnauthorized) _, _ = w.Write([]byte(`{"error":"` + msg + `"}`)) } + +func writeServiceUnavailable(w http.ResponseWriter, msg string) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusServiceUnavailable) + _, _ = w.Write([]byte(`{"error":"` + msg + `"}`)) +} + +func isNilTokenBlacklistChecker(checker tokenBlacklistChecker) bool { + if checker == nil { + return true + } + + value := reflect.ValueOf(checker) + switch value.Kind() { + case reflect.Chan, reflect.Func, reflect.Interface, reflect.Map, reflect.Pointer, reflect.Slice: + return value.IsNil() + default: + return false + } +} diff --git a/internal/api/middleware/auth_test.go b/internal/api/middleware/auth_test.go index 8cee102..92983a4 100644 --- a/internal/api/middleware/auth_test.go +++ b/internal/api/middleware/auth_test.go @@ -2,6 +2,7 @@ package middleware import ( "context" + "errors" "net/http" "net/http/httptest" "testing" @@ -10,6 +11,7 @@ import ( "github.com/golang-jwt/jwt/v5" "github.com/google/uuid" + "github.com/osama1998h/uniauth/internal/repository/cache" "github.com/osama1998h/uniauth/pkg/token" ) @@ -47,6 +49,15 @@ func nextHandler(reached *bool) http.Handler { }) } +type fakeTokenBlacklistChecker struct { + blacklisted bool + err error +} + +func (f *fakeTokenBlacklistChecker) IsTokenBlacklisted(context.Context, string) (bool, error) { + return f.blacklisted, f.err +} + func TestJWTAuth_ValidToken(t *testing.T) { maker := newTestMaker() userID := uuid.New() @@ -58,7 +69,7 @@ func TestJWTAuth_ValidToken(t *testing.T) { } reached := false - handler := JWTAuth(maker, nil)(nextHandler(&reached)) + handler := JWTAuth(maker, nil, nil)(nextHandler(&reached)) r := httptest.NewRequest("GET", "/", nil) r.Header.Set("Authorization", "Bearer "+tokenStr) @@ -74,6 +85,35 @@ func TestJWTAuth_ValidToken(t *testing.T) { } } +func TestJWTAuth_ValidToken_WithTypedNilBlacklistChecker(t *testing.T) { + maker := newTestMaker() + userID := uuid.New() + orgID := uuid.New() + + tokenStr, _, err := maker.CreateAccessToken(userID, orgID) + if err != nil { + t.Fatalf("create token: %v", err) + } + + var redisCache *cache.Cache + + reached := false + handler := JWTAuth(maker, redisCache, nil)(nextHandler(&reached)) + + r := httptest.NewRequest("GET", "/", nil) + r.Header.Set("Authorization", "Bearer "+tokenStr) + w := httptest.NewRecorder() + + handler.ServeHTTP(w, r) + + if w.Code != http.StatusOK { + t.Errorf("expected 200, got %d", w.Code) + } + if !reached { + t.Error("next handler was not called for valid token with typed-nil checker") + } +} + func TestJWTAuth_RejectsRefreshToken(t *testing.T) { maker := newTestMaker() userID := uuid.New() @@ -85,7 +125,7 @@ func TestJWTAuth_RejectsRefreshToken(t *testing.T) { } reached := false - handler := JWTAuth(maker, nil)(nextHandler(&reached)) + handler := JWTAuth(maker, nil, nil)(nextHandler(&reached)) r := httptest.NewRequest("GET", "/", nil) r.Header.Set("Authorization", "Bearer "+tokenStr) @@ -108,7 +148,7 @@ func TestJWTAuth_RejectsLegacyUntypedToken(t *testing.T) { tokenStr := createLegacyUntypedToken(t, userID, orgID, 15*time.Minute) reached := false - handler := JWTAuth(maker, nil)(nextHandler(&reached)) + handler := JWTAuth(maker, nil, nil)(nextHandler(&reached)) r := httptest.NewRequest("GET", "/", nil) r.Header.Set("Authorization", "Bearer "+tokenStr) @@ -138,7 +178,7 @@ func TestJWTAuth_ValidToken_InjectsContext(t *testing.T) { w.WriteHeader(http.StatusOK) }) - handler := JWTAuth(maker, nil)(captureHandler) + handler := JWTAuth(maker, nil, nil)(captureHandler) r := httptest.NewRequest("GET", "/", nil) r.Header.Set("Authorization", "Bearer "+tokenStr) w := httptest.NewRecorder() @@ -155,7 +195,7 @@ func TestJWTAuth_ValidToken_InjectsContext(t *testing.T) { func TestJWTAuth_MissingAuthorizationHeader(t *testing.T) { maker := newTestMaker() reached := false - handler := JWTAuth(maker, nil)(nextHandler(&reached)) + handler := JWTAuth(maker, nil, nil)(nextHandler(&reached)) r := httptest.NewRequest("GET", "/", nil) w := httptest.NewRecorder() @@ -172,7 +212,7 @@ func TestJWTAuth_MissingAuthorizationHeader(t *testing.T) { func TestJWTAuth_MalformedBearer(t *testing.T) { maker := newTestMaker() reached := false - handler := JWTAuth(maker, nil)(nextHandler(&reached)) + handler := JWTAuth(maker, nil, nil)(nextHandler(&reached)) r := httptest.NewRequest("GET", "/", nil) // Missing "Bearer " prefix @@ -198,7 +238,7 @@ func TestJWTAuth_ExpiredToken(t *testing.T) { reached := false // Use the valid maker for the middleware (correct secret, but token is expired) - handler := JWTAuth(validMaker, nil)(nextHandler(&reached)) + handler := JWTAuth(validMaker, nil, nil)(nextHandler(&reached)) r := httptest.NewRequest("GET", "/", nil) r.Header.Set("Authorization", "Bearer "+tokenStr) @@ -222,7 +262,7 @@ func TestJWTAuth_WrongSecret(t *testing.T) { tokenStr, _, _ := signingMaker.CreateAccessToken(userID, orgID) reached := false - handler := JWTAuth(verifyingMaker, nil)(nextHandler(&reached)) + handler := JWTAuth(verifyingMaker, nil, nil)(nextHandler(&reached)) r := httptest.NewRequest("GET", "/", nil) r.Header.Set("Authorization", "Bearer "+tokenStr) @@ -240,7 +280,7 @@ func TestJWTAuth_WrongSecret(t *testing.T) { func TestJWTAuth_InvalidTokenString(t *testing.T) { maker := newTestMaker() reached := false - handler := JWTAuth(maker, nil)(nextHandler(&reached)) + handler := JWTAuth(maker, nil, nil)(nextHandler(&reached)) r := httptest.NewRequest("GET", "/", nil) r.Header.Set("Authorization", "Bearer not-a-real-jwt") @@ -255,6 +295,60 @@ func TestJWTAuth_InvalidTokenString(t *testing.T) { } } +func TestJWTAuth_BlacklistedToken(t *testing.T) { + maker := newTestMaker() + userID := uuid.New() + orgID := uuid.New() + + tokenStr, _, err := maker.CreateAccessToken(userID, orgID) + if err != nil { + t.Fatalf("create token: %v", err) + } + + reached := false + handler := JWTAuth(maker, &fakeTokenBlacklistChecker{blacklisted: true}, nil)(nextHandler(&reached)) + + r := httptest.NewRequest("GET", "/", nil) + r.Header.Set("Authorization", "Bearer "+tokenStr) + w := httptest.NewRecorder() + + handler.ServeHTTP(w, r) + + if w.Code != http.StatusUnauthorized { + t.Fatalf("expected 401, got %d", w.Code) + } + if reached { + t.Fatal("next handler should not be called for blacklisted token") + } +} + +func TestJWTAuth_BlacklistLookupFailure(t *testing.T) { + maker := newTestMaker() + userID := uuid.New() + orgID := uuid.New() + + tokenStr, _, err := maker.CreateAccessToken(userID, orgID) + if err != nil { + t.Fatalf("create token: %v", err) + } + + reached := false + handler := JWTAuth(maker, &fakeTokenBlacklistChecker{err: errors.New("redis unavailable")}, nil)(nextHandler(&reached)) + + r := httptest.NewRequest("GET", "/", nil) + r.Header.Set("Authorization", "Bearer "+tokenStr) + w := httptest.NewRecorder() + + handler.ServeHTTP(w, r) + + if w.Code != http.StatusServiceUnavailable { + t.Fatalf("expected 503, got %d", w.Code) + } + if reached { + t.Fatal("next handler should not be called when blacklist lookup fails") + } +} + func TestGetUserID(t *testing.T) { t.Run("present in context", func(t *testing.T) { id := uuid.New() diff --git a/internal/api/middleware/ratelimit.go b/internal/api/middleware/ratelimit.go index c8a5f38..18cf285 100644 --- a/internal/api/middleware/ratelimit.go +++ b/internal/api/middleware/ratelimit.go @@ -3,8 +3,10 @@ package middleware import ( "context" "fmt" + "log/slog" "net/http" "reflect" + "strings" "time" ) @@ -13,7 +15,7 @@ type rateLimitCounter interface { } // RateLimit returns a middleware that limits requests per IP per minute. -func RateLimit(redisCache rateLimitCounter, requestsPerMinute int) func(next http.Handler) http.Handler { +func RateLimit(redisCache rateLimitCounter, requestsPerMinute int, logger *slog.Logger) func(next http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { ip := ClientIP(r) @@ -26,7 +28,13 @@ func RateLimit(redisCache rateLimitCounter, requestsPerMinute int) func(next htt count, err := redisCache.IncrRateLimit(r.Context(), key, time.Minute) if err != nil { - // Redis unavailable — allow the request but log + if logger != nil { + logger.WarnContext(r.Context(), "rate limiting unavailable", "path", r.URL.Path, "client_ip", ip, "error", err) + } + if shouldFailClosedOnRateLimit(r.URL.Path) { + writeServiceUnavailable(w, "rate limiting temporarily unavailable") + return + } next.ServeHTTP(w, r) return } @@ -44,6 +52,10 @@ func RateLimit(redisCache rateLimitCounter, requestsPerMinute int) func(next htt } } +func shouldFailClosedOnRateLimit(path string) bool { + return strings.HasPrefix(path, "/api/v1/auth/") +} + func isNilRateLimitCounter(counter rateLimitCounter) bool { if counter == nil { return true diff --git a/internal/api/middleware/ratelimit_test.go b/internal/api/middleware/ratelimit_test.go index eaa6c1f..a69cf66 100644 --- a/internal/api/middleware/ratelimit_test.go +++ b/internal/api/middleware/ratelimit_test.go @@ -2,6 +2,7 @@ package middleware import ( "context" + "errors" "net" "net/http" "net/http/httptest" @@ -100,13 +101,13 @@ func TestRateLimitUsesResolvedClientIP(t *testing.T) { t.Parallel() counter := newFakeRateLimitCounter() - handler := PopulateClientIP(NewClientIPResolver(nil))(RateLimit(counter, 10)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + handler := PopulateClientIP(NewClientIPResolver(nil))(RateLimit(counter, 10, nil)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) }))) requests := []*http.Request{ - newRateLimitedRequest("203.0.113.99:4000", map[string]string{"X-Forwarded-For": "198.51.100.1"}), - newRateLimitedRequest("203.0.113.99:4000", map[string]string{"X-Forwarded-For": "198.51.100.2"}), + newRateLimitedRequest("/", "203.0.113.99:4000", map[string]string{"X-Forwarded-For": "198.51.100.1"}), + newRateLimitedRequest("/", "203.0.113.99:4000", map[string]string{"X-Forwarded-For": "198.51.100.2"}), } for _, req := range requests { @@ -125,13 +126,13 @@ func TestRateLimitUsesResolvedClientIP(t *testing.T) { t.Parallel() counter := newFakeRateLimitCounter() - handler := PopulateClientIP(NewClientIPResolver(mustParseCIDRs(t, "10.0.0.0/8")))(RateLimit(counter, 10)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + handler := PopulateClientIP(NewClientIPResolver(mustParseCIDRs(t, "10.0.0.0/8")))(RateLimit(counter, 10, nil)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) }))) requests := []*http.Request{ - newRateLimitedRequest("10.0.0.5:4000", map[string]string{"X-Forwarded-For": "198.51.100.10"}), - newRateLimitedRequest("10.0.0.5:4000", map[string]string{"X-Forwarded-For": "198.51.100.11"}), + newRateLimitedRequest("/", "10.0.0.5:4000", map[string]string{"X-Forwarded-For": "198.51.100.10"}), + newRateLimitedRequest("/", "10.0.0.5:4000", map[string]string{"X-Forwarded-For": "198.51.100.11"}), } for _, req := range requests { @@ -151,7 +152,7 @@ func TestRateLimitAllowsTypedNilCounter(t *testing.T) { t.Parallel() var counter *fakeRateLimitCounter - handler := RateLimit(counter, 10)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + handler := RateLimit(counter, 10, nil)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusNoContent) })) @@ -162,6 +163,65 @@ func TestRateLimitAllowsTypedNilCounter(t *testing.T) { } } +func TestRateLimitReturnsServiceUnavailableForAuthRoutesOnRedisError(t *testing.T) { + t.Parallel() + + handler := RateLimit(errorRateLimitCounter{}, 10, nil)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + })) + + for _, path := range []string{"/api/v1/auth/login", "/api/v1/auth/password/reset-confirm"} { + t.Run(path, func(t *testing.T) { + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, newRateLimitedRequest(path, "203.0.113.20:4000", nil)) + + if rec.Code != http.StatusServiceUnavailable { + t.Fatalf("expected 503 for %s, got %d", path, rec.Code) + } + }) + } +} + +func TestRateLimitAllowsNonAuthRoutesOnRedisError(t *testing.T) { + t.Parallel() + + handler := RateLimit(errorRateLimitCounter{}, 10, nil)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + })) + + for _, path := range []string{"/health", "/swagger/index.html"} { + t.Run(path, func(t *testing.T) { + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, newRateLimitedRequest(path, "203.0.113.21:4000", nil)) + + if rec.Code != http.StatusOK { + t.Fatalf("expected 200 for %s, got %d", path, rec.Code) + } + }) + } +} + +func TestRateLimitReturnsTooManyRequestsWhenLimitExceeded(t *testing.T) { + t.Parallel() + + counter := newFakeRateLimitCounter() + handler := RateLimit(counter, 1, nil)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + })) + + first := httptest.NewRecorder() + handler.ServeHTTP(first, newRateLimitedRequest("/api/v1/auth/login", "203.0.113.25:4000", nil)) + if first.Code != http.StatusOK { + t.Fatalf("expected first request to pass, got %d", first.Code) + } + + second := httptest.NewRecorder() + handler.ServeHTTP(second, newRateLimitedRequest("/api/v1/auth/login", "203.0.113.25:4000", nil)) + if second.Code != http.StatusTooManyRequests { + t.Fatalf("expected second request to be rate limited, got %d", second.Code) + } +} + type fakeRateLimitCounter struct { counts map[string]int64 keys []string @@ -177,8 +237,14 @@ func (f *fakeRateLimitCounter) IncrRateLimit(_ context.Context, key string, _ ti return f.counts[key], nil } -func newRateLimitedRequest(remoteAddr string, headers map[string]string) *http.Request { - req := httptest.NewRequest(http.MethodGet, "/", nil) +type errorRateLimitCounter struct{} + +func (errorRateLimitCounter) IncrRateLimit(context.Context, string, time.Duration) (int64, error) { + return 0, errors.New("redis unavailable") +} + +func newRateLimitedRequest(path, remoteAddr string, headers map[string]string) *http.Request { + req := httptest.NewRequest(http.MethodGet, path, nil) req.RemoteAddr = remoteAddr for key, value := range headers { req.Header.Set(key, value) diff --git a/internal/api/router.go b/internal/api/router.go index f9dfbe3..8468dde 100644 --- a/internal/api/router.go +++ b/internal/api/router.go @@ -64,7 +64,7 @@ func NewRouter( AllowCredentials: true, MaxAge: 300, })) - r.Use(middleware.RateLimit(redisCache, cfg.Auth.RateLimitPerMinute)) + r.Use(middleware.RateLimit(redisCache, cfg.Auth.RateLimitPerMinute, logger)) // Health r.Get("/health", healthH.Live) @@ -87,7 +87,7 @@ func NewRouter( // Auth — requires JWT r.Group(func(r chi.Router) { - r.Use(middleware.JWTAuth(tokenMaker, redisCache)) + r.Use(middleware.JWTAuth(tokenMaker, redisCache, logger)) r.Post("/logout", authH.Logout) r.Post("/logout-all", authH.LogoutAll) r.Put("/password/change", authH.ChangePassword) @@ -96,7 +96,7 @@ func NewRouter( // All routes below require JWT auth r.Group(func(r chi.Router) { - r.Use(middleware.JWTAuth(tokenMaker, redisCache)) + r.Use(middleware.JWTAuth(tokenMaker, redisCache, logger)) // Users r.Route("/users", func(r chi.Router) { diff --git a/internal/domain/errors.go b/internal/domain/errors.go index 9e42378..ff1489b 100644 --- a/internal/domain/errors.go +++ b/internal/domain/errors.go @@ -4,17 +4,18 @@ import "errors" // Sentinel errors used across the service layer. var ( - ErrNotFound = errors.New("not found") - ErrAlreadyExists = errors.New("already exists") + ErrNotFound = errors.New("not found") + ErrAlreadyExists = errors.New("already exists") + ErrServiceUnavailable = errors.New("service unavailable") ErrInvalidCredentials = errors.New("invalid credentials") - ErrUnauthorized = errors.New("unauthorized") - ErrForbidden = errors.New("forbidden") - ErrTokenExpired = errors.New("token expired") - ErrTokenInvalid = errors.New("token invalid") - ErrUserInactive = errors.New("user account is inactive") - ErrOrgInactive = errors.New("organization is inactive") - ErrAPIKeyRevoked = errors.New("api key has been revoked") - ErrAPIKeyExpired = errors.New("api key has expired") - ErrWeakPassword = errors.New("password does not meet requirements") - ErrInvalidInput = errors.New("invalid input") + ErrUnauthorized = errors.New("unauthorized") + ErrForbidden = errors.New("forbidden") + ErrTokenExpired = errors.New("token expired") + ErrTokenInvalid = errors.New("token invalid") + ErrUserInactive = errors.New("user account is inactive") + ErrOrgInactive = errors.New("organization is inactive") + ErrAPIKeyRevoked = errors.New("api key has been revoked") + ErrAPIKeyExpired = errors.New("api key has expired") + ErrWeakPassword = errors.New("password does not meet requirements") + ErrInvalidInput = errors.New("invalid input") ) diff --git a/internal/service/auth.go b/internal/service/auth.go index ff15920..f8f38e6 100644 --- a/internal/service/auth.go +++ b/internal/service/auth.go @@ -14,16 +14,19 @@ import ( "github.com/osama1998h/uniauth/internal/config" "github.com/osama1998h/uniauth/internal/domain" - "github.com/osama1998h/uniauth/internal/repository/cache" db "github.com/osama1998h/uniauth/internal/repository/postgres" "github.com/osama1998h/uniauth/pkg/token" ) +type accessTokenBlacklistWriter interface { + BlacklistToken(ctx context.Context, tokenID string, ttl time.Duration) error +} + // AuthService handles authentication logic. type AuthService struct { store *db.Store tokenMaker *token.Maker - cache *cache.Cache + cache accessTokenBlacklistWriter auditSvc *AuditService webhookSvc *WebhookService emailSvc *EmailService @@ -34,7 +37,7 @@ type AuthService struct { func NewAuthService( store *db.Store, tokenMaker *token.Maker, - c *cache.Cache, + c accessTokenBlacklistWriter, auditSvc *AuditService, webhookSvc *WebhookService, emailSvc *EmailService, @@ -215,7 +218,9 @@ func (s *AuthService) Logout(ctx context.Context, refreshToken string, accessTok sess, err := s.store.GetSessionByTokenHash(ctx, tokenHash) if err != nil { if errors.Is(err, domain.ErrNotFound) { - s.blacklistAccessToken(ctx, accessTokenID, accessTokenExpiresAt) + if err := s.blacklistAccessToken(ctx, accessTokenID, accessTokenExpiresAt); err != nil { + return err + } return nil // already gone } return err @@ -223,20 +228,23 @@ func (s *AuthService) Logout(ctx context.Context, refreshToken string, accessTok if sess.UserID != claims.UserID { return domain.ErrTokenInvalid } + if err := s.blacklistAccessToken(ctx, accessTokenID, accessTokenExpiresAt); err != nil { + return err + } if !sess.IsValid() { - s.blacklistAccessToken(ctx, accessTokenID, accessTokenExpiresAt) return nil } if err := s.store.RevokeSession(ctx, sess.ID); err != nil { return err } - s.blacklistAccessToken(ctx, accessTokenID, accessTokenExpiresAt) return nil } // LogoutAll revokes all active sessions for a user and blacklists the current access token. func (s *AuthService) LogoutAll(ctx context.Context, userID uuid.UUID, accessTokenID uuid.UUID, accessTokenExpiresAt time.Time) error { - s.blacklistAccessToken(ctx, accessTokenID, accessTokenExpiresAt) + if err := s.blacklistAccessToken(ctx, accessTokenID, accessTokenExpiresAt); err != nil { + return err + } return s.store.RevokeAllUserSessions(ctx, userID) } @@ -328,26 +336,33 @@ func (s *AuthService) ChangePassword(ctx context.Context, orgID, userID uuid.UUI return fmt.Errorf("hash password: %w", err) } + if err := s.blacklistAccessToken(ctx, accessTokenID, accessTokenExpiresAt); err != nil { + return err + } + if err := s.store.UpdateUserPassword(ctx, userID, string(hashed)); err != nil { return fmt.Errorf("update password: %w", err) } // Revoke all existing sessions so old tokens can't be reused _ = s.store.RevokeAllUserSessions(ctx, userID) - s.blacklistAccessToken(ctx, accessTokenID, accessTokenExpiresAt) s.auditSvc.Log(&domain.AuditLog{OrgID: &user.OrgID, UserID: &userID, Action: domain.AuditActionPasswordChanged}) return nil } // blacklistAccessToken stores the token JTI in Redis so that all instances -// reject it immediately. Best-effort: errors are silently ignored because the -// short access token TTL (default 15 min) is already a reasonable security bound. -func (s *AuthService) blacklistAccessToken(ctx context.Context, tokenID uuid.UUID, expiresAt time.Time) { +// reject it immediately. Failures are surfaced so callers can abort before +// mutating DB-backed auth state when revocation guarantees cannot be enforced. +func (s *AuthService) blacklistAccessToken(ctx context.Context, tokenID uuid.UUID, expiresAt time.Time) error { ttl := time.Until(expiresAt) - if ttl > 0 { - _ = s.cache.BlacklistToken(ctx, tokenID.String(), ttl) + if s.cache == nil || ttl <= 0 { + return nil + } + if err := s.cache.BlacklistToken(ctx, tokenID.String(), ttl); err != nil { + return fmt.Errorf("%w: blacklist access token: %v", domain.ErrServiceUnavailable, err) } + return nil } // issueTokenPair creates access+refresh tokens and persists the session. diff --git a/internal/service/auth_redis_test.go b/internal/service/auth_redis_test.go new file mode 100644 index 0000000..e00345d --- /dev/null +++ b/internal/service/auth_redis_test.go @@ -0,0 +1,241 @@ +package service + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/google/uuid" + "golang.org/x/crypto/bcrypt" + + "github.com/osama1998h/uniauth/internal/domain" + db "github.com/osama1998h/uniauth/internal/repository/postgres" + "github.com/osama1998h/uniauth/internal/testutil" + "github.com/osama1998h/uniauth/pkg/token" +) + +const authRedisTestJWTSecret = "supersecretkey-at-least-32-chars!!" + +type fakeBlacklistWriter struct { + err error +} + +func (f fakeBlacklistWriter) BlacklistToken(context.Context, string, time.Duration) error { + return f.err +} + +func TestAuthServiceLogoutBlacklistFailureDoesNotRevokeSession(t *testing.T) { + store := testutil.RequireTestStore(t) + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + + svc := newRedisSensitiveAuthService(store, fakeBlacklistWriter{err: errors.New("redis unavailable")}) + user, _, refreshToken := createAuthRedisTestUserAndSession(t, ctx, store, svc.tokenMaker, "logout-blacklist-failure") + accessTokenID, accessTokenExpiry := createAccessTokenMetadata(t, svc.tokenMaker, user.ID, user.OrgID) + + if err := svc.Logout(ctx, refreshToken, accessTokenID, accessTokenExpiry); !errors.Is(err, domain.ErrServiceUnavailable) { + t.Fatalf("expected ErrServiceUnavailable, got %v", err) + } + + assertSessionActive(t, ctx, store, refreshToken) +} + +func TestAuthServiceLogoutRevokesSessionWhenBlacklistSucceeds(t *testing.T) { + store := testutil.RequireTestStore(t) + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + + svc := newRedisSensitiveAuthService(store, fakeBlacklistWriter{}) + user, _, refreshToken := createAuthRedisTestUserAndSession(t, ctx, store, svc.tokenMaker, "logout-success") + accessTokenID, accessTokenExpiry := createAccessTokenMetadata(t, svc.tokenMaker, user.ID, user.OrgID) + + if err := svc.Logout(ctx, refreshToken, accessTokenID, accessTokenExpiry); err != nil { + t.Fatalf("Logout() error = %v", err) + } + + assertSessionRevoked(t, ctx, store, refreshToken) +} + +func TestAuthServiceLogoutAllBlacklistFailureDoesNotRevokeSessions(t *testing.T) { + store := testutil.RequireTestStore(t) + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + + svc := newRedisSensitiveAuthService(store, fakeBlacklistWriter{err: errors.New("redis unavailable")}) + user, _, _ := createAuthRedisTestUserAndSession(t, ctx, store, svc.tokenMaker, "logout-all-blacklist-failure") + _, _, secondRefreshToken := createAuthRedisTestUserAndSessionForExistingUser(t, ctx, store, svc.tokenMaker, user, "logout-all-blacklist-failure-second") + accessTokenID, accessTokenExpiry := createAccessTokenMetadata(t, svc.tokenMaker, user.ID, user.OrgID) + + if err := svc.LogoutAll(ctx, user.ID, accessTokenID, accessTokenExpiry); !errors.Is(err, domain.ErrServiceUnavailable) { + t.Fatalf("expected ErrServiceUnavailable, got %v", err) + } + + assertSessionCount(t, ctx, store, user.ID, 2) + assertSessionActive(t, ctx, store, secondRefreshToken) +} + +func TestAuthServiceLogoutAllRevokesSessionsWhenBlacklistSucceeds(t *testing.T) { + store := testutil.RequireTestStore(t) + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + + svc := newRedisSensitiveAuthService(store, fakeBlacklistWriter{}) + user, _, firstRefreshToken := createAuthRedisTestUserAndSession(t, ctx, store, svc.tokenMaker, "logout-all-success") + _, _, secondRefreshToken := createAuthRedisTestUserAndSessionForExistingUser(t, ctx, store, svc.tokenMaker, user, "logout-all-success-second") + accessTokenID, accessTokenExpiry := createAccessTokenMetadata(t, svc.tokenMaker, user.ID, user.OrgID) + + if err := svc.LogoutAll(ctx, user.ID, accessTokenID, accessTokenExpiry); err != nil { + t.Fatalf("LogoutAll() error = %v", err) + } + + assertSessionRevoked(t, ctx, store, firstRefreshToken) + assertSessionRevoked(t, ctx, store, secondRefreshToken) +} + +func TestAuthServiceChangePasswordBlacklistFailureDoesNotMutateState(t *testing.T) { + store := testutil.RequireTestStore(t) + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + + svc := newRedisSensitiveAuthService(store, fakeBlacklistWriter{err: errors.New("redis unavailable")}) + user, currentPassword, refreshToken := createAuthRedisTestUserAndSession(t, ctx, store, svc.tokenMaker, "change-password-blacklist-failure") + accessTokenID, accessTokenExpiry := createAccessTokenMetadata(t, svc.tokenMaker, user.ID, user.OrgID) + + if err := svc.ChangePassword(ctx, user.OrgID, user.ID, currentPassword, "N3wPassword!2", accessTokenID, accessTokenExpiry); !errors.Is(err, domain.ErrServiceUnavailable) { + t.Fatalf("expected ErrServiceUnavailable, got %v", err) + } + + reloadedUser, err := store.GetUserByID(ctx, user.OrgID, user.ID) + if err != nil { + t.Fatalf("GetUserByID() error = %v", err) + } + if err := bcrypt.CompareHashAndPassword([]byte(reloadedUser.HashedPassword), []byte(currentPassword)); err != nil { + t.Fatalf("expected original password hash to remain valid, got %v", err) + } + assertSessionActive(t, ctx, store, refreshToken) +} + +func TestAuthServiceChangePasswordMutatesStateWhenBlacklistSucceeds(t *testing.T) { + store := testutil.RequireTestStore(t) + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + + svc := newRedisSensitiveAuthService(store, fakeBlacklistWriter{}) + user, currentPassword, refreshToken := createAuthRedisTestUserAndSession(t, ctx, store, svc.tokenMaker, "change-password-success") + accessTokenID, accessTokenExpiry := createAccessTokenMetadata(t, svc.tokenMaker, user.ID, user.OrgID) + newPassword := "N3wPassword!2" + + if err := svc.ChangePassword(ctx, user.OrgID, user.ID, currentPassword, newPassword, accessTokenID, accessTokenExpiry); err != nil { + t.Fatalf("ChangePassword() error = %v", err) + } + + reloadedUser, err := store.GetUserByID(ctx, user.OrgID, user.ID) + if err != nil { + t.Fatalf("GetUserByID() error = %v", err) + } + if err := bcrypt.CompareHashAndPassword([]byte(reloadedUser.HashedPassword), []byte(newPassword)); err != nil { + t.Fatalf("expected new password hash to be stored, got %v", err) + } + assertSessionRevoked(t, ctx, store, refreshToken) +} + +func newRedisSensitiveAuthService(store *db.Store, blacklistWriter accessTokenBlacklistWriter) *AuthService { + return &AuthService{ + store: store, + tokenMaker: token.NewMaker(authRedisTestJWTSecret, 15*time.Minute, 7*24*time.Hour), + cache: blacklistWriter, + auditSvc: NewAuditService(store, testutil.DiscardLogger()), + } +} + +func createAuthRedisTestUserAndSession(t *testing.T, ctx context.Context, store *db.Store, maker *token.Maker, prefix string) (*domain.User, string, string) { + t.Helper() + + org := testutil.CreateOrganization(t, store, prefix+"-org") + return createAuthRedisTestUserAndSessionInOrg(t, ctx, store, maker, org.ID, prefix) +} + +func createAuthRedisTestUserAndSessionForExistingUser(t *testing.T, ctx context.Context, store *db.Store, maker *token.Maker, user *domain.User, prefix string) (*domain.User, string, string) { + t.Helper() + return createSessionForExistingUser(t, ctx, store, maker, user) +} + +func createAuthRedisTestUserAndSessionInOrg(t *testing.T, ctx context.Context, store *db.Store, maker *token.Maker, orgID uuid.UUID, prefix string) (*domain.User, string, string) { + t.Helper() + + currentPassword := "Curr3ntPass!" + hashedPassword, err := bcrypt.GenerateFromPassword([]byte(currentPassword), bcrypt.DefaultCost) + if err != nil { + t.Fatalf("GenerateFromPassword() error = %v", err) + } + + email := prefix + "-" + uuid.NewString() + "@example.com" + user, err := store.CreateUser(ctx, orgID, email, string(hashedPassword), nil, false) + if err != nil { + t.Fatalf("CreateUser() error = %v", err) + } + + _, _, refreshToken := createSessionForExistingUser(t, ctx, store, maker, user) + return user, currentPassword, refreshToken +} + +func createSessionForExistingUser(t *testing.T, ctx context.Context, store *db.Store, maker *token.Maker, user *domain.User) (*domain.User, string, string) { + t.Helper() + + refreshToken, refreshClaims, err := maker.CreateRefreshToken(user.ID, user.OrgID) + if err != nil { + t.Fatalf("CreateRefreshToken() error = %v", err) + } + if _, err := store.CreateSession(ctx, user.ID, hashString(refreshToken), nil, nil, refreshClaims.ExpiresAt.Time); err != nil { + t.Fatalf("CreateSession() error = %v", err) + } + + return user, "", refreshToken +} + +func createAccessTokenMetadata(t *testing.T, maker *token.Maker, userID, orgID uuid.UUID) (uuid.UUID, time.Time) { + t.Helper() + + _, claims, err := maker.CreateAccessToken(userID, orgID) + if err != nil { + t.Fatalf("CreateAccessToken() error = %v", err) + } + return claims.TokenID, claims.ExpiresAt.Time +} + +func assertSessionActive(t *testing.T, ctx context.Context, store *db.Store, refreshToken string) { + t.Helper() + + sess, err := store.GetSessionByTokenHash(ctx, hashString(refreshToken)) + if err != nil { + t.Fatalf("GetSessionByTokenHash() error = %v", err) + } + if sess.RevokedAt != nil { + t.Fatalf("expected session %s to remain active", sess.ID) + } +} + +func assertSessionRevoked(t *testing.T, ctx context.Context, store *db.Store, refreshToken string) { + t.Helper() + + sess, err := store.GetSessionByTokenHash(ctx, hashString(refreshToken)) + if err != nil { + t.Fatalf("GetSessionByTokenHash() error = %v", err) + } + if sess.RevokedAt == nil { + t.Fatalf("expected session %s to be revoked", sess.ID) + } +} + +func assertSessionCount(t *testing.T, ctx context.Context, store *db.Store, userID uuid.UUID, want int) { + t.Helper() + + sessions, err := store.ListSessionsByUser(ctx, userID) + if err != nil { + t.Fatalf("ListSessionsByUser() error = %v", err) + } + if len(sessions) != want { + t.Fatalf("session count = %d, want %d", len(sessions), want) + } +}