package oauth import ( "context" "errors" "fmt" "time" "github.com/krow/krow-backend/go-api/internal/auth" "github.com/krow/krow-backend/go-api/internal/repo" ) // The persistence layer for clients, authorization codes and tokens. // // Two rules hold throughout this file and are worth stating once: // // 1. NO RAW CREDENTIAL IS EVER WRITTEN. Every code and token is hashed with // auth.HashToken before it reaches a statement. The CHECK constraints in // migrations 000013 and 000014 refuse anything that is not 64 hex // characters, so this is enforced twice — in Go where the mistake would be // made, and in the schema where it would land. // // 2. SINGLE-USE IS ENFORCED BY THE UPDATE, NOT BY A READ. Redeeming a code or // a refresh token is one statement that marks the row consumed and returns // it in the same breath. A SELECT followed by an UPDATE has a window // between them where two concurrent requests both see an unspent row and // both proceed, which is precisely the replay the single-use rule exists to // prevent. var ( // ErrNotFound covers a client, code or token that does not exist. ErrNotFound = errors.New("oauth: not found") // ErrGrantUnusable covers a code that is expired, already consumed, or // simply absent. ONE error for all three: distinguishing them tells a // caller whether a code they hold was ever real, which is an oracle. ErrGrantUnusable = errors.New("oauth: authorization code is not usable") // ErrTokenUnusable covers a token that is unknown, expired, revoked or // consumed. One error, same reasoning. ErrTokenUnusable = errors.New("oauth: token is not usable") // ErrRefreshReuse is raised when a CONSUMED refresh token is presented // again. It is distinct from ErrTokenUnusable internally because it // triggers family revocation — but the caller must still answer the client // with an indistinguishable error. ErrRefreshReuse = errors.New("oauth: refresh token reuse detected") ) // Store is the database-backed persistence for this package. type Store struct { db repo.Querier now func() time.Time } // NewStore builds a store over the existing pool. func NewStore(db repo.Querier) *Store { return &Store{db: db, now: time.Now} } // WithClock replaces the clock, so expiry can be tested without sleeping. func (s *Store) WithClock(now func() time.Time) *Store { s.now = now return s } /* ── Clients ────────────────────────────────────────────────────────────── */ // Client is a registered OAuth client. type Client struct { ClientID string ClientName string RedirectURIs []string GrantTypes []string Scopes []string DisabledAt *time.Time } // AllowsRedirect reports whether a redirect_uri is registered to this client. // // EXACT string equality. Not a prefix match, not a normalised comparison, not // "same host and port". Every relaxation of this check is an open redirect: a // prefix match lets `https://good.example/cb.attacker.com` through, and // normalising lets encoding tricks through. RFC 6749 section 3.1.2.3 says // exact, and exact is what this is. func (c Client) AllowsRedirect(uri string) bool { for _, registered := range c.RedirectURIs { if registered == uri { return true } } return false } // AllowsScopes reports whether every requested scope is registered. func (c Client) AllowsScopes(requested []string) bool { for _, want := range requested { found := false for _, have := range c.Scopes { if have == want { found = true break } } if !found { return false } } return true } // CreateClient registers a new public client. func (s *Store) CreateClient(ctx context.Context, c Client) error { _, err := s.db.Exec(ctx, `INSERT INTO oauth_clients (client_id, client_name, redirect_uris, grant_types, scopes) VALUES ($1, $2, $3, $4, $5)`, c.ClientID, c.ClientName, c.RedirectURIs, c.GrantTypes, c.Scopes) if err != nil { return fmt.Errorf("oauth: create client: %w", err) } return nil } // FindClient resolves a client_id. A disabled client is reported as not found: // whether it once existed is not the caller's business. func (s *Store) FindClient(ctx context.Context, clientID string) (Client, error) { var c Client err := s.db.QueryRow(ctx, `SELECT client_id, client_name, redirect_uris, grant_types, scopes, disabled_at FROM oauth_clients WHERE client_id = $1 AND disabled_at IS NULL`, clientID).Scan(&c.ClientID, &c.ClientName, &c.RedirectURIs, &c.GrantTypes, &c.Scopes, &c.DisabledAt) if err != nil { return Client{}, ErrNotFound } return c, nil } /* ── Authorization codes ────────────────────────────────────────────────── */ // Grant is an issued authorization code, as stored. type Grant struct { ID string ClientID string UserID string OrgID string RedirectURI string Scopes []string Resource string CodeChallenge string CodeChallengeMethod string ExpiresAt time.Time } // GrantTTL is how long an authorization code stays redeemable. // // Sixty seconds. RFC 6749 recommends a maximum of ten minutes and "a maximum // of 60 seconds is RECOMMENDED" for the code's lifetime in OAuth 2.1 guidance. // The code is in transit through a browser redirect and is exchanged // immediately by a client that is already waiting for it; a longer window buys // nothing and widens the replay opportunity. const GrantTTL = 60 * time.Second // CreateGrant stores an authorization code, returning the RAW code exactly // once. // // The raw value is returned and never persisted. The caller puts it in a // redirect and forgets it. func (s *Store) CreateGrant(ctx context.Context, g Grant) (rawCode string, err error) { rawCode, err = auth.GenerateToken() if err != nil { return "", fmt.Errorf("oauth: generate code: %w", err) } _, err = s.db.Exec(ctx, `INSERT INTO oauth_grants (code_hash, client_id, user_id, org_id, redirect_uri, scopes, resource, code_challenge, code_challenge_method, expires_at) VALUES ($1, $2, $3::uuid, $4::uuid, $5, $6, $7, $8, $9, $10)`, auth.HashToken(rawCode), g.ClientID, g.UserID, g.OrgID, g.RedirectURI, g.Scopes, g.Resource, g.CodeChallenge, g.CodeChallengeMethod, s.now().Add(GrantTTL)) if err != nil { return "", fmt.Errorf("oauth: create grant: %w", err) } return rawCode, nil } // RedeemGrant consumes an authorization code and returns what it was bound to. // // ONE STATEMENT. The UPDATE marks the row consumed and RETURNS it, so the read // and the write cannot be interleaved by a concurrent request. The predicate // carries the whole single-use rule: `consumed_at IS NULL` means a spent code // matches nothing, and `expires_at > now()` means an old one does too. A second // redemption of the same code updates zero rows and therefore fails, which is // what replay protection looks like when the database enforces it. func (s *Store) RedeemGrant(ctx context.Context, rawCode string) (Grant, error) { var g Grant err := s.db.QueryRow(ctx, `UPDATE oauth_grants SET consumed_at = now() WHERE code_hash = $1 AND consumed_at IS NULL AND expires_at > $2 RETURNING id::text, client_id, user_id::text, org_id::text, redirect_uri, scopes, resource, code_challenge, code_challenge_method, expires_at`, auth.HashToken(rawCode), s.now()). Scan(&g.ID, &g.ClientID, &g.UserID, &g.OrgID, &g.RedirectURI, &g.Scopes, &g.Resource, &g.CodeChallenge, &g.CodeChallengeMethod, &g.ExpiresAt) if err != nil { // No row: unknown, expired or already spent. Indistinguishable on // purpose — see ErrGrantUnusable. return Grant{}, ErrGrantUnusable } return g, nil } /* ── Tokens ─────────────────────────────────────────────────────────────── */ // Token is an issued access or refresh token, as stored. type Token struct { ID string Type string FamilyID string ClientID string UserID string OrgID string Scopes []string Audience string ExpiresAt time.Time } // Token lifetimes. // // Fifteen minutes for an access token is the number the MCP plan committed to, // and the reasoning is that an access token travels on every single request: it // is the most exposed credential in the system and the one with the least need // to be long-lived, because a refresh token exists precisely so the client can // get another without troubling the user. // // Thirty days for a refresh token matches the session's own "remember me" // ceiling, so a connected client and a remembered browser lapse on the same // schedule rather than on two different ones nobody can remember. const ( AccessTokenTTL = 15 * time.Minute RefreshTokenTTL = 30 * 24 * time.Hour ) // TokenPair is what a successful token request produces. // // The raw values are here and nowhere else: they are returned to the client in // the token response and are never stored, logged or re-derivable. type TokenPair struct { AccessToken string RefreshToken string ExpiresIn int Scopes []string FamilyID string } // IssuePair mints an access and refresh token in one family. // // familyID empty starts a new lineage; a supplied one continues an existing // lineage through a rotation, which is what lets reuse detection revoke every // descendant of a stolen token. func (s *Store) IssuePair(ctx context.Context, t Token, familyID string) (TokenPair, error) { if familyID == "" { generated, err := newUUID() if err != nil { return TokenPair{}, err } familyID = generated } access, err := auth.GenerateToken() if err != nil { return TokenPair{}, fmt.Errorf("oauth: generate access token: %w", err) } refresh, err := auth.GenerateToken() if err != nil { return TokenPair{}, fmt.Errorf("oauth: generate refresh token: %w", err) } now := s.now() for _, row := range []struct { raw string kind string expires time.Time }{ {access, "access", now.Add(AccessTokenTTL)}, {refresh, "refresh", now.Add(RefreshTokenTTL)}, } { if _, err := s.db.Exec(ctx, `INSERT INTO oauth_tokens (token_hash, token_type, family_id, client_id, user_id, org_id, scopes, audience, expires_at) VALUES ($1, $2, $3::uuid, $4, $5::uuid, $6::uuid, $7, $8, $9)`, auth.HashToken(row.raw), row.kind, familyID, t.ClientID, t.UserID, t.OrgID, t.Scopes, t.Audience, row.expires); err != nil { return TokenPair{}, fmt.Errorf("oauth: store %s token: %w", row.kind, err) } } return TokenPair{ AccessToken: access, RefreshToken: refresh, ExpiresIn: int(AccessTokenTTL.Seconds()), Scopes: t.Scopes, FamilyID: familyID, }, nil } // FindAccessToken resolves a raw access token for validation. // // Read-only: validation happens on every MCP request and must not write. The // predicate does the whole job — unknown, expired, revoked and wrong-type all // return no row and therefore the same error. func (s *Store) FindAccessToken(ctx context.Context, raw string) (Token, error) { var t Token err := s.db.QueryRow(ctx, `SELECT id::text, token_type, family_id::text, client_id, user_id::text, org_id::text, scopes, audience, expires_at FROM oauth_tokens WHERE token_hash = $1 AND token_type = 'access' AND revoked_at IS NULL AND expires_at > $2`, auth.HashToken(raw), s.now()). Scan(&t.ID, &t.Type, &t.FamilyID, &t.ClientID, &t.UserID, &t.OrgID, &t.Scopes, &t.Audience, &t.ExpiresAt) if err != nil { return Token{}, ErrTokenUnusable } return t, nil } // RedeemRefreshToken consumes a refresh token, or detects its reuse. // // The two-step here is deliberate and is the heart of reuse detection: // // 1. Try to consume an unspent, unexpired, unrevoked refresh token. One // statement, same single-use reasoning as RedeemGrant. // 2. If that matched nothing, look again WITHOUT the `consumed_at IS NULL` // predicate. A row that exists but was already consumed is not an ordinary // failure — it means someone presented a token that had already been // rotated away, and there is no way to tell the legitimate client retrying // from an attacker replaying a stolen token. // // OAuth 2.1's answer to that ambiguity is to assume the worse case and revoke // the whole family. The attacker loses access; the legitimate client is pushed // through a fresh authorization it can complete. Doing nothing would leave a // thief with a working credential. func (s *Store) RedeemRefreshToken(ctx context.Context, raw string) (Token, error) { hash := auth.HashToken(raw) var t Token err := s.db.QueryRow(ctx, `UPDATE oauth_tokens SET consumed_at = now(), last_used_at = now() WHERE token_hash = $1 AND token_type = 'refresh' AND consumed_at IS NULL AND revoked_at IS NULL AND expires_at > $2 RETURNING id::text, token_type, family_id::text, client_id, user_id::text, org_id::text, scopes, audience, expires_at`, hash, s.now()). Scan(&t.ID, &t.Type, &t.FamilyID, &t.ClientID, &t.UserID, &t.OrgID, &t.Scopes, &t.Audience, &t.ExpiresAt) if err == nil { return t, nil } // Step 2: was this a token that HAD been valid and is now spent? var familyID string if probeErr := s.db.QueryRow(ctx, `SELECT family_id::text FROM oauth_tokens WHERE token_hash = $1 AND token_type = 'refresh' AND consumed_at IS NOT NULL`, hash).Scan(&familyID); probeErr == nil { // Reuse. Revoke the lineage and report it, so the caller can log it at // a level that gets noticed — while still answering the client with an // indistinguishable error. _ = s.RevokeFamily(ctx, familyID, "refresh_token_reuse") return Token{}, ErrRefreshReuse } return Token{}, ErrTokenUnusable } /* ── Revocation ─────────────────────────────────────────────────────────── */ // RevokeFamily revokes every token in a rotation lineage. // // Idempotent, and it does not care whether the rows were already revoked: the // predicate narrows to unrevoked rows so a second call is a no-op rather than // an error, which matters because this is called from an error path. func (s *Store) RevokeFamily(ctx context.Context, familyID, reason string) error { _, err := s.db.Exec(ctx, `UPDATE oauth_tokens SET revoked_at = now(), revoked_reason = $2 WHERE family_id = $1::uuid AND revoked_at IS NULL`, familyID, reason) if err != nil { return fmt.Errorf("oauth: revoke family: %w", err) } return nil } // RevokeToken revokes one token by its raw value, and its family with it. // // Revoking the family rather than the single row is what makes "disconnect" // mean what a person expects. Revoking one access token would leave the // refresh token alive to mint another within seconds, so the button that says // "disconnect Claude" would not disconnect Claude. func (s *Store) RevokeToken(ctx context.Context, raw, reason string) error { var familyID string if err := s.db.QueryRow(ctx, `SELECT family_id::text FROM oauth_tokens WHERE token_hash = $1`, auth.HashToken(raw)).Scan(&familyID); err != nil { // RFC 7009: revoking an unknown token is a success. Saying otherwise // turns the revocation endpoint into a way to test whether a token // exists. return nil } return s.RevokeFamily(ctx, familyID, reason) } // RevokeAllForUser revokes every token a user holds. // // Called when an account is suspended or a person disconnects every app. Token // validation already re-reads the user and refuses a suspended one, so this is // belt to that braces: it stops the tokens existing rather than relying on // every future validation to notice. func (s *Store) RevokeAllForUser(ctx context.Context, userID, reason string) error { _, err := s.db.Exec(ctx, `UPDATE oauth_tokens SET revoked_at = now(), revoked_reason = $2 WHERE user_id = $1::uuid AND revoked_at IS NULL`, userID, reason) if err != nil { return fmt.Errorf("oauth: revoke user tokens: %w", err) } return nil }