Files
krow_backend/go-api/internal/oauth/store.go
Aravind f2aa3b3ad8
Some checks failed
CI / fixture (push) Has been cancelled
CI / test (push) Has been cancelled
mcp connection
2026-09-22 10:58:02 +05:30

445 lines
16 KiB
Go

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
}