mcp connection
This commit is contained in:
444
go-api/internal/oauth/store.go
Normal file
444
go-api/internal/oauth/store.go
Normal file
@@ -0,0 +1,444 @@
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user