445 lines
16 KiB
Go
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
|
|
}
|