Files
ResolveSpec/pkg/security/lookup/direct/oauth.go
T
Hein c9fa8c60f2 refactor(security): move all database access into pkg/security/lookup
pkg/security no longer contains SQL. Every provider calls a store interface
from lookup, implemented by a procedure backend (Postgres stored procedures,
the default there) and a direct backend (dialect-driven SQL for postgres,
sqlite, mysql and mssql with configurable table and column names).

- add sectypes, lookup, lookup/{dialect,procedure,direct,backends,ddl,conformance}
- split totp and providers sub packages out of the core package
- replace SQLNames/TableNames/QueryMode with lookup.Config (see breaking_changes.md)
- direct backend now covers column/row security and API-key login
- move txsettings SQL to lookup.ApplyTxSettings; remove password.go
- move schema scripts under lookup/, add reference DDL per dialect
- add a shared conformance suite; run it on sqlite, and on Postgres in a
  podman/docker container (RESOLVESPEC_TEST_CONTAINERS=1)
- fix procedure schema bugs found on real Postgres: duplicate p_data
  parameter, JSON null arrays, expires_at timezone casts, passkey list
  GROUP BY, missing resolvespec_passkey_login; accept zone-less timestamps
2026-10-01 13:19:44 +02:00

381 lines
14 KiB
Go

package direct
import (
"context"
"database/sql"
"errors"
"fmt"
"strings"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
// nullIfEmpty keeps optional TEXT columns (e.g. client_secret_hash of a public client) NULL
// rather than "".
func nullIfEmpty(s string) any {
if s == "" {
return nil
}
return s
}
// OAuthClients implements lookup.OAuthClientStore. Array columns (redirect_uris, grant_types,
// allowed_scopes, scopes) are JSON through the dialect.
type OAuthClients struct{ *Base }
var _ lookup.OAuthClientStore = (*OAuthClients)(nil)
// NewOAuthClients creates the direct OAuthClientStore.
func NewOAuthClients(b *Base) *OAuthClients { return &OAuthClients{Base: b} }
// RegisterClient implements lookup.OAuthClientStore.
func (o *OAuthClients) RegisterClient(ctx context.Context, client *sectypes.OAuthServerClient) (*sectypes.OAuthServerClient, error) {
grantTypes := client.GrantTypes
if len(grantTypes) == 0 {
grantTypes = []string{"authorization_code"}
}
allowedScopes := client.AllowedScopes
if len(allowedScopes) == 0 {
allowedScopes = []string{"openid", "profile", "email"}
}
authMethod := client.TokenEndpointAuthMethod
if authMethod == "" {
authMethod = "none"
}
redirects, err := o.d.EncodeJSON(client.RedirectURIs)
if err != nil {
return nil, fmt.Errorf("failed to marshal redirect_uris: %w", err)
}
if redirects == nil { // the column is NOT NULL
redirects = "[]"
}
grants, err := o.d.EncodeJSON(grantTypes)
if err != nil {
return nil, fmt.Errorf("failed to marshal grant_types: %w", err)
}
scopes, err := o.d.EncodeJSON(allowedScopes)
if err != nil {
return nil, fmt.Errorf("failed to marshal allowed_scopes: %w", err)
}
err = o.do(func(q Querier) error {
return o.Insert(lookup.EntityOAuthClients).Set(
Set(lookup.OAuthClientsClientID, client.ClientID),
Set(lookup.OAuthClientsRedirectURIs, redirects),
Set(lookup.OAuthClientsClientName, client.ClientName),
Set(lookup.OAuthClientsGrantTypes, grants),
Set(lookup.OAuthClientsAllowedScopes, scopes),
Set(lookup.OAuthClientsClientSecretHash, nullIfEmpty(client.ClientSecretHash)),
Set(lookup.OAuthClientsTokenEndpointAuthMethod, authMethod),
Set(lookup.OAuthClientsIsActive, true),
Set(lookup.OAuthClientsCreatedAt, o.Now()),
).Exec(ctx, q)
})
if err != nil {
return nil, fmt.Errorf("failed to register client: %w", err)
}
return &sectypes.OAuthServerClient{
ClientID: client.ClientID,
RedirectURIs: client.RedirectURIs,
ClientName: client.ClientName,
GrantTypes: grantTypes,
AllowedScopes: allowedScopes,
ClientSecretHash: client.ClientSecretHash,
TokenEndpointAuthMethod: authMethod,
}, nil
}
// GetClient implements lookup.OAuthClientStore.
func (o *OAuthClients) GetClient(ctx context.Context, clientID string) (*sectypes.OAuthServerClient, error) {
var redirects, grants, scopes any
var name, secret, method sql.NullString
err := o.do(func(q Querier) error {
return o.From(lookup.EntityOAuthClients).
Cols(lookup.OAuthClientsRedirectURIs, lookup.OAuthClientsClientName, lookup.OAuthClientsGrantTypes,
lookup.OAuthClientsAllowedScopes, lookup.OAuthClientsClientSecretHash, lookup.OAuthClientsTokenEndpointAuthMethod).
Where(Eq(lookup.OAuthClientsClientID, clientID), Eq(lookup.OAuthClientsIsActive, true)).
QueryRow(ctx, q, &redirects, &name, &grants, &scopes, &secret, &method)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, fmt.Errorf("client not found")
}
return nil, fmt.Errorf("failed to get client: %w", err)
}
res := &sectypes.OAuthServerClient{
ClientID: clientID,
ClientName: name.String,
ClientSecretHash: secret.String,
TokenEndpointAuthMethod: method.String,
}
_ = o.d.DecodeJSON(redirects, &res.RedirectURIs)
_ = o.d.DecodeJSON(grants, &res.GrantTypes)
_ = o.d.DecodeJSON(scopes, &res.AllowedScopes)
return res, nil
}
// SaveCode implements lookup.OAuthClientStore.
func (o *OAuthClients) SaveCode(ctx context.Context, code *sectypes.OAuthCode) error {
scopes, err := o.d.EncodeJSON(code.Scopes)
if err != nil {
return fmt.Errorf("failed to marshal scopes: %w", err)
}
method := code.CodeChallengeMethod
if method == "" {
method = "S256"
}
return o.do(func(q Querier) error {
return o.Insert(lookup.EntityOAuthCodes).Set(
Set(lookup.OAuthCodesCode, code.Code),
Set(lookup.OAuthCodesClientID, code.ClientID),
Set(lookup.OAuthCodesRedirectURI, code.RedirectURI),
Set(lookup.OAuthCodesClientState, code.ClientState),
Set(lookup.OAuthCodesCodeChallenge, code.CodeChallenge),
Set(lookup.OAuthCodesCodeChallengeMethod, method),
Set(lookup.OAuthCodesSessionToken, code.SessionToken),
Set(lookup.OAuthCodesRefreshToken, code.RefreshToken),
Set(lookup.OAuthCodesScopes, scopes),
Set(lookup.OAuthCodesExpiresAt, code.ExpiresAt),
Set(lookup.OAuthCodesCreatedAt, o.Now()),
).Exec(ctx, q)
})
}
// ExchangeCode implements lookup.OAuthClientStore: the code is consumed in a transaction and
// only the caller whose delete removes the row gets it, so a code cannot be redeemed twice.
func (o *OAuthClients) ExchangeCode(ctx context.Context, code string) (*sectypes.OAuthCode, error) {
var res sectypes.OAuthCode
var state, refresh sql.NullString
var scopes any
err := o.tx(ctx, func(q Querier) error {
err := o.From(lookup.EntityOAuthCodes).
Cols(lookup.OAuthCodesClientID, lookup.OAuthCodesRedirectURI, lookup.OAuthCodesClientState,
lookup.OAuthCodesCodeChallenge, lookup.OAuthCodesCodeChallengeMethod, lookup.OAuthCodesSessionToken,
lookup.OAuthCodesRefreshToken, lookup.OAuthCodesScopes).
Where(Eq(lookup.OAuthCodesCode, code), Gt(lookup.OAuthCodesExpiresAt, o.Now())).
QueryRow(ctx, q, &res.ClientID, &res.RedirectURI, &state, &res.CodeChallenge, &res.CodeChallengeMethod,
&res.SessionToken, &refresh, &scopes)
if err != nil {
return err
}
n, err := o.Delete(lookup.EntityOAuthCodes).Where(Eq(lookup.OAuthCodesCode, code)).Exec(ctx, q)
if err != nil {
return err
}
if n == 0 {
return sql.ErrNoRows
}
return nil
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, fmt.Errorf("invalid or expired code")
}
return nil, fmt.Errorf("failed to exchange code: %w", err)
}
res.Code = code
res.ClientState = state.String
res.RefreshToken = refresh.String
_ = o.d.DecodeJSON(scopes, &res.Scopes)
return &res, nil
}
// Introspect implements lookup.OAuthClientStore (RFC 7662). An unknown or expired token is
// {active:false}, not an error.
func (o *OAuthClients) Introspect(ctx context.Context, token string) (*sectypes.OAuthTokenInfo, error) {
var info sectypes.OAuthTokenInfo
var userID int
var username, email, roles sql.NullString
var level sql.NullInt64
var exp, iat time.Time
err := o.do(func(q Querier) error {
return o.From(lookup.EntityUserSessions).
Cols(lookup.UsersID, lookup.UsersUsername, lookup.UsersEmail, lookup.UsersUserLevel, lookup.UsersRoles,
lookup.SessionsExpiresAt, lookup.SessionsCreatedAt).
Join(lookup.EntityUsers, EqCol(lookup.UsersID, lookup.SessionsUserID)).
Where(Eq(lookup.SessionsToken, token), Gt(lookup.SessionsExpiresAt, o.Now()), Eq(lookup.UsersIsActive, true)).
QueryRow(ctx, q, &userID, &username, &email, &level, &roles, o.timeDest(&exp), o.timeDest(&iat))
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return &sectypes.OAuthTokenInfo{Active: false}, nil
}
return nil, fmt.Errorf("failed to introspect token: %w", err)
}
info.Active = true
info.Sub = fmt.Sprintf("%d", userID)
info.Username = username.String
info.Email = email.String
info.UserLevel = int(level.Int64)
info.Roles = ParseRoles(roles.String)
if !exp.IsZero() {
info.Exp = exp.Unix()
}
if !iat.IsZero() {
info.Iat = iat.Unix()
}
return &info, nil
}
// Revoke implements lookup.OAuthClientStore (RFC 7009): the session is deleted; an unknown
// token is not an error.
func (o *OAuthClients) Revoke(ctx context.Context, token string) error {
return o.do(func(q Querier) error {
_, err := o.Delete(lookup.EntityUserSessions).Where(Eq(lookup.SessionsToken, token)).Exec(ctx, q)
return err
})
}
// OAuthUsers implements lookup.OAuthUserStore.
type OAuthUsers struct{ *Base }
var _ lookup.OAuthUserStore = (*OAuthUsers)(nil)
// NewOAuthUsers creates the direct OAuthUserStore.
func NewOAuthUsers(b *Base) *OAuthUsers { return &OAuthUsers{Base: b} }
// GetOrCreateUser implements lookup.OAuthUserStore: select by email, then update or insert,
// in one transaction (no upsert).
func (o *OAuthUsers) GetOrCreateUser(ctx context.Context, user *sectypes.UserContext, provider string) (int, error) {
roles := strings.Join(user.Roles, ",")
var userID int
err := o.tx(ctx, func(q Querier) error {
now := o.Now()
var remoteID, authProvider sql.NullString
err := o.From(lookup.EntityUsers).Cols(lookup.UsersID, lookup.UsersRemoteID, lookup.UsersAuthProvider).
Where(Eq(lookup.UsersEmail, user.Email)).QueryRow(ctx, q, &userID, &remoteID, &authProvider)
if err == nil {
// remote_id and auth_provider are only filled when still unset.
sets := []Assignment{Set(lookup.UsersLastLoginAt, now), Set(lookup.UsersUpdatedAt, now)}
if !remoteID.Valid {
sets = append(sets, Set(lookup.UsersRemoteID, user.RemoteID))
}
if !authProvider.Valid {
sets = append(sets, Set(lookup.UsersAuthProvider, provider))
}
_, err := o.Update(lookup.EntityUsers).Set(sets...).Where(Eq(lookup.UsersID, userID)).Exec(ctx, q)
return err
}
if !errors.Is(err, sql.ErrNoRows) {
return err
}
id, err := o.Insert(lookup.EntityUsers).Set(
Set(lookup.UsersUsername, user.UserName),
Set(lookup.UsersEmail, user.Email),
Set(lookup.UsersPassword, nil),
Set(lookup.UsersUserLevel, user.UserLevel),
Set(lookup.UsersRoles, roles),
Set(lookup.UsersIsActive, true),
Set(lookup.UsersCreatedAt, now),
Set(lookup.UsersUpdatedAt, now),
Set(lookup.UsersLastLoginAt, now),
Set(lookup.UsersRemoteID, user.RemoteID),
Set(lookup.UsersAuthProvider, provider),
).ExecID(ctx, q, lookup.UsersID)
userID = int(id)
return err
})
if err != nil {
return 0, fmt.Errorf("failed to get or create user: %w", err)
}
return userID, nil
}
// CreateSession implements lookup.OAuthUserStore: insert, or update when the token exists.
func (o *OAuthUsers) CreateSession(ctx context.Context, s lookup.OAuthSession) error {
return o.tx(ctx, func(q Querier) error {
now := o.Now()
exists, err := o.From(lookup.EntityUserSessions).Cols(lookup.SessionsID).Where(Eq(lookup.SessionsToken, s.SessionToken)).Exists(ctx, q)
if err != nil {
return err
}
if exists {
_, err := o.Update(lookup.EntityUserSessions).Set(
Set(lookup.SessionsAccessToken, s.AccessToken),
Set(lookup.SessionsRefreshToken, s.RefreshToken),
Set(lookup.SessionsTokenType, s.TokenType),
Set(lookup.SessionsExpiresAt, s.ExpiresAt),
Set(lookup.SessionsLastActivityAt, now),
).Where(Eq(lookup.SessionsToken, s.SessionToken)).Exec(ctx, q)
return err
}
return o.Insert(lookup.EntityUserSessions).Set(
Set(lookup.SessionsToken, s.SessionToken),
Set(lookup.SessionsUserID, s.UserID),
Set(lookup.SessionsExpiresAt, s.ExpiresAt),
Set(lookup.SessionsCreatedAt, now),
Set(lookup.SessionsLastActivityAt, now),
Set(lookup.SessionsAccessToken, s.AccessToken),
Set(lookup.SessionsRefreshToken, s.RefreshToken),
Set(lookup.SessionsTokenType, s.TokenType),
Set(lookup.SessionsAuthProvider, s.Provider),
).Exec(ctx, q)
})
}
// GetByRefreshToken implements lookup.OAuthUserStore.
func (o *OAuthUsers) GetByRefreshToken(ctx context.Context, refreshToken string) (*lookup.OAuthRefreshSession, error) {
var s lookup.OAuthRefreshSession
var access, tokenType sql.NullString
err := o.do(func(q Querier) error {
return o.From(lookup.EntityUserSessions).
Cols(lookup.SessionsUserID, lookup.SessionsAccessToken, lookup.SessionsTokenType, lookup.SessionsExpiresAt).
Where(Eq(lookup.SessionsRefreshToken, refreshToken), Gt(lookup.SessionsExpiresAt, o.Now())).
QueryRow(ctx, q, &s.UserID, &access, &tokenType, o.timeDest(&s.Expiry))
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, fmt.Errorf("refresh token not found or expired")
}
return nil, fmt.Errorf("failed to get session by refresh token: %w", err)
}
s.AccessToken = access.String
s.TokenType = tokenType.String
return &s, nil
}
// UpdateRefreshToken implements lookup.OAuthUserStore.
func (o *OAuthUsers) UpdateRefreshToken(ctx context.Context, userID int, oldRefreshToken, newSessionToken, newAccessToken, newRefreshToken string, expiresAt time.Time) error {
var rows int64
err := o.do(func(q Querier) error {
var err error
rows, err = o.Update(lookup.EntityUserSessions).Set(
Set(lookup.SessionsToken, newSessionToken),
Set(lookup.SessionsAccessToken, newAccessToken),
Set(lookup.SessionsRefreshToken, newRefreshToken),
Set(lookup.SessionsExpiresAt, expiresAt),
Set(lookup.SessionsLastActivityAt, o.Now()),
).Where(Eq(lookup.SessionsUserID, userID), Eq(lookup.SessionsRefreshToken, oldRefreshToken)).Exec(ctx, q)
return err
})
if err != nil {
return fmt.Errorf("failed to update session: %w", err)
}
if rows == 0 {
return fmt.Errorf("session not found")
}
return nil
}
// GetUser implements lookup.OAuthUserStore.
func (o *OAuthUsers) GetUser(ctx context.Context, userID int) (*sectypes.UserContext, error) {
var u userRow
err := o.do(func(q Querier) error {
return o.From(lookup.EntityUsers).
Cols(lookup.UsersUsername, lookup.UsersEmail, lookup.UsersUserLevel, lookup.UsersRoles,
lookup.UsersProgramUserID, lookup.UsersProgramUserTable).
Where(Eq(lookup.UsersID, userID), Eq(lookup.UsersIsActive, true)).
QueryRow(ctx, q, &u.username, &u.email, &u.userLevel, &u.roles, &u.programUserID, &u.programUserTable)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, fmt.Errorf("user not found")
}
return nil, fmt.Errorf("failed to get user data: %w", err)
}
u.id = userID
return u.context(""), nil
}