package oauth import ( "context" "net/http" "net/http/httptest" "net/url" "strings" "sync" "testing" ) // Abuse cases and input limits. // // The property every test here shares: a refusal must not teach the caller // anything. Not whether a token existed, not whether an account is suspended // rather than deleted, not what the database is called, and never the value of // a credential that was presented. /* ── Nothing sensitive reaches a response ───────────────────────────────── */ // The broadest check in this file: drive every failure path with known secret // values, and assert none of them comes back. func TestNoSecretEverAppearsInAResponse(t *testing.T) { h := newHarness(t) clientID := h.register() const ( secretVerifier = "SENTINELverifier0123456789abcdefghijklmnop" secretCode = "SENTINELcodevalue" secretToken = "SENTINELtokenvalue" ) bodies := map[string]string{} // A failed exchange, with a sentinel code and verifier. bodies["bad code"] = h.exchange(clientID, secretCode, secretVerifier).Body.String() // A real code with the wrong verifier. real := h.authorizeOK(clientID, verifier43) bodies["bad verifier"] = h.exchange(clientID, real, secretVerifier).Body.String() // A refresh with a sentinel token. bodies["bad refresh"] = h.token(url.Values{ "grant_type": {"refresh_token"}, "refresh_token": {secretToken}, "client_id": {clientID}, }).Body.String() // Revocation of an unknown token. form := url.Values{"token": {secretToken}} req := httptest.NewRequest(http.MethodPost, "/oauth/revoke", strings.NewReader(form.Encode())) req.Header.Set("Content-Type", "application/x-www-form-urlencoded") rec := httptest.NewRecorder() h.server.RevokeHandler().ServeHTTP(rec, req) bodies["revoke unknown"] = rec.Body.String() for where, body := range bodies { for what, secret := range map[string]string{ "code_verifier": secretVerifier, "authorization code": secretCode, "token": secretToken, } { if strings.Contains(body, secret) { t.Errorf("%s: the response echoes the presented %s:\n%s", where, what, body) } } // Nor may it leak the shape of the system. for _, tell := range []string{"SQLSTATE", "pq:", "pgx", "oauth_tokens", "oauth_grants", "password", "Krow-force", "relation", "column"} { if strings.Contains(body, tell) { t.Errorf("%s: the response leaks an internal detail (%q):\n%s", where, tell, body) } } } } /* ── Replay ─────────────────────────────────────────────────────────────── */ // Ten attempts to spend one code. Exactly one may succeed. func TestAuthorizationCodeReplayUnderLoad(t *testing.T) { h := newHarness(t) clientID := h.register() code := h.authorizeOK(clientID, verifier43) succeeded := 0 for i := 0; i < 10; i++ { if h.exchange(clientID, code, verifier43).Code == http.StatusOK { succeeded++ } } if succeeded != 1 { t.Errorf("%d of 10 exchanges of the same code succeeded, want exactly 1", succeeded) } } // The same, concurrently. A single-use credential redeemed by two racing // callers must be spent exactly once — this is the property that // UPDATE … RETURNING buys, and the one a SELECT-then-UPDATE would lose. // Run with -race. func TestConcurrentCodeRedemptionSpendsItOnce(t *testing.T) { h := newHarness(t) clientID := h.register() code := h.authorizeOK(clientID, verifier43) const racers = 8 var wg sync.WaitGroup var mu sync.Mutex succeeded := 0 for i := 0; i < racers; i++ { wg.Add(1) go func() { defer wg.Done() if h.exchange(clientID, code, verifier43).Code == http.StatusOK { mu.Lock() succeeded++ mu.Unlock() } }() } wg.Wait() if succeeded != 1 { t.Errorf("%d of %d concurrent redemptions succeeded, want exactly 1", succeeded, racers) } } // Concurrent refresh of the same token: one rotation, not several. Two // successes would mean two live families from one credential. // Run with -race. func TestConcurrentRefreshRotatesOnce(t *testing.T) { h := newHarness(t) clientID := h.register() code := h.authorizeOK(clientID, verifier43) tokens := decodeTokens(t, h.exchange(clientID, code, verifier43)) const racers = 8 var wg sync.WaitGroup var mu sync.Mutex succeeded := 0 for i := 0; i < racers; i++ { wg.Add(1) go func() { defer wg.Done() rec := h.token(url.Values{ "grant_type": {"refresh_token"}, "refresh_token": {tokens.RefreshToken}, "client_id": {clientID}, }) if rec.Code == http.StatusOK { mu.Lock() succeeded++ mu.Unlock() } }() } wg.Wait() if succeeded != 1 { t.Errorf("%d of %d concurrent refreshes succeeded, want exactly 1 — "+ "a refresh token was double-spent", succeeded, racers) } } /* ── Open redirect ──────────────────────────────────────────────────────── */ // Every shape of redirect tampering, each of which has been a real CVE // somewhere. None may be honoured, and none may be answered WITH a redirect — // redirecting an error to an unvalidated URI is the open redirect itself. func TestOpenRedirectAttempts(t *testing.T) { h := newHarness(t) clientID := h.register() for name, redirect := range map[string]string{ "different host": "https://attacker.example/cb", "prefix extension": testRedirect + ".attacker.example", "path traversal": testRedirect + "/../../evil", "userinfo trick": "https://claude.example.test@attacker.example/cb", "added query": testRedirect + "?next=https://attacker.example", "protocol swap": strings.Replace(testRedirect, "https", "http", 1), "case variation": strings.ToUpper(testRedirect), "trailing slash": testRedirect + "/", "double slash": "//attacker.example/cb", "encoded traversal": testRedirect + "/%2e%2e/evil", "null byte": testRedirect + "\x00.attacker.example", "newline injection": testRedirect + "\nLocation: https://attacker.example", "javascript": "javascript:alert(1)", "completely missing": "", } { t.Run(name, func(t *testing.T) { rec := h.authorize(map[string]string{ "client_id": clientID, "redirect_uri": redirect, "response_type": "code", "state": "xyz", "code_challenge": ChallengeFor(verifier43), "code_challenge_method": "S256", "resource": testResource, }) if rec.Code == http.StatusFound { location := rec.Header().Get("Location") t.Fatalf("answered with a redirect to %q — an unregistered target "+ "must produce a direct error, never a redirect", location) } if rec.Code != http.StatusBadRequest { t.Errorf("status = %d, want 400", rec.Code) } if strings.Contains(rec.Body.String(), "attacker.example") { t.Error("the error echoes the attacker's host back") } }) } } /* ── Malformed and oversized input ──────────────────────────────────────── */ func TestRegistrationRejectsMalformedBodies(t *testing.T) { h := newHarness(t) for name, body := range map[string]string{ "not json": `not json at all`, "truncated": `{"client_name":`, "null": `null`, "array": `[]`, "deeply nested": `{"client_name":` + strings.Repeat(`[`, 2000) + strings.Repeat(`]`, 2000) + `}`, "empty": ``, "wrong types": `{"client_name":123,"redirect_uris":"not-an-array"}`, "huge name": `{"client_name":"` + strings.Repeat("A", 100_000) + `","redirect_uris":["https://a.test/cb"]}`, "too many uris": `{"client_name":"x","redirect_uris":[` + strings.TrimSuffix(strings.Repeat(`"https://a.test/cb",`, 50), ",") + `]}`, "oversized body": `{"client_name":"` + strings.Repeat("A", 20<<10) + `"}`, } { t.Run(name, func(t *testing.T) { req := httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body)) rec := httptest.NewRecorder() // A panic here fails the test by crashing it, which is the assertion. h.server.RegisterHandler().ServeHTTP(rec, req) if rec.Code == http.StatusCreated { // Only the "huge name" case may legitimately succeed, truncated. if name != "huge name" { t.Errorf("status = %d; a malformed registration was accepted", rec.Code) } return } if rec.Code < 400 || rec.Code >= 500 { t.Errorf("status = %d, want a 4xx", rec.Code) } }) } } // A client name is stored, shown on a consent screen, and attacker-controlled. // It must be bounded, or registration becomes free storage. func TestClientNameIsBounded(t *testing.T) { h := newHarness(t) body := `{"client_name":"` + strings.Repeat("A", 5000) + `","redirect_uris":["https://a.test/cb"]}` rec := httptest.NewRecorder() h.server.RegisterHandler().ServeHTTP(rec, httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body))) if rec.Code != http.StatusCreated { t.Fatalf("status = %d", rec.Code) } var stored string if err := h.h.Pool.QueryRow(context.Background(), `SELECT client_name FROM oauth_clients ORDER BY created_date DESC LIMIT 1`).Scan(&stored); err != nil { t.Fatalf("read: %v", err) } if len(stored) > 200 { t.Errorf("stored client_name is %d characters; the column's CHECK allows 200", len(stored)) } } func TestTokenEndpointRejectsMalformedRequests(t *testing.T) { h := newHarness(t) for name, tc := range map[string]struct { body string contentType string }{ "no content type": {"grant_type=authorization_code", ""}, "json body": {`{"grant_type":"authorization_code"}`, "application/json"}, "empty": {"", "application/x-www-form-urlencoded"}, "garbage": {"%%%%", "application/x-www-form-urlencoded"}, "huge": {"grant_type=authorization_code&code=" + strings.Repeat("A", 200_000), "application/x-www-form-urlencoded"}, "repeated params": {"grant_type=authorization_code&grant_type=password", "application/x-www-form-urlencoded"}, "null grant": {"grant_type=", "application/x-www-form-urlencoded"}, "unknown grant": {"grant_type=magic", "application/x-www-form-urlencoded"}, "injection in code": {"grant_type=authorization_code&code=' OR 1=1 --&client_id=x&redirect_uri=y&code_verifier=" + strings.Repeat("a", 43), "application/x-www-form-urlencoded"}, } { t.Run(name, func(t *testing.T) { req := httptest.NewRequest(http.MethodPost, "/oauth/token", strings.NewReader(tc.body)) if tc.contentType != "" { req.Header.Set("Content-Type", tc.contentType) } rec := httptest.NewRecorder() h.server.TokenHandler().ServeHTTP(rec, req) if rec.Code == http.StatusOK { t.Errorf("a malformed token request succeeded: %s", rec.Body.String()) } if rec.Code >= 500 { t.Errorf("status = %d; a malformed request must not be an internal error: %s", rec.Code, rec.Body.String()) } }) } } /* ── Scope escalation ───────────────────────────────────────────────────── */ // krow.write must be unreachable from every angle: registration, authorization, // and the consent POST. func TestWriteScopeIsUnreachable(t *testing.T) { h := newHarness(t) t.Run("at registration", func(t *testing.T) { body := `{"client_name":"x","redirect_uris":["` + testRedirect + `"],"scope":"` + ScopeWrite + `"}` rec := httptest.NewRecorder() h.server.RegisterHandler().ServeHTTP(rec, httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body))) if rec.Code == http.StatusCreated { t.Error("a client registered for krow.write") } }) t.Run("at authorization", func(t *testing.T) { clientID := h.register() rec := h.authorize(map[string]string{ "client_id": clientID, "redirect_uri": testRedirect, "response_type": "code", "state": "xyz", "code_challenge": ChallengeFor(verifier43), "code_challenge_method": "S256", "resource": testResource, "scope": ScopeRead + " " + ScopeWrite, }) loc, _ := url.Parse(rec.Header().Get("Location")) if loc.Query().Get("code") != "" { t.Error("an authorization requesting krow.write produced a code") } }) t.Run("no issued token carries it", func(t *testing.T) { clientID := h.register() code := h.authorizeOK(clientID, verifier43) tokens := decodeTokens(t, h.exchange(clientID, code, verifier43)) if strings.Contains(tokens.Scope, ScopeWrite) { t.Errorf("an issued token carries %q", tokens.Scope) } }) } /* ── Cache and transport headers ────────────────────────────────────────── */ // A credential-bearing response must never be cached, and no endpoint may put // a token in a URL. func TestSensitiveResponsesAreNotCacheable(t *testing.T) { h := newHarness(t) clientID := h.register() code := h.authorizeOK(clientID, verifier43) rec := h.exchange(clientID, code, verifier43) for header, want := range map[string]string{ "Cache-Control": "no-store", "Pragma": "no-cache", } { if got := rec.Header().Get(header); !strings.Contains(got, want) { t.Errorf("%s = %q, want it to contain %q", header, got, want) } } // The authorization redirect carries a code in its query — that is the // protocol — but it must never carry a token. approved := h.authorize(authorizeParamsFor(clientID, verifier43)) if strings.Contains(approved.Header().Get("Location"), "access_token") { t.Error("an access token appeared in a redirect URL") } }