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
This commit is contained in:
Hein
2026-10-01 13:19:44 +02:00
parent 60bd0a6dd3
commit c9fa8c60f2
118 changed files with 11218 additions and 5565 deletions
+608
View File
@@ -0,0 +1,608 @@
package direct
import (
"context"
"crypto/rand"
"crypto/sha256"
"database/sql"
"encoding/hex"
"errors"
"fmt"
"strings"
"time"
"github.com/bitechdev/ResolveSpec/pkg/logger"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
const sessionLifetime = 24 * time.Hour
// AuthOptions tunes Auth.
type AuthOptions struct {
// UpgradePasswordHash rewrites legacy cleartext passwords as bcrypt on a successful login.
UpgradePasswordHash bool
}
// Auth implements lookup.AuthStore on the tables. Passwords are verified with bcrypt;
// legacy cleartext rows are accepted at login and only rewritten when UpgradePasswordHash
// is set. Registration never honours client-supplied user_level or roles. Multi-step writes
// (login, register, refresh, reset) run in one transaction.
type Auth struct {
*Base
opts AuthOptions
}
var _ lookup.AuthStore = (*Auth)(nil)
// NewAuth creates the direct AuthStore.
func NewAuth(b *Base, opts AuthOptions) *Auth { return &Auth{Base: b, opts: opts} }
// GenerateSessionToken returns "sess_<64 hex>_<unix>".
func GenerateSessionToken() (string, error) {
buf := make([]byte, 32)
if _, err := rand.Read(buf); err != nil {
return "", err
}
return fmt.Sprintf("sess_%s_%d", hex.EncodeToString(buf), time.Now().Unix()), nil
}
// ParseRoles splits the comma-separated roles column.
func ParseRoles(s string) []string {
if s == "" {
return []string{}
}
return strings.Split(s, ",")
}
func claimStrings(claims map[string]any) (ip, ua string) {
if claims == nil {
return "", ""
}
if v, ok := claims["ip_address"].(string); ok {
ip = v
}
if v, ok := claims["user_agent"].(string); ok {
ua = v
}
return ip, ua
}
func sha256Hex(s string) string {
h := sha256.Sum256([]byte(s))
return hex.EncodeToString(h[:])
}
// userRow is the users columns every session-bearing response needs.
type userRow struct {
id int
username sql.NullString
email sql.NullString
roles sql.NullString
programUserTable sql.NullString
userLevel sql.NullInt64
programUserID sql.NullInt64
}
func (u *userRow) context(sessionID string) *sectypes.UserContext {
return &sectypes.UserContext{
UserID: u.id,
UserName: u.username.String,
Email: u.email.String,
UserLevel: int(u.userLevel.Int64),
SessionID: sessionID,
Roles: ParseRoles(u.roles.String),
ProgramUserID: int(u.programUserID.Int64),
ProgramUserTable: u.programUserTable.String,
}
}
// insertSession writes a session row and stamps the user's last login.
func (a *Auth) insertSession(ctx context.Context, q Querier, token string, userID int64, expiresAt time.Time, ip, ua string, now time.Time) error {
err := a.Insert(lookup.EntityUserSessions).Set(
Set(lookup.SessionsToken, token),
Set(lookup.SessionsUserID, userID),
Set(lookup.SessionsExpiresAt, expiresAt),
Set(lookup.SessionsIPAddress, ip),
Set(lookup.SessionsUserAgent, ua),
Set(lookup.SessionsLastActivityAt, now),
Set(lookup.SessionsCreatedAt, now),
).Exec(ctx, q)
if err != nil {
return err
}
return a.touchLastLogin(ctx, q, userID, now)
}
func (a *Auth) touchLastLogin(ctx context.Context, q Querier, userID int64, now time.Time) error {
_, err := a.Update(lookup.EntityUsers).Set(Set(lookup.UsersLastLoginAt, now)).Where(Eq(lookup.UsersID, userID)).Exec(ctx, q)
return err
}
// Login implements lookup.AuthStore.
func (a *Auth) Login(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) {
var userID int
var email, roles, programUserTable, storedPassword sql.NullString
var userLevel, programUserID sql.NullInt64
err := a.do(func(q Querier) error {
return a.From(lookup.EntityUsers).
Cols(lookup.UsersID, lookup.UsersEmail, lookup.UsersUserLevel, lookup.UsersRoles,
lookup.UsersProgramUserID, lookup.UsersProgramUserTable, lookup.UsersPassword).
Where(Eq(lookup.UsersUsername, req.Username), Eq(lookup.UsersIsActive, true)).
QueryRow(ctx, q, &userID, &email, &userLevel, &roles, &programUserID, &programUserTable, &storedPassword)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
BurnPasswordCheck(req.Password)
return nil, fmt.Errorf("invalid credentials")
}
return nil, fmt.Errorf("login query failed: %w", err)
}
ok, needsRehash := VerifyPassword(storedPassword.String, req.Password)
if !ok {
if storedPassword.String == "" {
BurnPasswordCheck(req.Password)
}
return nil, fmt.Errorf("invalid credentials")
}
if needsRehash && a.opts.UpgradePasswordHash {
a.upgradePasswordHash(ctx, userID, req.Password)
}
token, err := GenerateSessionToken()
if err != nil {
return nil, fmt.Errorf("failed to generate session token: %w", err)
}
now := a.Now()
ip, ua := claimStrings(req.Claims)
err = a.tx(ctx, func(q Querier) error {
return a.insertSession(ctx, q, token, int64(userID), now.Add(sessionLifetime), ip, ua, now)
})
if err != nil {
return nil, fmt.Errorf("login query failed: %w", err)
}
return &sectypes.LoginResponse{
Token: token,
User: &sectypes.UserContext{
UserID: userID,
UserName: req.Username,
Email: email.String,
UserLevel: int(userLevel.Int64),
Roles: ParseRoles(roles.String),
SessionID: token,
ProgramUserID: int(programUserID.Int64),
ProgramUserTable: programUserTable.String,
},
ExpiresIn: int64(sessionLifetime.Seconds()),
}, nil
}
// upgradePasswordHash replaces a legacy cleartext password with a bcrypt hash. Failure is
// logged and ignored: the login itself already succeeded.
func (a *Auth) upgradePasswordHash(ctx context.Context, userID int, password string) {
h, err := HashPassword(password)
if err != nil {
return
}
err = a.do(func(q Querier) error {
_, err := a.Update(lookup.EntityUsers).
Set(Set(lookup.UsersPassword, h), Set(lookup.UsersUpdatedAt, a.Now())).
Where(Eq(lookup.UsersID, userID)).Exec(ctx, q)
return err
})
if err != nil {
logger.Warn("failed to upgrade legacy password hash for user %d: %v", userID, err)
}
}
// Register implements lookup.AuthStore.
func (a *Auth) Register(ctx context.Context, req sectypes.RegisterRequest) (*sectypes.LoginResponse, error) {
if req.Username == "" {
return nil, fmt.Errorf("username is required")
}
if req.Email == "" {
return nil, fmt.Errorf("email is required")
}
if req.Password == "" {
return nil, fmt.Errorf("password is required")
}
hash, err := HashPassword(req.Password)
if err != nil {
return nil, err
}
token, err := GenerateSessionToken()
if err != nil {
return nil, fmt.Errorf("failed to generate session token: %w", err)
}
// Privileges are never taken from the request: self-registration always creates an
// unprivileged user.
const userLevel = 0
now := a.Now()
ip, ua := claimStrings(req.Claims)
var userID int64
err = a.tx(ctx, func(q Querier) error {
exists, err := a.From(lookup.EntityUsers).Cols(lookup.UsersID).Where(Eq(lookup.UsersUsername, req.Username)).Exists(ctx, q)
if err != nil {
return err
}
if exists {
return lookup.ErrUsernameExists
}
exists, err = a.From(lookup.EntityUsers).Cols(lookup.UsersID).Where(Eq(lookup.UsersEmail, req.Email)).Exists(ctx, q)
if err != nil {
return err
}
if exists {
return lookup.ErrEmailExists
}
userID, err = a.Insert(lookup.EntityUsers).Set(
Set(lookup.UsersUsername, req.Username),
Set(lookup.UsersEmail, req.Email),
Set(lookup.UsersPassword, hash),
Set(lookup.UsersUserLevel, userLevel),
Set(lookup.UsersRoles, ""),
Set(lookup.UsersIsActive, true),
Set(lookup.UsersCreatedAt, now),
Set(lookup.UsersUpdatedAt, now),
Set(lookup.UsersProgramUserID, 0),
Set(lookup.UsersProgramUserTable, ""),
).ExecID(ctx, q, lookup.UsersID)
if err != nil {
return err
}
return a.insertSession(ctx, q, token, userID, now.Add(sessionLifetime), ip, ua, now)
})
if err != nil {
if errors.Is(err, lookup.ErrUsernameExists) || errors.Is(err, lookup.ErrEmailExists) {
return nil, err
}
return nil, fmt.Errorf("register query failed: %w", err)
}
return &sectypes.LoginResponse{
Token: token,
User: &sectypes.UserContext{
UserID: int(userID),
UserName: req.Username,
Email: req.Email,
UserLevel: userLevel,
Roles: ParseRoles(""),
SessionID: token,
},
ExpiresIn: int64(sessionLifetime.Seconds()),
}, nil
}
// Logout implements lookup.AuthStore.
func (a *Auth) Logout(ctx context.Context, req sectypes.LogoutRequest) error {
token := strings.TrimPrefix(strings.TrimPrefix(req.Token, "Bearer "), "bearer ")
var rows int64
err := a.do(func(q Querier) error {
var err error
rows, err = a.Delete(lookup.EntityUserSessions).
Where(Eq(lookup.SessionsToken, token), Eq(lookup.SessionsUserID, req.UserID)).Exec(ctx, q)
return err
})
if err != nil {
return fmt.Errorf("logout query failed: %w", err)
}
if rows == 0 {
return fmt.Errorf("session not found")
}
return nil
}
// sessionUser selects the user behind a live session token.
func (a *Auth) sessionUser(ctx context.Context, q Querier, token string, extra ...lookup.Column) (*userRow, []any, error) {
var u userRow
dest := []any{&u.id, &u.username, &u.email, &u.userLevel, &u.roles, &u.programUserID, &u.programUserTable}
cols := []lookup.Column{lookup.SessionsUserID, lookup.UsersUsername, lookup.UsersEmail, lookup.UsersUserLevel,
lookup.UsersRoles, lookup.UsersProgramUserID, lookup.UsersProgramUserTable}
extras := make([]any, len(extra))
for i, c := range extra {
cols = append(cols, c)
extras[i] = new(sql.NullString)
dest = append(dest, extras[i])
}
err := a.From(lookup.EntityUserSessions).Cols(cols...).
Join(lookup.EntityUsers, EqCol(lookup.SessionsUserID, lookup.UsersID)).
Where(Eq(lookup.SessionsToken, token), Gt(lookup.SessionsExpiresAt, a.Now()), Eq(lookup.UsersIsActive, true)).
QueryRow(ctx, q, dest...)
return &u, extras, err
}
// Session implements lookup.AuthStore. reference is only meaningful to the procedure backend.
func (a *Auth) Session(ctx context.Context, token, _ string) (*sectypes.UserContext, error) {
var u *userRow
err := a.do(func(q Querier) error {
var err error
u, _, err = a.sessionUser(ctx, q, token)
return err
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, fmt.Errorf("invalid or expired session")
}
return nil, fmt.Errorf("session query failed: %w", err)
}
return u.context(token), nil
}
// TouchSession implements lookup.AuthStore.
func (a *Auth) TouchSession(ctx context.Context, token string, _ *sectypes.UserContext) error {
return a.do(func(q Querier) error {
now := a.Now()
_, err := a.Update(lookup.EntityUserSessions).Set(Set(lookup.SessionsLastActivityAt, now)).
Where(Eq(lookup.SessionsToken, token), Gt(lookup.SessionsExpiresAt, now)).Exec(ctx, q)
return err
})
}
// Refresh implements lookup.AuthStore: the old session is replaced by a new one.
func (a *Auth) Refresh(ctx context.Context, oldToken string) (*sectypes.LoginResponse, error) {
var u *userRow
var extras []any
err := a.do(func(q Querier) error {
var err error
u, extras, err = a.sessionUser(ctx, q, oldToken, lookup.SessionsIPAddress, lookup.SessionsUserAgent)
return err
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, fmt.Errorf("invalid or expired refresh token")
}
return nil, fmt.Errorf("refresh token query failed: %w", err)
}
ip := extras[0].(*sql.NullString).String
ua := extras[1].(*sql.NullString).String
newToken, err := GenerateSessionToken()
if err != nil {
return nil, fmt.Errorf("failed to generate session token: %w", err)
}
now := a.Now()
err = a.tx(ctx, func(q Querier) error {
err := a.Insert(lookup.EntityUserSessions).Set(
Set(lookup.SessionsToken, newToken),
Set(lookup.SessionsUserID, u.id),
Set(lookup.SessionsExpiresAt, now.Add(sessionLifetime)),
Set(lookup.SessionsIPAddress, ip),
Set(lookup.SessionsUserAgent, ua),
Set(lookup.SessionsLastActivityAt, now),
Set(lookup.SessionsCreatedAt, now),
).Exec(ctx, q)
if err != nil {
return err
}
_, err = a.Delete(lookup.EntityUserSessions).Where(Eq(lookup.SessionsToken, oldToken)).Exec(ctx, q)
return err
})
if err != nil {
return nil, fmt.Errorf("refresh token generation failed: %w", err)
}
return &sectypes.LoginResponse{
Token: newToken,
User: u.context(newToken),
ExpiresIn: int64(sessionLifetime.Seconds()),
}, nil
}
// apiKeyTypes are the key types accepted by LoginAPIKey.
var apiKeyTypes = []any{string(sectypes.KeyTypeHeaderAPI), string(sectypes.KeyTypeGenericAPI)}
// LoginAPIKey implements lookup.AuthStore. Unknown, expired, inactive and wrong-type keys
// (and inactive users) all return lookup.ErrInvalidAPIKey; the raw key is never logged.
func (a *Auth) LoginAPIKey(ctx context.Context, rawKey string, claims map[string]any) (*sectypes.LoginResponse, error) {
if rawKey == "" {
return nil, lookup.ErrInvalidAPIKey
}
now := a.Now()
var keyID int64
var u userRow
err := a.do(func(q Querier) error {
return a.From(lookup.EntityUserKeys).
Cols(lookup.KeysID, lookup.UsersID, lookup.UsersUsername, lookup.UsersEmail, lookup.UsersUserLevel,
lookup.UsersRoles, lookup.UsersProgramUserID, lookup.UsersProgramUserTable).
Join(lookup.EntityUsers, EqCol(lookup.KeysUserID, lookup.UsersID)).
Where(
Eq(lookup.KeysKeyHash, sectypes.HashKey(rawKey)),
In(lookup.KeysKeyType, apiKeyTypes...),
Eq(lookup.KeysIsActive, true),
Or(IsNull(lookup.KeysExpiresAt), Gt(lookup.KeysExpiresAt, now)),
Eq(lookup.UsersIsActive, true),
).QueryRow(ctx, q, &keyID, &u.id, &u.username, &u.email, &u.userLevel, &u.roles, &u.programUserID, &u.programUserTable)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, lookup.ErrInvalidAPIKey
}
return nil, fmt.Errorf("api key login query failed: %w", err)
}
token, err := GenerateSessionToken()
if err != nil {
return nil, fmt.Errorf("failed to generate session token: %w", err)
}
ip, ua := claimStrings(claims)
err = a.tx(ctx, func(q Querier) error {
if err := a.insertSession(ctx, q, token, int64(u.id), now.Add(sessionLifetime), ip, ua, now); err != nil {
return err
}
_, err := a.Update(lookup.EntityUserKeys).Set(Set(lookup.KeysLastUsedAt, now)).Where(Eq(lookup.KeysID, keyID)).Exec(ctx, q)
return err
})
if err != nil {
return nil, fmt.Errorf("api key login query failed: %w", err)
}
return &sectypes.LoginResponse{
Token: token,
User: u.context(token),
ExpiresIn: int64(sessionLifetime.Seconds()),
}, nil
}
// JWTLogin implements lookup.AuthStore (mirrors resolvespec_jwt_login). The token is a
// placeholder until JWT signing is wired in.
func (a *Auth) JWTLogin(ctx context.Context, req sectypes.LoginRequest) (*sectypes.LoginResponse, error) {
var userID int
var email, roles, storedPassword sql.NullString
var userLevel sql.NullInt64
err := a.do(func(q Querier) error {
return a.From(lookup.EntityUsers).
Cols(lookup.UsersID, lookup.UsersEmail, lookup.UsersUserLevel, lookup.UsersRoles, lookup.UsersPassword).
Where(Eq(lookup.UsersUsername, req.Username), Eq(lookup.UsersIsActive, true)).
QueryRow(ctx, q, &userID, &email, &userLevel, &roles, &storedPassword)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
BurnPasswordCheck(req.Password)
return nil, fmt.Errorf("invalid credentials")
}
return nil, fmt.Errorf("login query failed: %w", err)
}
ok, needsRehash := VerifyPassword(storedPassword.String, req.Password)
if !ok {
if storedPassword.String == "" {
BurnPasswordCheck(req.Password)
}
return nil, fmt.Errorf("invalid credentials")
}
if needsRehash && a.opts.UpgradePasswordHash {
a.upgradePasswordHash(ctx, userID, req.Password)
}
expiresAt := a.Now().Add(sessionLifetime)
return &sectypes.LoginResponse{
Token: fmt.Sprintf("token_%d_%d", userID, expiresAt.Unix()),
User: &sectypes.UserContext{
UserID: userID,
UserName: req.Username,
Email: email.String,
UserLevel: int(userLevel.Int64),
Roles: ParseRoles(roles.String),
},
ExpiresIn: int64(sessionLifetime.Seconds()),
}, nil
}
// JWTLogout implements lookup.AuthStore: the token goes on the blacklist.
func (a *Auth) JWTLogout(ctx context.Context, req sectypes.LogoutRequest) error {
now := a.Now()
err := a.do(func(q Querier) error {
return a.Insert(lookup.EntityTokenBlacklist).Set(
Set(lookup.BlacklistToken, req.Token),
Set(lookup.BlacklistUserID, req.UserID),
Set(lookup.BlacklistExpiresAt, now.Add(sessionLifetime)),
Set(lookup.BlacklistCreatedAt, now),
).Exec(ctx, q)
})
if err != nil {
return fmt.Errorf("logout query failed: %w", err)
}
return nil
}
// ResetRequest implements lookup.AuthStore. An unknown user yields a generic empty success
// so accounts cannot be enumerated.
func (a *Auth) ResetRequest(ctx context.Context, req sectypes.PasswordResetRequest) (*sectypes.PasswordResetResponse, error) {
if req.Email == "" && req.Username == "" {
return nil, fmt.Errorf("email or username is required")
}
var userID int
err := a.do(func(q Querier) error {
lookupCol, val := lookup.UsersUsername, req.Username
if req.Email != "" {
lookupCol, val = lookup.UsersEmail, req.Email
}
return a.From(lookup.EntityUsers).Cols(lookup.UsersID).
Where(Eq(lookupCol, val), Eq(lookup.UsersIsActive, true)).QueryRow(ctx, q, &userID)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return &sectypes.PasswordResetResponse{Token: "", ExpiresIn: 0}, nil
}
return nil, fmt.Errorf("password reset request query failed: %w", err)
}
raw := make([]byte, 32)
if _, err := rand.Read(raw); err != nil {
return nil, fmt.Errorf("failed to generate reset token: %w", err)
}
rawToken := hex.EncodeToString(raw)
now := a.Now()
err = a.tx(ctx, func(q Querier) error {
if _, err := a.Delete(lookup.EntityUserPasswordResets).
Where(Eq(lookup.ResetsUserID, userID), Eq(lookup.ResetsUsed, false)).Exec(ctx, q); err != nil {
return err
}
return a.Insert(lookup.EntityUserPasswordResets).Set(
Set(lookup.ResetsUserID, userID),
Set(lookup.ResetsTokenHash, sha256Hex(rawToken)),
Set(lookup.ResetsExpiresAt, now.Add(time.Hour)),
Set(lookup.ResetsCreatedAt, now),
Set(lookup.ResetsUsed, false),
).Exec(ctx, q)
})
if err != nil {
return nil, fmt.Errorf("password reset request query failed: %w", err)
}
return &sectypes.PasswordResetResponse{Token: rawToken, ExpiresIn: 3600}, nil
}
// ResetComplete implements lookup.AuthStore: sets the new password, ends every session of
// the user and consumes the reset token, atomically.
func (a *Auth) ResetComplete(ctx context.Context, req sectypes.PasswordResetCompleteRequest) error {
if req.Token == "" {
return fmt.Errorf("token is required")
}
if req.NewPassword == "" {
return fmt.Errorf("new_password is required")
}
newHash, err := HashPassword(req.NewPassword)
if err != nil {
return err
}
tokenHash := sha256Hex(req.Token)
now := a.Now()
var resetID, userID int
var expiresAt time.Time
err = a.do(func(q Querier) error {
return a.From(lookup.EntityUserPasswordResets).
Cols(lookup.ResetsID, lookup.ResetsUserID, lookup.ResetsExpiresAt).
Where(Eq(lookup.ResetsTokenHash, tokenHash), Eq(lookup.ResetsUsed, false)).
QueryRow(ctx, q, &resetID, &userID, a.timeDest(&expiresAt))
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return fmt.Errorf("invalid or expired token")
}
return fmt.Errorf("password reset complete query failed: %w", err)
}
if !expiresAt.After(now) {
return fmt.Errorf("invalid or expired token")
}
err = a.tx(ctx, func(q Querier) error {
if _, err := a.Update(lookup.EntityUsers).
Set(Set(lookup.UsersPassword, newHash), Set(lookup.UsersUpdatedAt, now)).
Where(Eq(lookup.UsersID, userID)).Exec(ctx, q); err != nil {
return err
}
if _, err := a.Delete(lookup.EntityUserSessions).Where(Eq(lookup.SessionsUserID, userID)).Exec(ctx, q); err != nil {
return err
}
_, err := a.Update(lookup.EntityUserPasswordResets).
Set(Set(lookup.ResetsUsed, true), Set(lookup.ResetsUsedAt, now)).
Where(Eq(lookup.ResetsID, resetID)).Exec(ctx, q)
return err
})
if err != nil {
return fmt.Errorf("password reset complete query failed: %w", err)
}
return nil
}
+268
View File
@@ -0,0 +1,268 @@
package direct
import (
"context"
"database/sql"
"errors"
"strings"
"testing"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
func newAuth(t *testing.T, opts AuthOptions) (*Auth, *sql.DB) {
db := newTestDB(t)
return NewAuth(newTestBase(t, db, nil), opts), db
}
func registerUser(t *testing.T, a *Auth, name string) *sectypes.LoginResponse {
t.Helper()
resp, err := a.Register(context.Background(), sectypes.RegisterRequest{Username: name, Email: name + "@x.io", Password: "pw-" + name})
if err != nil {
t.Fatal(err)
}
return resp
}
func TestRegisterLoginSessionFlow(t *testing.T) {
ctx := context.Background()
a, db := newAuth(t, AuthOptions{})
reg, err := a.Register(ctx, sectypes.RegisterRequest{
Username: "ann", Email: "ann@x.io", Password: "secret",
UserLevel: 99, Roles: []string{"admin"}, // must be ignored
})
if err != nil {
t.Fatal(err)
}
if reg.User.UserLevel != 0 || len(reg.User.Roles) != 0 {
t.Fatalf("register honoured privileges: %+v", reg.User)
}
var stored string
if err := db.QueryRow(`SELECT password FROM users WHERE username='ann'`).Scan(&stored); err != nil || !strings.HasPrefix(stored, "$2") {
t.Fatalf("password not bcrypt: %q %v", stored, err)
}
if _, err := a.Register(ctx, sectypes.RegisterRequest{Username: "ann", Email: "other@x.io", Password: "x"}); !errors.Is(err, lookup.ErrUsernameExists) {
t.Fatalf("got %v", err)
}
if _, err := a.Register(ctx, sectypes.RegisterRequest{Username: "bob", Email: "ann@x.io", Password: "x"}); !errors.Is(err, lookup.ErrEmailExists) {
t.Fatalf("got %v", err)
}
var n int
_ = db.QueryRow(`SELECT COUNT(*) FROM users`).Scan(&n)
if n != 1 {
t.Fatalf("failed register left a row: %d", n)
}
login, err := a.Login(ctx, sectypes.LoginRequest{Username: "ann", Password: "secret", Claims: map[string]any{"ip_address": "1.2.3.4", "user_agent": "ua"}})
if err != nil {
t.Fatal(err)
}
if !strings.HasPrefix(login.Token, "sess_") || login.ExpiresIn != 86400 || login.User.Email != "ann@x.io" {
t.Fatalf("login: %+v", login)
}
var ip string
_ = db.QueryRow(`SELECT ip_address FROM user_sessions WHERE session_token=?`, login.Token).Scan(&ip)
if ip != "1.2.3.4" {
t.Fatalf("ip %q", ip)
}
if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "ann", Password: "wrong"}); err == nil || err.Error() != "invalid credentials" {
t.Fatalf("got %v", err)
}
if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "nobody", Password: "x"}); err == nil || err.Error() != "invalid credentials" {
t.Fatalf("got %v", err)
}
u, err := a.Session(ctx, login.Token, "authenticate")
if err != nil || u.UserName != "ann" || u.SessionID != login.Token {
t.Fatalf("session: %+v %v", u, err)
}
if err := a.TouchSession(ctx, login.Token, u); err != nil {
t.Fatal(err)
}
if _, err := a.Session(ctx, "nope", "authenticate"); err == nil || err.Error() != "invalid or expired session" {
t.Fatalf("got %v", err)
}
ref, err := a.Refresh(ctx, login.Token)
if err != nil || ref.Token == login.Token {
t.Fatalf("refresh: %+v %v", ref, err)
}
if _, err := a.Session(ctx, login.Token, ""); err == nil {
t.Fatal("old session still valid after refresh")
}
if _, err := a.Refresh(ctx, login.Token); err == nil || err.Error() != "invalid or expired refresh token" {
t.Fatalf("got %v", err)
}
if err := a.Logout(ctx, sectypes.LogoutRequest{Token: "Bearer " + ref.Token, UserID: ref.User.UserID}); err != nil {
t.Fatal(err)
}
if err := a.Logout(ctx, sectypes.LogoutRequest{Token: ref.Token, UserID: ref.User.UserID}); err == nil || err.Error() != "session not found" {
t.Fatalf("got %v", err)
}
}
func TestExpiredSessionRejected(t *testing.T) {
ctx := context.Background()
a, _ := newAuth(t, AuthOptions{})
resp := registerUser(t, a, "eve")
a.Now = func() time.Time { return time.Now().Add(48 * time.Hour) }
if _, err := a.Session(ctx, resp.Token, ""); err == nil {
t.Fatal("expired session accepted")
}
}
func TestLegacyPasswordUpgradeIsOptIn(t *testing.T) {
ctx := context.Background()
for _, upgrade := range []bool{false, true} {
a, db := newAuth(t, AuthOptions{UpgradePasswordHash: upgrade})
_, err := db.Exec(`INSERT INTO users (username, email, password, user_level, roles, is_active) VALUES ('old','o@x.io','clear',1,'a,b',1)`)
if err != nil {
t.Fatal(err)
}
resp, err := a.Login(ctx, sectypes.LoginRequest{Username: "old", Password: "clear"})
if err != nil || len(resp.User.Roles) != 2 {
t.Fatalf("login: %+v %v", resp, err)
}
var stored string
_ = db.QueryRow(`SELECT password FROM users WHERE username='old'`).Scan(&stored)
if got := strings.HasPrefix(stored, "$2"); got != upgrade {
t.Fatalf("upgrade=%v stored=%q", upgrade, stored)
}
}
}
func TestInactiveUserCannotLogin(t *testing.T) {
ctx := context.Background()
a, db := newAuth(t, AuthOptions{})
resp := registerUser(t, a, "ian")
_, _ = db.Exec(`UPDATE users SET is_active = 0`)
if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "ian", Password: "pw-ian"}); err == nil {
t.Fatal("inactive login accepted")
}
if _, err := a.Session(ctx, resp.Token, ""); err == nil {
t.Fatal("inactive session accepted")
}
}
func TestLoginAPIKey(t *testing.T) {
ctx := context.Background()
a, db := newAuth(t, AuthOptions{})
reg := registerUser(t, a, "kim")
insert := func(raw, typ string, active int, expires any) {
t.Helper()
_, err := db.Exec(`INSERT INTO user_keys (user_id, key_type, key_hash, name, is_active, expires_at) VALUES (?,?,?,?,?,?)`,
reg.User.UserID, typ, sectypes.HashKey(raw), "k", active, expires)
if err != nil {
t.Fatal(err)
}
}
insert("good", "header_api", 1, nil)
insert("generic", "api", 1, nil)
insert("jwt", "jwt_secret", 1, nil)
insert("off", "api", 0, nil)
insert("old", "api", 1, time.Now().Add(-time.Hour))
for _, k := range []string{"good", "generic"} {
resp, err := a.LoginAPIKey(ctx, k, map[string]any{"ip_address": "9.9.9.9"})
if err != nil || resp.User.UserName != "kim" || !strings.HasPrefix(resp.Token, "sess_") {
t.Fatalf("%s: %+v %v", k, resp, err)
}
if _, err := a.Session(ctx, resp.Token, ""); err != nil {
t.Fatal(err)
}
}
var used sql.NullString
_ = db.QueryRow(`SELECT last_used_at FROM user_keys WHERE key_hash = ?`, sectypes.HashKey("good")).Scan(&used)
if !used.Valid {
t.Fatal("last_used_at not stamped")
}
for _, k := range []string{"", "missing", "jwt", "off", "old"} {
if _, err := a.LoginAPIKey(ctx, k, nil); !errors.Is(err, lookup.ErrInvalidAPIKey) {
t.Fatalf("%q: got %v", k, err)
}
}
_, _ = db.Exec(`UPDATE users SET is_active = 0`)
if _, err := a.LoginAPIKey(ctx, "good", nil); !errors.Is(err, lookup.ErrInvalidAPIKey) {
t.Fatalf("inactive user: got %v", err)
}
}
func TestPasswordReset(t *testing.T) {
ctx := context.Background()
a, _ := newAuth(t, AuthOptions{})
reg := registerUser(t, a, "rae")
empty, err := a.ResetRequest(ctx, sectypes.PasswordResetRequest{Email: "none@x.io"})
if err != nil || empty.Token != "" {
t.Fatalf("enumeration leak: %+v %v", empty, err)
}
if _, err := a.ResetRequest(ctx, sectypes.PasswordResetRequest{}); err == nil {
t.Fatal("expected error")
}
r1, err := a.ResetRequest(ctx, sectypes.PasswordResetRequest{Email: "rae@x.io"})
if err != nil || r1.Token == "" {
t.Fatal(err)
}
r2, _ := a.ResetRequest(ctx, sectypes.PasswordResetRequest{Username: "rae"})
if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: r1.Token, NewPassword: "n"}); err == nil {
t.Fatal("superseded token accepted")
}
if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: r2.Token, NewPassword: "newpw"}); err != nil {
t.Fatal(err)
}
if err := a.ResetComplete(ctx, sectypes.PasswordResetCompleteRequest{Token: r2.Token, NewPassword: "again"}); err == nil {
t.Fatal("token reused")
}
if _, err := a.Session(ctx, reg.Token, ""); err == nil {
t.Fatal("sessions survived reset")
}
if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "rae", Password: "newpw"}); err != nil {
t.Fatal(err)
}
}
func TestJWTLoginLogout(t *testing.T) {
ctx := context.Background()
a, db := newAuth(t, AuthOptions{})
reg := registerUser(t, a, "jay")
resp, err := a.JWTLogin(ctx, sectypes.LoginRequest{Username: "jay", Password: "pw-jay"})
if err != nil || !strings.HasPrefix(resp.Token, "token_") {
t.Fatalf("%+v %v", resp, err)
}
if err := a.JWTLogout(ctx, sectypes.LogoutRequest{Token: "tok", UserID: reg.User.UserID}); err != nil {
t.Fatal(err)
}
var n int
_ = db.QueryRow(`SELECT COUNT(*) FROM token_blacklist WHERE token='tok'`).Scan(&n)
if n != 1 {
t.Fatal("token not blacklisted")
}
}
func TestCustomSchemaNames(t *testing.T) {
db := newTestDB(t, `
CREATE TABLE app_users (uid INTEGER PRIMARY KEY AUTOINCREMENT, login TEXT, email TEXT, password TEXT,
user_level INTEGER, roles TEXT, is_active INTEGER, created_at DATETIME, updated_at DATETIME,
last_login_at DATETIME, program_user_id INTEGER, program_user_table TEXT, remote_id TEXT, auth_provider TEXT,
totp_secret TEXT, totp_enabled INTEGER, totp_enabled_at DATETIME);`)
schema := lookup.Schema{lookup.EntityUsers: {Name: "app_users", Columns: map[string]string{"id": "uid", "username": "login"}}}
a := NewAuth(newTestBase(t, db, schema), AuthOptions{})
ctx := context.Background()
if _, err := a.Register(ctx, sectypes.RegisterRequest{Username: "zed", Email: "z@x.io", Password: "p"}); err != nil {
t.Fatal(err)
}
if _, err := a.Login(ctx, sectypes.LoginRequest{Username: "zed", Password: "p"}); err != nil {
t.Fatal(err)
}
var login string
if err := db.QueryRow(`SELECT login FROM app_users`).Scan(&login); err != nil || login != "zed" {
t.Fatalf("%q %v", login, err)
}
}
+462
View File
@@ -0,0 +1,462 @@
// Package direct is the table-backed implementation of the lookup stores. SQL is built
// from the configured Schema (table and column names) and Dialect (placeholders, quoting,
// booleans, insert-returning-id); no statement is written per database and no ORM is used.
package direct
import (
"context"
"database/sql"
"fmt"
"strings"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect"
)
// Runner runs a database operation, reconnecting once when the *sql.DB has been closed.
// procedure.Runner (and procedure.DB) satisfy it.
type Runner interface {
Run(run func(*sql.DB) error) error
}
// Querier is implemented by *sql.DB and *sql.Tx.
type Querier interface {
ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error)
QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error)
QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row
}
// Base is the state shared by every direct store: the runner, dialect, schema and clock.
type Base struct {
run Runner
d dialect.Dialect
schema lookup.Schema
// Now is the clock; tests replace it.
Now func() time.Time
}
// NewBase creates the shared state. The schema is merged with the defaults and validated.
func NewBase(run Runner, d dialect.Dialect, schema lookup.Schema) (*Base, error) {
if run == nil {
return nil, fmt.Errorf("direct: nil runner")
}
if d == nil {
return nil, fmt.Errorf("direct: nil dialect")
}
merged := lookup.DefaultSchema().Merge(schema)
if err := merged.Validate(); err != nil {
return nil, err
}
return &Base{run: run, d: d, schema: merged, Now: time.Now}, nil
}
// Dialect returns the dialect in use.
func (b *Base) Dialect() dialect.Dialect { return b.d }
// do runs fn against the database without a transaction.
func (b *Base) do(fn func(q Querier) error) error {
return b.run.Run(func(db *sql.DB) error { return fn(db) })
}
// tx runs fn in one transaction; an error rolls back.
func (b *Base) tx(ctx context.Context, fn func(q Querier) error) error {
return b.run.Run(func(db *sql.DB) error {
tx, err := db.BeginTx(ctx, nil)
if err != nil {
return err
}
if err := fn(tx); err != nil {
_ = tx.Rollback()
return err
}
return tx.Commit()
})
}
// tableRef returns the (possibly schema-qualified) physical table name of an entity.
func (b *Base) tableRef(e lookup.Entity) string {
t := b.schema[e]
name := t.Name
if name == "" {
name = string(e)
}
if t.Schema != "" {
return t.Schema + "." + name
}
return name
}
// colName returns the physical column name of a logical column.
func (b *Base) colName(c lookup.Column) string {
if t, ok := b.schema[c.Entity]; ok {
if n := t.Columns[c.Name]; n != "" {
return n
}
}
return c.Name
}
// arg converts a Go value to a bind argument (booleans go through the dialect).
func (b *Base) arg(v any) any {
if bv, ok := v.(bool); ok {
return b.d.Bool(bv)
}
return v
}
// timeDest scans a time column through the dialect, so drivers returning strings work.
type timeDest struct {
d dialect.Dialect
v *time.Time
}
func (t timeDest) Scan(src any) error {
v, err := t.d.ScanTime(src)
if err != nil {
return err
}
*t.v = v
return nil
}
type boolDest struct {
d dialect.Dialect
v *bool
}
func (t boolDest) Scan(src any) error {
v, err := t.d.ScanBool(src)
if err != nil {
return err
}
*t.v = v
return nil
}
func (b *Base) timeDest(v *time.Time) sql.Scanner { return timeDest{d: b.d, v: v} }
func (b *Base) boolDest(v *bool) sql.Scanner { return boolDest{d: b.d, v: v} }
// --- query builder --------------------------------------------------------
// builder accumulates bind arguments and renders column references.
type builder struct {
b *Base
args []any
aliases map[lookup.Entity]string
nalias int
}
func (bl *builder) ph(v any) string {
bl.args = append(bl.args, bl.b.arg(v))
return bl.b.d.Placeholder(len(bl.args))
}
// col renders a column; with aliases set (select queries) it is qualified by its table alias.
func (bl *builder) col(c lookup.Column) string {
name := bl.b.d.Quote(bl.b.colName(c))
if bl.aliases != nil {
if a, ok := bl.aliases[c.Entity]; ok {
return a + "." + name
}
}
return name
}
// Cond renders one boolean condition.
type Cond func(*builder) string
// Eq is `col = value`.
func Eq(c lookup.Column, v any) Cond {
return func(bl *builder) string { return bl.col(c) + " = " + bl.ph(v) }
}
// EqFold is a case-insensitive `LOWER(col) = value` match (the value is lowered in Go).
func EqFold(c lookup.Column, v string) Cond {
return func(bl *builder) string { return "LOWER(" + bl.col(c) + ") = " + bl.ph(strings.ToLower(v)) }
}
// Ne is `col <> value`.
func Ne(c lookup.Column, v any) Cond {
return func(bl *builder) string { return bl.col(c) + " <> " + bl.ph(v) }
}
// Gt is `col > value`.
func Gt(c lookup.Column, v any) Cond {
return func(bl *builder) string { return bl.col(c) + " > " + bl.ph(v) }
}
// IsNull is `col IS NULL`.
func IsNull(c lookup.Column) Cond {
return func(bl *builder) string { return bl.col(c) + " IS NULL" }
}
// EqCol is `a = b` between two columns (join conditions).
func EqCol(a, c lookup.Column) Cond {
return func(bl *builder) string { return bl.col(a) + " = " + bl.col(c) }
}
// In is `col IN (v...)`; an empty list renders a condition that is never true.
func In(c lookup.Column, vs ...any) Cond {
return func(bl *builder) string {
if len(vs) == 0 {
return "1 = 0"
}
ph := make([]string, len(vs))
for i, v := range vs {
ph[i] = bl.ph(v)
}
return bl.col(c) + " IN (" + strings.Join(ph, ", ") + ")"
}
}
// InSelect is `col IN (subselect)`; the subselect's arguments share the outer numbering.
func InSelect(c lookup.Column, sub *Select) Cond {
return func(bl *builder) string { return bl.col(c) + " IN (" + sub.render(bl) + ")" }
}
// Or joins conditions with OR inside parentheses.
func Or(cs ...Cond) Cond { return joinConds("OR", cs) }
// And joins conditions with AND inside parentheses.
func And(cs ...Cond) Cond { return joinConds("AND", cs) }
func joinConds(op string, cs []Cond) Cond {
return func(bl *builder) string {
parts := make([]string, len(cs))
for i, c := range cs {
parts[i] = c(bl)
}
return "(" + strings.Join(parts, " "+op+" ") + ")"
}
}
func (bl *builder) where(cs []Cond) string {
if len(cs) == 0 {
return ""
}
parts := make([]string, len(cs))
for i, c := range cs {
parts[i] = c(bl)
}
return " WHERE " + strings.Join(parts, " AND ")
}
// Stmt is a rendered statement.
type Stmt struct {
SQL string
Args []any
}
// Select builds a SELECT.
type Select struct {
b *Base
from lookup.Entity
joins []join
cols []lookup.Column
conds []Cond
order []lookup.Column
}
type join struct {
e lookup.Entity
on Cond
}
// From starts a SELECT on e.
func (b *Base) From(e lookup.Entity) *Select { return &Select{b: b, from: e} }
// Cols sets the selected columns.
func (s *Select) Cols(cs ...lookup.Column) *Select { s.cols = cs; return s }
// Join adds `JOIN e ON on`.
func (s *Select) Join(e lookup.Entity, on Cond) *Select {
s.joins = append(s.joins, join{e: e, on: on})
return s
}
// Where adds AND-ed conditions.
func (s *Select) Where(cs ...Cond) *Select { s.conds = append(s.conds, cs...); return s }
// OrderBy adds ascending order columns.
func (s *Select) OrderBy(cs ...lookup.Column) *Select { s.order = append(s.order, cs...); return s }
// Build renders the statement.
func (s *Select) Build() Stmt {
bl := &builder{b: s.b}
sqlText := s.render(bl)
return Stmt{SQL: sqlText, Args: bl.args}
}
// render writes the select into bl, giving every table a fresh alias so a subselect cannot
// clash with the statement around it.
func (s *Select) render(bl *builder) string {
saved := bl.aliases
defer func() { bl.aliases = saved }()
bl.aliases = map[lookup.Entity]string{}
alias := func() string { a := fmt.Sprintf("t%d", bl.nalias); bl.nalias++; return a }
bl.aliases[s.from] = alias()
for _, j := range s.joins {
bl.aliases[j.e] = alias()
}
sel := make([]string, len(s.cols))
for i, c := range s.cols {
sel[i] = bl.col(c)
}
var sb strings.Builder
sb.WriteString("SELECT " + strings.Join(sel, ", "))
sb.WriteString(" FROM " + s.b.d.Quote(s.b.tableRef(s.from)) + " " + bl.aliases[s.from])
for _, j := range s.joins {
sb.WriteString(" JOIN " + s.b.d.Quote(s.b.tableRef(j.e)) + " " + bl.aliases[j.e] + " ON " + j.on(bl))
}
sb.WriteString(bl.where(s.conds))
if len(s.order) > 0 {
o := make([]string, len(s.order))
for i, c := range s.order {
o[i] = bl.col(c)
}
sb.WriteString(" ORDER BY " + strings.Join(o, ", "))
}
return sb.String()
}
// QueryRow runs the select and scans the first row into dest.
func (s *Select) QueryRow(ctx context.Context, q Querier, dest ...any) error {
st := s.Build()
return q.QueryRowContext(ctx, st.SQL, st.Args...).Scan(dest...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters
}
// Query runs the select.
func (s *Select) Query(ctx context.Context, q Querier) (*sql.Rows, error) {
st := s.Build()
return q.QueryContext(ctx, st.SQL, st.Args...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters
}
// Exists reports whether the select returns at least one row.
func (s *Select) Exists(ctx context.Context, q Querier) (bool, error) {
s.cols = []lookup.Column{s.firstCol()}
rows, err := s.Query(ctx, q)
if err != nil {
return false, err
}
defer func() { _ = rows.Close() }()
ok := rows.Next()
return ok, rows.Err()
}
func (s *Select) firstCol() lookup.Column {
if len(s.cols) > 0 {
return s.cols[0]
}
return lookup.Column{Entity: s.from, Name: lookup.FirstColumn(s.from)}
}
// Assignment is one `col = value` of an UPDATE or INSERT.
type Assignment struct {
Col lookup.Column
Val any
}
// Set builds an Assignment.
func Set(c lookup.Column, v any) Assignment { return Assignment{Col: c, Val: v} }
// Update builds an UPDATE.
type Update struct {
b *Base
e lookup.Entity
sets []Assignment
conds []Cond
}
// Update starts an UPDATE of e.
func (b *Base) Update(e lookup.Entity) *Update { return &Update{b: b, e: e} }
// Set adds assignments.
func (u *Update) Set(as ...Assignment) *Update { u.sets = append(u.sets, as...); return u }
// Where adds AND-ed conditions.
func (u *Update) Where(cs ...Cond) *Update { u.conds = append(u.conds, cs...); return u }
// Exec runs the update and returns the affected row count.
func (u *Update) Exec(ctx context.Context, q Querier) (int64, error) {
bl := &builder{b: u.b}
set := make([]string, len(u.sets))
for i, a := range u.sets {
set[i] = bl.col(a.Col) + " = " + bl.ph(a.Val)
}
sqlText := "UPDATE " + u.b.d.Quote(u.b.tableRef(u.e)) + " SET " + strings.Join(set, ", ") + bl.where(u.conds)
res, err := q.ExecContext(ctx, sqlText, bl.args...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters
if err != nil {
return 0, err
}
return res.RowsAffected()
}
// Delete builds a DELETE.
type Delete struct {
b *Base
e lookup.Entity
conds []Cond
}
// Delete starts a DELETE on e.
func (b *Base) Delete(e lookup.Entity) *Delete { return &Delete{b: b, e: e} }
// Where adds AND-ed conditions.
func (d *Delete) Where(cs ...Cond) *Delete { d.conds = append(d.conds, cs...); return d }
// Exec runs the delete and returns the affected row count.
func (d *Delete) Exec(ctx context.Context, q Querier) (int64, error) {
bl := &builder{b: d.b}
sqlText := "DELETE FROM " + d.b.d.Quote(d.b.tableRef(d.e)) + bl.where(d.conds)
res, err := q.ExecContext(ctx, sqlText, bl.args...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters
if err != nil {
return 0, err
}
return res.RowsAffected()
}
// Insert builds an INSERT.
type Insert struct {
b *Base
e lookup.Entity
sets []Assignment
}
// Insert starts an INSERT into e.
func (b *Base) Insert(e lookup.Entity) *Insert { return &Insert{b: b, e: e} }
// Set adds assignments.
func (i *Insert) Set(as ...Assignment) *Insert { i.sets = append(i.sets, as...); return i }
func (i *Insert) colsAndArgs() ([]string, []any) {
cols := make([]string, len(i.sets))
args := make([]any, len(i.sets))
for n, a := range i.sets {
cols[n] = i.b.colName(a.Col)
args[n] = i.b.arg(a.Val)
}
return cols, args
}
// Exec runs the insert.
func (i *Insert) Exec(ctx context.Context, q Querier) error {
cols, args := i.colsAndArgs()
ph := make([]string, len(cols))
qc := make([]string, len(cols))
for n, c := range cols {
qc[n] = i.b.d.Quote(c)
ph[n] = i.b.d.Placeholder(n + 1)
}
sqlText := "INSERT INTO " + i.b.d.Quote(i.b.tableRef(i.e)) + " (" + strings.Join(qc, ", ") + ") VALUES (" + strings.Join(ph, ", ") + ")"
_, err := q.ExecContext(ctx, sqlText, args...) //nolint:gosec // G701: identifiers come from the validated schema, values are bound parameters
return err
}
// ExecID runs the insert and returns the generated value of idCol, using the dialect's
// insert-returning-id strategy.
func (i *Insert) ExecID(ctx context.Context, q Querier, idCol lookup.Column) (int64, error) {
cols, args := i.colsAndArgs()
ins := i.b.d.InsertReturningID(i.b.tableRef(i.e), cols, i.b.colName(idCol))
return ins.Run(ctx, q, args...)
}
@@ -0,0 +1,75 @@
package direct
import (
"database/sql"
"testing"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect"
)
type noRun struct{}
func (noRun) Run(func(*sql.DB) error) error { return nil }
func TestBuilderRendersPerDialect(t *testing.T) {
schema := lookup.Schema{lookup.EntityUserSessions: {Schema: "auth", Name: "sessions"}}
cases := map[string]string{
"postgres": `SELECT t0."session_token" FROM "auth"."sessions" t0`,
"mysql": "SELECT t0.`session_token` FROM `auth`.`sessions` t0",
"mssql": `SELECT t0.[session_token] FROM [auth].[sessions] t0`,
}
tails := map[string]string{
"postgres": ` WHERE t0."user_id" = $1 AND t0."session_token" IN ($2, $3)`,
"mysql": " WHERE t0.`user_id` = ? AND t0.`session_token` IN (?, ?)",
"mssql": ` WHERE t0.[user_id] = @p1 AND t0.[session_token] IN (@p2, @p3)`,
}
for name, head := range cases {
d, err := dialect.Get(name)
if err != nil {
t.Fatal(err)
}
b, err := NewBase(noRun{}, d, schema)
if err != nil {
t.Fatal(err)
}
st := b.From(lookup.EntityUserSessions).Cols(lookup.SessionsToken).
Where(Eq(lookup.SessionsUserID, 7), In(lookup.SessionsToken, "a", "b")).Build()
if st.SQL != head+tails[name] || len(st.Args) != 3 {
t.Errorf("%s:\n got %s\n want %s", name, st.SQL, head+tails[name])
}
}
}
func TestBuilderBoolsGoThroughDialect(t *testing.T) {
d, _ := dialect.Get("sqlite")
b, _ := NewBase(noRun{}, d, nil)
st := b.From(lookup.EntityUsers).Cols(lookup.UsersID).Where(Eq(lookup.UsersIsActive, true)).Build()
if st.Args[0] != d.Bool(true) {
t.Fatalf("bool not converted: %#v", st.Args[0])
}
}
func TestBuilderSubselectSharesArguments(t *testing.T) {
d, _ := dialect.Get("postgres")
b, _ := NewBase(noRun{}, d, nil)
sub := b.From(lookup.EntitySecGroupMembers).Cols(lookup.GroupMembersGroupID).Where(Eq(lookup.GroupMembersUserID, 5))
st := b.From(lookup.EntitySecRowRules).Cols(lookup.RowRulesID).
Where(Eq(lookup.RowRulesTableName, "t"), InSelect(lookup.RowRulesGroupID, sub), Eq(lookup.RowRulesSchemaName, "s")).Build()
want := `SELECT t0."id" FROM "sec_row_rules" t0 WHERE t0."table_name" = $1 AND t0."group_id" IN (SELECT t1."group_id" FROM "sec_group_members" t1 WHERE t1."user_id" = $2) AND t0."schema_name" = $3`
if st.SQL != want || len(st.Args) != 3 {
t.Fatalf("got %s\nwant %s", st.SQL, want)
}
}
func TestSchemaRejectsUnsafeIdentifiers(t *testing.T) {
d, _ := dialect.Get("postgres")
bad := lookup.Schema{lookup.EntityUsers: {Name: `users"; DROP TABLE x; --`}}
if _, err := NewBase(noRun{}, d, bad); err == nil {
t.Fatal("unsafe table name accepted")
}
bad = lookup.Schema{lookup.EntityUsers: {Columns: map[string]string{"username": "a b"}}}
if _, err := NewBase(noRun{}, d, bad); err == nil {
t.Fatal("unsafe column name accepted")
}
}
@@ -0,0 +1,46 @@
package direct
import (
"database/sql"
"testing"
_ "github.com/glebarez/go-sqlite"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/ddl"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/dialect"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup/procedure"
)
func newTestDB(t *testing.T, extraDDL ...string) *sql.DB {
t.Helper()
db, err := sql.Open("sqlite", ":memory:")
if err != nil {
t.Fatal(err)
}
db.SetMaxOpenConns(1)
t.Cleanup(func() { _ = db.Close() })
ref, err := ddl.SQL("sqlite")
if err != nil {
t.Fatal(err)
}
for _, s := range append([]string{ref}, extraDDL...) {
if _, err := db.Exec(s); err != nil {
t.Fatalf("ddl: %v", err)
}
}
return db
}
func newTestBase(t *testing.T, db *sql.DB, schema lookup.Schema) *Base {
t.Helper()
d, err := dialect.Detect(db)
if err != nil {
t.Fatal(err)
}
b, err := NewBase(procedure.NewDB(db, nil, nil), d, schema)
if err != nil {
t.Fatal(err)
}
return b
}
+197
View File
@@ -0,0 +1,197 @@
package direct
import (
"context"
"database/sql"
"errors"
"fmt"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
// Keys implements lookup.KeyStore on the user keys table. scopes and meta are stored as
// JSON through the dialect (native JSON column or TEXT).
type Keys struct{ *Base }
var _ lookup.KeyStore = (*Keys)(nil)
// NewKeys creates the direct KeyStore.
func NewKeys(b *Base) *Keys { return &Keys{Base: b} }
// Create implements lookup.KeyStore.
func (k *Keys) Create(ctx context.Context, req sectypes.CreateKeyRequest, keyHash string) (*sectypes.UserKey, error) {
scopes, err := k.d.EncodeJSON(req.Scopes)
if err != nil {
return nil, fmt.Errorf("failed to marshal scopes: %w", err)
}
meta, err := k.d.EncodeJSON(req.Meta)
if err != nil {
return nil, fmt.Errorf("failed to marshal meta: %w", err)
}
now := k.Now()
var id int64
err = k.do(func(q Querier) error {
var err error
id, err = k.Insert(lookup.EntityUserKeys).Set(
Set(lookup.KeysUserID, req.UserID),
Set(lookup.KeysKeyType, string(req.KeyType)),
Set(lookup.KeysKeyHash, keyHash),
Set(lookup.KeysName, req.Name),
Set(lookup.KeysScopes, scopes),
Set(lookup.KeysMeta, meta),
Set(lookup.KeysExpiresAt, req.ExpiresAt),
Set(lookup.KeysCreatedAt, now),
Set(lookup.KeysIsActive, true),
).ExecID(ctx, q, lookup.KeysID)
return err
})
if err != nil {
return nil, fmt.Errorf("create key query failed: %w", err)
}
return &sectypes.UserKey{
ID: id,
UserID: req.UserID,
KeyType: req.KeyType,
KeyHash: keyHash,
Name: req.Name,
Scopes: req.Scopes,
Meta: req.Meta,
ExpiresAt: req.ExpiresAt,
CreatedAt: now,
IsActive: true,
}, nil
}
// keyScan holds the destinations for one key row.
type keyScan struct {
k sectypes.UserKey
keyType string
scopes, meta any
expiresAt, created, lastU time.Time
active bool
}
func (k *Keys) keyCols(withLastUsed bool) []lookup.Column {
cols := []lookup.Column{lookup.KeysID, lookup.KeysUserID, lookup.KeysKeyType, lookup.KeysName, lookup.KeysScopes,
lookup.KeysMeta, lookup.KeysExpiresAt, lookup.KeysCreatedAt, lookup.KeysIsActive}
if withLastUsed {
cols = append(cols, lookup.KeysLastUsedAt)
}
return cols
}
func (k *Keys) dest(s *keyScan, withLastUsed bool) []any {
d := []any{&s.k.ID, &s.k.UserID, &s.keyType, &s.k.Name, &s.scopes, &s.meta,
k.timeDest(&s.expiresAt), k.timeDest(&s.created), k.boolDest(&s.active)}
if withLastUsed {
d = append(d, k.timeDest(&s.lastU))
}
return d
}
func (k *Keys) finish(s *keyScan) sectypes.UserKey {
out := s.k
out.KeyType = sectypes.KeyType(s.keyType)
out.CreatedAt = s.created
out.IsActive = s.active
_ = k.d.DecodeJSON(s.scopes, &out.Scopes)
_ = k.d.DecodeJSON(s.meta, &out.Meta)
if !s.expiresAt.IsZero() {
t := s.expiresAt
out.ExpiresAt = &t
}
if !s.lastU.IsZero() {
t := s.lastU
out.LastUsedAt = &t
}
return out
}
// List implements lookup.KeyStore: active, non-expired keys; an empty keyType means all types.
func (k *Keys) List(ctx context.Context, userID int, keyType sectypes.KeyType) ([]sectypes.UserKey, error) {
keys := []sectypes.UserKey{}
conds := []Cond{
Eq(lookup.KeysUserID, userID),
Eq(lookup.KeysIsActive, true),
Or(IsNull(lookup.KeysExpiresAt), Gt(lookup.KeysExpiresAt, k.Now())),
}
if keyType != "" {
conds = append(conds, Eq(lookup.KeysKeyType, string(keyType)))
}
err := k.do(func(q Querier) error {
keys = keys[:0]
rows, err := k.From(lookup.EntityUserKeys).Cols(k.keyCols(true)...).Where(conds...).OrderBy(lookup.KeysID).Query(ctx, q)
if err != nil {
return err
}
defer func() { _ = rows.Close() }()
for rows.Next() {
var s keyScan
if err := rows.Scan(k.dest(&s, true)...); err != nil {
return err
}
keys = append(keys, k.finish(&s))
}
return rows.Err()
})
if err != nil {
return nil, fmt.Errorf("get user keys query failed: %w", err)
}
return keys, nil
}
// Delete implements lookup.KeyStore: soft-deletes the key after checking ownership and
// returns its hash.
func (k *Keys) Delete(ctx context.Context, userID int, keyID int64) (string, error) {
var keyHash string
err := k.tx(ctx, func(q Querier) error {
match := []Cond{Eq(lookup.KeysID, keyID), Eq(lookup.KeysUserID, userID), Eq(lookup.KeysIsActive, true)}
if err := k.From(lookup.EntityUserKeys).Cols(lookup.KeysKeyHash).Where(match...).QueryRow(ctx, q, &keyHash); err != nil {
return err
}
_, err := k.Update(lookup.EntityUserKeys).Set(Set(lookup.KeysIsActive, false)).Where(match...).Exec(ctx, q)
return err
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return "", errors.New("key not found or already deleted")
}
return "", fmt.Errorf("delete key query failed: %w", err)
}
return keyHash, nil
}
// Validate implements lookup.KeyStore: finds an active, non-expired key by hash (optionally of
// one type) and stamps last_used_at.
func (k *Keys) Validate(ctx context.Context, keyHash string, keyType sectypes.KeyType) (*sectypes.UserKey, error) {
conds := []Cond{
Eq(lookup.KeysKeyHash, keyHash),
Eq(lookup.KeysIsActive, true),
Or(IsNull(lookup.KeysExpiresAt), Gt(lookup.KeysExpiresAt, k.Now())),
}
if keyType != "" {
conds = append(conds, Eq(lookup.KeysKeyType, string(keyType)))
}
var s keyScan
now := k.Now()
err := k.tx(ctx, func(q Querier) error {
s = keyScan{}
if err := k.From(lookup.EntityUserKeys).Cols(k.keyCols(false)...).Where(conds...).QueryRow(ctx, q, k.dest(&s, false)...); err != nil {
return err
}
_, err := k.Update(lookup.EntityUserKeys).Set(Set(lookup.KeysLastUsedAt, now)).Where(Eq(lookup.KeysID, s.k.ID)).Exec(ctx, q)
return err
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, errors.New("invalid or expired key")
}
return nil, fmt.Errorf("validate key query failed: %w", err)
}
out := k.finish(&s)
out.KeyHash = keyHash
out.LastUsedAt = &now
return &out, nil
}
+71
View File
@@ -0,0 +1,71 @@
package direct
import (
"context"
"testing"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
func TestKeysLifecycle(t *testing.T) {
ctx := context.Background()
db := newTestDB(t)
k := NewKeys(newTestBase(t, db, nil))
_, _ = db.Exec(`INSERT INTO users (username,email,password,is_active) VALUES ('u','u@x.io','x',1)`)
exp := time.Now().Add(time.Hour)
created, err := k.Create(ctx, sectypes.CreateKeyRequest{
UserID: 1, KeyType: sectypes.KeyTypeHeaderAPI, Name: "ci",
Scopes: []string{"read", "write"}, Meta: map[string]any{"env": "prod"}, ExpiresAt: &exp,
}, sectypes.HashKey("raw1"))
if err != nil || created.ID == 0 {
t.Fatalf("%+v %v", created, err)
}
// nil scopes/meta must not store the JSON text "null"
if _, err := k.Create(ctx, sectypes.CreateKeyRequest{UserID: 1, KeyType: sectypes.KeyTypeJWTSecret, Name: "j"}, sectypes.HashKey("raw2")); err != nil {
t.Fatal(err)
}
if _, err := k.Create(ctx, sectypes.CreateKeyRequest{UserID: 1, KeyType: sectypes.KeyTypeGenericAPI, Name: "old", ExpiresAt: ptr(time.Now().Add(-time.Hour))}, sectypes.HashKey("raw3")); err != nil {
t.Fatal(err)
}
all, err := k.List(ctx, 1, "")
if err != nil || len(all) != 2 {
t.Fatalf("list all: %d %v", len(all), err)
}
one, _ := k.List(ctx, 1, sectypes.KeyTypeHeaderAPI)
if len(one) != 1 || one[0].Name != "ci" || len(one[0].Scopes) != 2 || one[0].Meta["env"] != "prod" || one[0].ExpiresAt == nil {
t.Fatalf("list typed: %+v", one)
}
if other, _ := k.List(ctx, 2, ""); len(other) != 0 {
t.Fatal("other user's keys listed")
}
got, err := k.Validate(ctx, sectypes.HashKey("raw1"), sectypes.KeyTypeHeaderAPI)
if err != nil || got.UserID != 1 || got.KeyHash != sectypes.HashKey("raw1") || got.LastUsedAt == nil {
t.Fatalf("%+v %v", got, err)
}
if _, err := k.Validate(ctx, sectypes.HashKey("raw1"), sectypes.KeyTypeGenericAPI); err == nil || err.Error() != "invalid or expired key" {
t.Fatalf("wrong type: %v", err)
}
if _, err := k.Validate(ctx, sectypes.HashKey("raw3"), ""); err == nil {
t.Fatal("expired key validated")
}
if _, err := k.Delete(ctx, 2, created.ID); err == nil || err.Error() != "key not found or already deleted" {
t.Fatalf("foreign delete: %v", err)
}
hash, err := k.Delete(ctx, 1, created.ID)
if err != nil || hash != sectypes.HashKey("raw1") {
t.Fatalf("%q %v", hash, err)
}
if _, err := k.Delete(ctx, 1, created.ID); err == nil {
t.Fatal("double delete succeeded")
}
if _, err := k.Validate(ctx, sectypes.HashKey("raw1"), ""); err == nil {
t.Fatal("deleted key validated")
}
}
func ptr[T any](v T) *T { return &v }
+380
View File
@@ -0,0 +1,380 @@
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
}
+130
View File
@@ -0,0 +1,130 @@
package direct
import (
"context"
"testing"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
func TestOAuthClientAndCodeFlow(t *testing.T) {
ctx := context.Background()
db := newTestDB(t)
b := newTestBase(t, db, nil)
c := NewOAuthClients(b)
reg, err := c.RegisterClient(ctx, &sectypes.OAuthServerClient{ClientID: "cid", RedirectURIs: []string{"https://a/cb"}, ClientName: "App"})
if err != nil || reg.TokenEndpointAuthMethod != "none" || len(reg.GrantTypes) != 1 || len(reg.AllowedScopes) != 3 {
t.Fatalf("%+v %v", reg, err)
}
got, err := c.GetClient(ctx, "cid")
if err != nil || got.ClientName != "App" || got.RedirectURIs[0] != "https://a/cb" || got.ClientSecretHash != "" {
t.Fatalf("%+v %v", got, err)
}
if _, err := c.GetClient(ctx, "nope"); err == nil || err.Error() != "client not found" {
t.Fatalf("got %v", err)
}
_, _ = db.Exec(`UPDATE oauth_clients SET is_active = 0`)
if _, err := c.GetClient(ctx, "cid"); err == nil {
t.Fatal("inactive client returned")
}
code := &sectypes.OAuthCode{Code: "c1", ClientID: "cid", RedirectURI: "https://a/cb", CodeChallenge: "ch",
SessionToken: "st", Scopes: []string{"openid"}, ExpiresAt: time.Now().Add(time.Minute)}
if err := c.SaveCode(ctx, code); err != nil {
t.Fatal(err)
}
ex, err := c.ExchangeCode(ctx, "c1")
if err != nil || ex.Code != "c1" || ex.CodeChallengeMethod != "S256" || ex.SessionToken != "st" || len(ex.Scopes) != 1 {
t.Fatalf("%+v %v", ex, err)
}
if _, err := c.ExchangeCode(ctx, "c1"); err == nil || err.Error() != "invalid or expired code" {
t.Fatalf("code reused: %v", err)
}
code.Code, code.ExpiresAt = "c2", time.Now().Add(-time.Minute)
_ = c.SaveCode(ctx, code)
if _, err := c.ExchangeCode(ctx, "c2"); err == nil {
t.Fatal("expired code exchanged")
}
}
func TestOAuthIntrospectRevoke(t *testing.T) {
ctx := context.Background()
a, db := newAuth(t, AuthOptions{})
reg := registerUser(t, a, "oli")
_, _ = db.Exec(`UPDATE users SET roles='r1,r2', user_level=3`)
c := NewOAuthClients(a.Base)
info, err := c.Introspect(ctx, reg.Token)
if err != nil || !info.Active || info.Username != "oli" || info.UserLevel != 3 || len(info.Roles) != 2 || info.Exp == 0 || info.Iat == 0 || info.Sub != "1" {
t.Fatalf("%+v %v", info, err)
}
if err := c.Revoke(ctx, reg.Token); err != nil {
t.Fatal(err)
}
if info, err := c.Introspect(ctx, reg.Token); err != nil || info.Active {
t.Fatalf("%+v %v", info, err)
}
if err := c.Revoke(ctx, "unknown"); err != nil {
t.Fatal(err)
}
}
func TestOAuthUsers(t *testing.T) {
ctx := context.Background()
db := newTestDB(t)
o := NewOAuthUsers(newTestBase(t, db, nil))
id, err := o.GetOrCreateUser(ctx, &sectypes.UserContext{UserName: "gh", Email: "g@x.io", RemoteID: "r-1", Roles: []string{"a"}}, "github")
if err != nil || id == 0 {
t.Fatalf("%d %v", id, err)
}
// second login: same user; existing remote_id/auth_provider are kept
id2, err := o.GetOrCreateUser(ctx, &sectypes.UserContext{UserName: "gh", Email: "g@x.io", RemoteID: "r-2"}, "google")
if err != nil || id2 != id {
t.Fatalf("%d %v", id2, err)
}
var remote, prov string
_ = db.QueryRow(`SELECT remote_id, auth_provider FROM users WHERE id=?`, id).Scan(&remote, &prov)
if remote != "r-1" || prov != "github" {
t.Fatalf("overwrote: %q %q", remote, prov)
}
exp := time.Now().Add(time.Hour)
s := lookup.OAuthSession{SessionToken: "s1", UserID: id, AccessToken: "a1", RefreshToken: "r1", TokenType: "Bearer", ExpiresAt: exp, Provider: "github"}
if err := o.CreateSession(ctx, s); err != nil {
t.Fatal(err)
}
s.AccessToken = "a1b" // same token: updated, not duplicated
if err := o.CreateSession(ctx, s); err != nil {
t.Fatal(err)
}
var n int
_ = db.QueryRow(`SELECT COUNT(*) FROM user_sessions`).Scan(&n)
if n != 1 {
t.Fatalf("sessions: %d", n)
}
ref, err := o.GetByRefreshToken(ctx, "r1")
if err != nil || ref.UserID != id || ref.AccessToken != "a1b" || ref.TokenType != "Bearer" || ref.Expiry.IsZero() {
t.Fatalf("%+v %v", ref, err)
}
if _, err := o.GetByRefreshToken(ctx, "zzz"); err == nil {
t.Fatal("unknown refresh token accepted")
}
if err := o.UpdateRefreshToken(ctx, id, "r1", "s2", "a2", "r2", exp); err != nil {
t.Fatal(err)
}
if err := o.UpdateRefreshToken(ctx, id, "r1", "s3", "a3", "r3", exp); err == nil || err.Error() != "session not found" {
t.Fatalf("got %v", err)
}
u, err := o.GetUser(ctx, id)
if err != nil || u.UserName != "gh" || u.UserID != id {
t.Fatalf("%+v %v", u, err)
}
if _, err := o.GetUser(ctx, 999); err == nil || err.Error() != "user not found" {
t.Fatalf("got %v", err)
}
}
+294
View File
@@ -0,0 +1,294 @@
package direct
import (
"context"
"database/sql"
"encoding/base64"
"errors"
"fmt"
"time"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
// Passkey implements lookup.PasskeyStore. credential_id, public_key and aaguid are base64
// TEXT (not native bytea) and transports is JSON, so one schema works on every dialect.
type Passkey struct{ *Base }
var _ lookup.PasskeyStore = (*Passkey)(nil)
// NewPasskey creates the direct PasskeyStore.
func NewPasskey(b *Base) *Passkey { return &Passkey{Base: b} }
// Store implements lookup.PasskeyStore.
func (p *Passkey) Store(ctx context.Context, rec lookup.PasskeyCredentialRecord) (int64, error) {
transports, err := p.d.EncodeJSON(rec.Transports)
if err != nil {
return 0, fmt.Errorf("failed to marshal transports: %w", err)
}
var id int64
err = p.tx(ctx, func(q Querier) error {
exists, err := p.From(lookup.EntityUserPasskeyCredentials).Cols(lookup.PasskeyID).
Where(Eq(lookup.PasskeyCredentialID, rec.CredentialID)).Exists(ctx, q)
if err != nil {
return err
}
if exists {
return fmt.Errorf("credential already exists")
}
userExists, err := p.From(lookup.EntityUsers).Cols(lookup.UsersID).Where(Eq(lookup.UsersID, rec.UserID)).Exists(ctx, q)
if err != nil {
return err
}
if !userExists {
return fmt.Errorf("user not found")
}
now := p.Now()
id, err = p.Insert(lookup.EntityUserPasskeyCredentials).Set(
Set(lookup.PasskeyUserID, rec.UserID),
Set(lookup.PasskeyCredentialID, rec.CredentialID),
Set(lookup.PasskeyPublicKey, rec.PublicKey),
Set(lookup.PasskeyAttestationType, rec.AttestationType),
Set(lookup.PasskeyAAGUID, ""),
Set(lookup.PasskeySignCount, int64(rec.SignCount)),
Set(lookup.PasskeyTransports, transports),
Set(lookup.PasskeyBackupEligible, rec.BackupEligible),
Set(lookup.PasskeyBackupState, rec.BackupState),
Set(lookup.PasskeyName, rec.Name),
Set(lookup.PasskeyCreatedAt, now),
Set(lookup.PasskeyLastUsedAt, now),
).ExecID(ctx, q, lookup.PasskeyID)
return err
})
if err != nil {
return 0, err
}
return id, nil
}
// Get implements lookup.PasskeyStore.
func (p *Passkey) Get(ctx context.Context, credentialID string) (int, uint32, error) {
var userID int
var count int64
err := p.do(func(q Querier) error {
return p.From(lookup.EntityUserPasskeyCredentials).Cols(lookup.PasskeyUserID, lookup.PasskeySignCount).
Where(Eq(lookup.PasskeyCredentialID, credentialID)).QueryRow(ctx, q, &userID, &count)
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return 0, 0, fmt.Errorf("credential not found")
}
return 0, 0, fmt.Errorf("failed to get credential: %w", err)
}
return userID, uint32(count), nil //nolint:gosec // sign counters are stored from uint32 values
}
// UpdateCounter implements lookup.PasskeyStore. A counter that did not advance flags the
// credential as possibly cloned and leaves the stored counter unchanged.
func (p *Passkey) UpdateCounter(ctx context.Context, credentialID string, newCounter uint32) (bool, error) {
var clone bool
err := p.tx(ctx, func(q Querier) error {
match := Eq(lookup.PasskeyCredentialID, credentialID)
var old int64
if err := p.From(lookup.EntityUserPasskeyCredentials).Cols(lookup.PasskeySignCount).Where(match).QueryRow(ctx, q, &old); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return fmt.Errorf("credential not found")
}
return err
}
if int64(newCounter) <= old {
clone = true
_, err := p.Update(lookup.EntityUserPasskeyCredentials).Set(Set(lookup.PasskeyCloneWarning, true)).Where(match).Exec(ctx, q)
return err
}
_, err := p.Update(lookup.EntityUserPasskeyCredentials).
Set(Set(lookup.PasskeySignCount, int64(newCounter)), Set(lookup.PasskeyLastUsedAt, p.Now())).
Where(match).Exec(ctx, q)
return err
})
return clone, err
}
// List implements lookup.PasskeyStore, newest first.
func (p *Passkey) List(ctx context.Context, userID int) ([]sectypes.PasskeyCredential, error) {
var out []sectypes.PasskeyCredential
err := p.do(func(q Querier) error {
out = make([]sectypes.PasskeyCredential, 0)
rows, err := p.From(lookup.EntityUserPasskeyCredentials).
Cols(lookup.PasskeyID, lookup.PasskeyUserID, lookup.PasskeyCredentialID, lookup.PasskeyPublicKey,
lookup.PasskeyAttestationType, lookup.PasskeyAAGUID, lookup.PasskeySignCount, lookup.PasskeyCloneWarning,
lookup.PasskeyTransports, lookup.PasskeyBackupEligible, lookup.PasskeyBackupState, lookup.PasskeyName,
lookup.PasskeyCreatedAt, lookup.PasskeyLastUsedAt).
Where(Eq(lookup.PasskeyUserID, userID)).OrderBy(lookup.PasskeyCreatedAt).Query(ctx, q)
if err != nil {
return err
}
defer func() { _ = rows.Close() }()
for rows.Next() {
var id, uid int
var credB64, pubB64 string
var attestation, aaguidB64, name sql.NullString
var count sql.NullInt64
var clone, eligible, state bool
var transports any
var created, last time.Time
if err := rows.Scan(&id, &uid, &credB64, &pubB64, &attestation, &aaguidB64, &count, p.boolDest(&clone),
&transports, p.boolDest(&eligible), p.boolDest(&state), &name, p.timeDest(&created), p.timeDest(&last)); err != nil {
return err
}
credID, err := base64.StdEncoding.DecodeString(credB64)
if err != nil {
continue
}
pub, err := base64.StdEncoding.DecodeString(pubB64)
if err != nil {
continue
}
aaguid, _ := base64.StdEncoding.DecodeString(aaguidB64.String)
c := sectypes.PasskeyCredential{
ID: fmt.Sprintf("%d", id),
UserID: uid,
CredentialID: credID,
PublicKey: pub,
AttestationType: attestation.String,
AAGUID: aaguid,
SignCount: uint32(count.Int64), //nolint:gosec // stored from uint32 values
CloneWarning: clone,
BackupEligible: eligible,
BackupState: state,
Name: name.String,
CreatedAt: created,
LastUsedAt: last,
}
_ = p.d.DecodeJSON(transports, &c.Transports)
out = append(out, c)
}
return rows.Err()
})
if err != nil {
return nil, fmt.Errorf("failed to get credentials: %w", err)
}
// newest first
for i, j := 0, len(out)-1; i < j; i, j = i+1, j-1 {
out[i], out[j] = out[j], out[i]
}
return out, nil
}
// Delete implements lookup.PasskeyStore.
func (p *Passkey) Delete(ctx context.Context, userID int, credentialID string) error {
var rows int64
err := p.do(func(q Querier) error {
var err error
rows, err = p.Base.Delete(lookup.EntityUserPasskeyCredentials).
Where(Eq(lookup.PasskeyUserID, userID), Eq(lookup.PasskeyCredentialID, credentialID)).Exec(ctx, q)
return err
})
if err != nil {
return err
}
if rows == 0 {
return fmt.Errorf("credential not found")
}
return nil
}
// Rename implements lookup.PasskeyStore.
func (p *Passkey) Rename(ctx context.Context, userID int, credentialID, name string) error {
var rows int64
err := p.do(func(q Querier) error {
var err error
rows, err = p.Update(lookup.EntityUserPasskeyCredentials).Set(Set(lookup.PasskeyName, name)).
Where(Eq(lookup.PasskeyUserID, userID), Eq(lookup.PasskeyCredentialID, credentialID)).Exec(ctx, q)
return err
})
if err != nil {
return err
}
if rows == 0 {
return fmt.Errorf("credential not found")
}
return nil
}
// ByUsername implements lookup.PasskeyStore.
func (p *Passkey) ByUsername(ctx context.Context, username string) (int, []lookup.PasskeyCredentialRef, error) {
var userID int
var creds []lookup.PasskeyCredentialRef
err := p.do(func(q Querier) error {
creds = make([]lookup.PasskeyCredentialRef, 0)
if err := p.From(lookup.EntityUsers).Cols(lookup.UsersID).
Where(Eq(lookup.UsersUsername, username), Eq(lookup.UsersIsActive, true)).QueryRow(ctx, q, &userID); err != nil {
return err
}
rows, err := p.From(lookup.EntityUserPasskeyCredentials).Cols(lookup.PasskeyCredentialID, lookup.PasskeyTransports).
Where(Eq(lookup.PasskeyUserID, userID)).Query(ctx, q)
if err != nil {
return err
}
defer func() { _ = rows.Close() }()
for rows.Next() {
var ref lookup.PasskeyCredentialRef
var transports any
if err := rows.Scan(&ref.CredentialID, &transports); err != nil {
return err
}
_ = p.d.DecodeJSON(transports, &ref.Transports)
creds = append(creds, ref)
}
return rows.Err()
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return 0, nil, fmt.Errorf("user not found")
}
return 0, nil, fmt.Errorf("failed to get credentials: %w", err)
}
return userID, creds, nil
}
// Login implements lookup.PasskeyStore: it creates the session for a user whose passkey
// assertion was already verified.
func (p *Passkey) Login(ctx context.Context, userID int, claims map[string]any) (*sectypes.LoginResponse, error) {
var u userRow
err := p.do(func(q Querier) error {
return p.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("passkey login query failed: %w", err)
}
u.id = userID
token, err := GenerateSessionToken()
if err != nil {
return nil, fmt.Errorf("failed to generate session token: %w", err)
}
now := p.Now()
ip, ua := claimStrings(claims)
err = p.tx(ctx, func(q Querier) error {
if err := p.Insert(lookup.EntityUserSessions).Set(
Set(lookup.SessionsToken, token),
Set(lookup.SessionsUserID, userID),
Set(lookup.SessionsExpiresAt, now.Add(sessionLifetime)),
Set(lookup.SessionsIPAddress, ip),
Set(lookup.SessionsUserAgent, ua),
Set(lookup.SessionsLastActivityAt, now),
Set(lookup.SessionsCreatedAt, now),
).Exec(ctx, q); err != nil {
return err
}
_, err := p.Update(lookup.EntityUsers).Set(Set(lookup.UsersLastLoginAt, now)).Where(Eq(lookup.UsersID, userID)).Exec(ctx, q)
return err
})
if err != nil {
return nil, fmt.Errorf("passkey login query failed: %w", err)
}
return &sectypes.LoginResponse{Token: token, User: u.context(token), ExpiresIn: int64(sessionLifetime.Seconds())}, nil
}
@@ -0,0 +1,162 @@
package direct
import (
"context"
"encoding/base64"
"testing"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
)
func b64(s string) string { return base64.StdEncoding.EncodeToString([]byte(s)) }
func TestPasskeyLifecycle(t *testing.T) {
ctx := context.Background()
a, _ := newAuth(t, AuthOptions{})
reg := registerUser(t, a, "pat")
uid := reg.User.UserID
p := NewPasskey(a.Base)
rec := lookup.PasskeyCredentialRecord{UserID: uid, CredentialID: b64("cred1"), PublicKey: b64("pk"), AttestationType: "none",
Transports: []string{"usb", "nfc"}, Name: "Key 1"}
id, err := p.Store(ctx, rec)
if err != nil || id == 0 {
t.Fatalf("%d %v", id, err)
}
if _, err := p.Store(ctx, rec); err == nil || err.Error() != "credential already exists" {
t.Fatalf("dup: %v", err)
}
rec.CredentialID, rec.UserID = b64("cred2"), 999
if _, err := p.Store(ctx, rec); err == nil || err.Error() != "user not found" {
t.Fatalf("no user: %v", err)
}
rec.UserID, rec.Name = uid, "Key 2"
if _, err := p.Store(ctx, rec); err != nil {
t.Fatal(err)
}
owner, count, err := p.Get(ctx, b64("cred1"))
if err != nil || owner != uid || count != 0 {
t.Fatalf("%d %d %v", owner, count, err)
}
if _, _, err := p.Get(ctx, b64("zzz")); err == nil || err.Error() != "credential not found" {
t.Fatalf("got %v", err)
}
if clone, err := p.UpdateCounter(ctx, b64("cred1"), 5); err != nil || clone {
t.Fatalf("%v %v", clone, err)
}
if clone, err := p.UpdateCounter(ctx, b64("cred1"), 5); err != nil || !clone {
t.Fatalf("replayed counter must flag clone: %v %v", clone, err)
}
if _, count, _ = p.Get(ctx, b64("cred1")); count != 5 {
t.Fatalf("counter changed on clone: %d", count)
}
if _, err := p.UpdateCounter(ctx, b64("missing"), 1); err == nil {
t.Fatal("expected not found")
}
list, err := p.List(ctx, uid)
if err != nil || len(list) != 2 {
t.Fatalf("%d %v", len(list), err)
}
for _, c := range list {
if string(c.CredentialID) == "cred1" {
if !c.CloneWarning || c.SignCount != 5 || len(c.Transports) != 2 || c.Name != "Key 1" {
t.Fatalf("%+v", c)
}
}
}
if err := p.Rename(ctx, uid, b64("cred1"), "Renamed"); err != nil {
t.Fatal(err)
}
if err := p.Rename(ctx, uid+1, b64("cred1"), "x"); err == nil {
t.Fatal("renamed another user's credential")
}
gotID, refs, err := p.ByUsername(ctx, "pat")
if err != nil || gotID != uid || len(refs) != 2 {
t.Fatalf("%d %+v %v", gotID, refs, err)
}
if _, _, err := p.ByUsername(ctx, "ghost"); err == nil || err.Error() != "user not found" {
t.Fatalf("got %v", err)
}
resp, err := p.Login(ctx, uid, map[string]any{"ip_address": "1.1.1.1"})
if err != nil || resp.User.UserName != "pat" || resp.ExpiresIn != 86400 {
t.Fatalf("%+v %v", resp, err)
}
if _, err := a.Session(ctx, resp.Token, ""); err != nil {
t.Fatal(err)
}
if err := p.Delete(ctx, uid+1, b64("cred1")); err == nil {
t.Fatal("deleted another user's credential")
}
if err := p.Delete(ctx, uid, b64("cred1")); err != nil {
t.Fatal(err)
}
if err := p.Delete(ctx, uid, b64("cred1")); err == nil || err.Error() != "credential not found" {
t.Fatalf("got %v", err)
}
}
func TestTOTPLifecycle(t *testing.T) {
ctx := context.Background()
a, _ := newAuth(t, AuthOptions{})
uid := registerUser(t, a, "tom").User.UserID
s := NewTOTP(a.Base)
if on, err := s.Status(ctx, uid); err != nil || on {
t.Fatalf("%v %v", on, err)
}
if _, err := s.Secret(ctx, uid); err == nil || err.Error() != "TOTP not enabled for user" {
t.Fatalf("got %v", err)
}
if err := s.RegenerateBackupCodes(ctx, uid, []string{"h"}); err == nil {
t.Fatal("regenerate without 2FA")
}
if err := s.Enable(ctx, 999, "S", nil); err == nil || err.Error() != "user not found" {
t.Fatalf("got %v", err)
}
if err := s.Enable(ctx, uid, "SECRET", []string{"h1", "h2"}); err != nil {
t.Fatal(err)
}
if on, _ := s.Status(ctx, uid); !on {
t.Fatal("not enabled")
}
if sec, err := s.Secret(ctx, uid); err != nil || sec != "SECRET" {
t.Fatalf("%q %v", sec, err)
}
if ok, err := s.ValidateBackupCode(ctx, uid, "h1"); err != nil || !ok {
t.Fatalf("%v %v", ok, err)
}
if _, err := s.ValidateBackupCode(ctx, uid, "h1"); err == nil || err.Error() != "backup code already used" {
t.Fatalf("reuse: %v", err)
}
if ok, err := s.ValidateBackupCode(ctx, uid, "nope"); err != nil || ok {
t.Fatalf("%v %v", ok, err)
}
if err := s.RegenerateBackupCodes(ctx, uid, []string{"n1"}); err != nil {
t.Fatal(err)
}
if ok, _ := s.ValidateBackupCode(ctx, uid, "h2"); ok {
t.Fatal("old code survived regenerate")
}
if ok, _ := s.ValidateBackupCode(ctx, uid, "n1"); !ok {
t.Fatal("new code rejected")
}
if err := s.Disable(ctx, uid); err != nil {
t.Fatal(err)
}
if on, _ := s.Status(ctx, uid); on {
t.Fatal("still enabled")
}
if ok, _ := s.ValidateBackupCode(ctx, uid, "n1"); ok {
t.Fatal("codes survived disable")
}
}
+67
View File
@@ -0,0 +1,67 @@
package direct
import (
"crypto/subtle"
"errors"
"strings"
"sync"
"golang.org/x/crypto/bcrypt"
)
// bcrypt only considers the first 72 bytes of input; longer passwords are
// rejected rather than silently truncated.
const maxPasswordBytes = 72
var errPasswordTooLong = errors.New("password must be at most 72 bytes")
// HashPassword returns the bcrypt hash of password.
func HashPassword(password string) (string, error) {
if len(password) > maxPasswordBytes {
return "", errPasswordTooLong
}
h, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
if err != nil {
return "", err
}
return string(h), nil
}
func isBcryptHash(s string) bool {
return strings.HasPrefix(s, "$2a$") || strings.HasPrefix(s, "$2b$") || strings.HasPrefix(s, "$2y$")
}
// VerifyPassword checks supplied against the stored value. A stored bcrypt hash is compared
// with bcrypt. A legacy cleartext value (written before hashing was implemented) is compared
// in constant time and, on a match, needsRehash is true so the caller can upgrade the row.
// An empty stored value (e.g. an OAuth2-only user) never matches.
func VerifyPassword(stored, supplied string) (ok, needsRehash bool) {
if stored == "" || supplied == "" || len(supplied) > maxPasswordBytes {
return false, false
}
if isBcryptHash(stored) {
return bcrypt.CompareHashAndPassword([]byte(stored), []byte(supplied)) == nil, false
}
if subtle.ConstantTimeCompare([]byte(stored), []byte(supplied)) == 1 {
return true, true
}
return false, false
}
var (
dummyHashOnce sync.Once
dummyHash string
)
// BurnPasswordCheck spends roughly one bcrypt comparison so an unknown username costs about
// the same as a wrong password.
func BurnPasswordCheck(supplied string) {
dummyHashOnce.Do(func() {
h, _ := bcrypt.GenerateFromPassword([]byte("resolvespec-dummy"), bcrypt.DefaultCost)
dummyHash = string(h)
})
if len(supplied) > maxPasswordBytes {
supplied = supplied[:maxPasswordBytes]
}
_ = bcrypt.CompareHashAndPassword([]byte(dummyHash), []byte(supplied))
}
@@ -0,0 +1,19 @@
package direct
import "testing"
func TestVerifyPasswordEdgeCases(t *testing.T) {
h, _ := HashPassword("pw")
if ok, _ := VerifyPassword(h, "pw"); !ok {
t.Error("bcrypt match failed")
}
if ok, _ := VerifyPassword("", "pw"); ok {
t.Error("empty stored must not match")
}
if ok, _ := VerifyPassword("pw", ""); ok {
t.Error("empty supplied must not match")
}
if _, err := HashPassword(string(make([]byte, 73))); err == nil {
t.Error("73-byte password must be rejected")
}
}
+200
View File
@@ -0,0 +1,200 @@
package direct
import (
"context"
"database/sql"
"fmt"
"strconv"
"strings"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
)
// PolicyOptions tunes Policy.
type PolicyOptions struct {
// NoGroups skips the group membership table: only rules addressed to the user directly
// apply. Use it when the sec_group_members table is not deployed.
NoGroups bool
}
// Policy implements lookup.PolicyStore on the rule tables.
//
// Applicable rules are the active rules whose user_id is the caller plus the rules of every
// group the caller belongs to; schema and table match case-insensitively and exactly (never a
// prefix). Column security returns the union of the matching rules. Row security: any
// applicable has_block rule wins, otherwise the templates are combined with AND, each in
// parentheses. No rule is an empty result; failures are errors so callers fail closed.
type Policy struct {
*Base
opts PolicyOptions
}
var _ lookup.PolicyStore = (*Policy)(nil)
// NewPolicy creates the direct PolicyStore.
func NewPolicy(b *Base, opts PolicyOptions) *Policy { return &Policy{Base: b, opts: opts} }
// applicable restricts a rule query to the rules that apply to userID.
func (p *Policy) applicable(userCol, groupCol lookup.Column, userID int64) Cond {
if p.opts.NoGroups {
return Eq(userCol, userID)
}
members := p.From(lookup.EntitySecGroupMembers).Cols(lookup.GroupMembersGroupID).Where(Eq(lookup.GroupMembersUserID, userID))
return Or(Eq(userCol, userID), InSelect(groupCol, members))
}
// ColumnSecurity implements lookup.PolicyStore.
func (p *Policy) ColumnSecurity(ctx context.Context, userID int, schema, table string) ([]sectypes.ColumnSecurity, error) {
var rules []sectypes.ColumnSecurity
err := p.do(func(q Querier) error {
rules = nil
rows, err := p.From(lookup.EntitySecColumnRules).
Cols(lookup.ColRulesID, lookup.ColRulesColumnPath, lookup.ColRulesAccessType, lookup.ColRulesMaskStart,
lookup.ColRulesMaskEnd, lookup.ColRulesMaskInvert, lookup.ColRulesMaskChar, lookup.ColRulesExtraFilters).
Where(
Eq(lookup.ColRulesIsActive, true),
EqFold(lookup.ColRulesSchemaName, schema),
EqFold(lookup.ColRulesTableName, table),
p.applicable(lookup.ColRulesUserID, lookup.ColRulesGroupID, int64(userID)),
).OrderBy(lookup.ColRulesID).Query(ctx, q)
if err != nil {
return err
}
defer func() { _ = rows.Close() }()
for rows.Next() {
var id int
var path, access string
var start, end sql.NullInt64
var invert sql.NullBool
var maskChar sql.NullString
var extra any
var inv any
if err := rows.Scan(&id, &path, &access, &start, &end, &inv, &maskChar, &extra); err != nil {
return err
}
if inv != nil {
b, err := p.d.ScanBool(inv)
if err != nil {
return err
}
invert = sql.NullBool{Bool: b, Valid: true}
}
rule := sectypes.ColumnSecurity{
ID: id,
Schema: schema,
Tablename: table,
Path: strings.Split(path, "."),
Accesstype: access,
UserID: userID,
MaskStart: int(start.Int64),
MaskEnd: int(end.Int64),
MaskInvert: invert.Bool,
MaskChar: "*",
Control: schema + "." + table + "." + path,
}
if maskChar.Valid && maskChar.String != "" {
rule.MaskChar = maskChar.String
}
if err := p.d.DecodeJSON(extra, &rule.ExtraFilters); err != nil {
return err
}
rules = append(rules, rule)
}
return rows.Err()
})
if err != nil {
return nil, fmt.Errorf("failed to load column security: %w", err)
}
return rules, nil
}
// numericUser reduces a user reference to the integer the rule tables key on. Structured
// values are rejected, and so are non-numeric strings: a reference that cannot be matched
// must fail closed rather than silently load no rules.
func numericUser(ref any) (int64, error) {
switch v := ref.(type) {
case *sectypes.UserContext:
if v == nil {
return 0, fmt.Errorf("row security: nil user context")
}
return int64(v.UserID), nil
case sectypes.UserContext:
return int64(v.UserID), nil
case int:
return int64(v), nil
case int8:
return int64(v), nil
case int16:
return int64(v), nil
case int32:
return int64(v), nil
case int64:
return v, nil
case uint:
return int64(v), nil //nolint:gosec // user ids fit int64
case uint8:
return int64(v), nil
case uint16:
return int64(v), nil
case uint32:
return int64(v), nil
case uint64:
return int64(v), nil //nolint:gosec // user ids fit int64
case string:
n, err := strconv.ParseInt(strings.TrimSpace(v), 10, 64)
if err != nil {
return 0, fmt.Errorf("row security: user reference %q is not a numeric id", v)
}
return n, nil
case nil:
return 0, fmt.Errorf("row security: no user reference")
}
return 0, fmt.Errorf("row security: unsupported user reference type %T", ref)
}
// RowSecurity implements lookup.PolicyStore.
func (p *Policy) RowSecurity(ctx context.Context, userRef any, schema, table string) (sectypes.RowSecurity, error) {
uid, err := numericUser(userRef)
if err != nil {
return sectypes.RowSecurity{}, err
}
var templates []string
block := false
err = p.do(func(q Querier) error {
templates, block = nil, false
rows, err := p.From(lookup.EntitySecRowRules).Cols(lookup.RowRulesTemplate, lookup.RowRulesHasBlock).
Where(
Eq(lookup.RowRulesIsActive, true),
EqFold(lookup.RowRulesSchemaName, schema),
EqFold(lookup.RowRulesTableName, table),
p.applicable(lookup.RowRulesUserID, lookup.RowRulesGroupID, uid),
).OrderBy(lookup.RowRulesID).Query(ctx, q)
if err != nil {
return err
}
defer func() { _ = rows.Close() }()
for rows.Next() {
var tpl sql.NullString
var hb bool
if err := rows.Scan(&tpl, p.boolDest(&hb)); err != nil {
return err
}
if hb {
block = true
}
if t := strings.TrimSpace(tpl.String); t != "" {
templates = append(templates, "("+t+")")
}
}
return rows.Err()
})
if err != nil {
return sectypes.RowSecurity{}, fmt.Errorf("failed to load row security: %w", err)
}
rs := sectypes.RowSecurity{Schema: schema, Tablename: table, UserID: userRef, HasBlock: block}
if !block {
rs.Template = strings.Join(templates, " AND ")
}
return rs, nil
}
+93
View File
@@ -0,0 +1,93 @@
package direct
import (
"context"
"testing"
)
func TestPolicyColumnAndRowSecurity(t *testing.T) {
ctx := context.Background()
db := newTestDB(t)
p := NewPolicy(newTestBase(t, db, nil), PolicyOptions{})
exec := func(q string, args ...any) {
t.Helper()
if _, err := db.Exec(q, args...); err != nil {
t.Fatal(err)
}
}
exec(`INSERT INTO sec_group_members (group_id, user_id) VALUES (10, 1), (10, 3)`)
// column rules: user 1 direct, group 10 (user 1 and 3), another user, inactive, other table, prefix table
exec(`INSERT INTO sec_column_rules (user_id, group_id, schema_name, table_name, column_path, access_type, mask_start, mask_end, mask_invert, mask_char, extra_filters, is_active)
VALUES (1, NULL, 'Public', 'Users', 'email', 'mask', 2, 1, 1, '#', '{"k":"v"}', 1),
(NULL, 10, 'public', 'users', 'profile.ssn', 'hide', NULL, NULL, NULL, NULL, NULL, 1),
(2, NULL, 'public', 'users', 'other', 'hide', 0, 0, 0, '*', NULL, 1),
(1, NULL, 'public', 'users', 'off', 'hide', 0, 0, 0, '*', NULL, 0),
(1, NULL, 'public', 'orders', 'x', 'hide', 0, 0, 0, '*', NULL, 1),
(1, NULL, 'public', 'users_archive', 'y', 'hide', 0, 0, 0, '*', NULL, 1)`)
rules, err := p.ColumnSecurity(ctx, 1, "public", "users")
if err != nil || len(rules) != 2 {
t.Fatalf("%d %v %+v", len(rules), err, rules)
}
m := rules[0]
if m.Accesstype != "mask" || m.MaskStart != 2 || m.MaskEnd != 1 || !m.MaskInvert || m.MaskChar != "#" ||
m.ExtraFilters["k"] != "v" || len(m.Path) != 1 || m.Path[0] != "email" || m.UserID != 1 {
t.Fatalf("%+v", m)
}
h := rules[1]
if len(h.Path) != 2 || h.Path[1] != "ssn" || h.MaskChar != "*" || h.Accesstype != "hide" {
t.Fatalf("%+v", h)
}
if r, err := p.ColumnSecurity(ctx, 3, "public", "users"); err != nil || len(r) != 1 {
t.Fatalf("group member: %d %v", len(r), err)
}
if r, err := p.ColumnSecurity(ctx, 99, "public", "users"); err != nil || len(r) != 0 {
t.Fatalf("no rules must be empty: %d %v", len(r), err)
}
// row rules
exec(`INSERT INTO sec_row_rules (user_id, group_id, schema_name, table_name, template, has_block, is_active) VALUES
(1, NULL, 'public', 'orders', 'owner_id = {UserID}', 0, 1),
(NULL, 10, 'public', 'orders', 'region = 1', 0, 1),
(NULL, 10, 'public', 'orders', 'ignored', 0, 0),
(3, NULL, 'public', 'secret', NULL, 1, 1),
(NULL, 10, 'public', 'secret', 'x = 1', 0, 1)`)
rs, err := p.RowSecurity(ctx, 1, "public", "orders")
if err != nil || rs.Template != "(owner_id = {UserID}) AND (region = 1)" || rs.HasBlock {
t.Fatalf("%+v %v", rs, err)
}
if rs, err := p.RowSecurity(ctx, "3", "PUBLIC", "Secret"); err != nil || !rs.HasBlock || rs.Template != "" {
t.Fatalf("block must win: %+v %v", rs, err)
}
if rs, err := p.RowSecurity(ctx, 99, "public", "orders"); err != nil || rs.Template != "" || rs.HasBlock {
t.Fatalf("%+v %v", rs, err)
}
for _, bad := range []any{nil, "abc", []int{1}, 1.5} {
if _, err := p.RowSecurity(ctx, bad, "public", "orders"); err == nil {
t.Fatalf("user ref %#v accepted", bad)
}
}
}
func TestPolicyNoGroups(t *testing.T) {
ctx := context.Background()
db := newTestDB(t)
p := NewPolicy(newTestBase(t, db, nil), PolicyOptions{NoGroups: true})
_, _ = db.Exec(`DROP TABLE sec_group_members`)
_, _ = db.Exec(`INSERT INTO sec_column_rules (user_id, group_id, schema_name, table_name, column_path, access_type, is_active) VALUES
(1, NULL, 's', 't', 'a', 'hide', 1), (NULL, 5, 's', 't', 'b', 'hide', 1)`)
rules, err := p.ColumnSecurity(ctx, 1, "s", "t")
if err != nil || len(rules) != 1 || rules[0].Path[0] != "a" {
t.Fatalf("%+v %v", rules, err)
}
}
func TestPolicyFailsClosedOnMissingTable(t *testing.T) {
db := newTestDB(t)
p := NewPolicy(newTestBase(t, db, nil), PolicyOptions{})
_, _ = db.Exec(`DROP TABLE sec_row_rules`)
if _, err := p.RowSecurity(context.Background(), 1, "s", "t"); err == nil {
t.Fatal("expected error for missing table")
}
}
+159
View File
@@ -0,0 +1,159 @@
package direct
import (
"context"
"database/sql"
"errors"
"fmt"
"github.com/bitechdev/ResolveSpec/pkg/security/lookup"
)
// TOTP implements lookup.TOTPStore: the secret and enabled flag live on the users table,
// backup code hashes in their own table.
type TOTP struct{ *Base }
var _ lookup.TOTPStore = (*TOTP)(nil)
// NewTOTP creates the direct TOTPStore.
func NewTOTP(b *Base) *TOTP { return &TOTP{Base: b} }
func (t *TOTP) replaceBackupCodes(ctx context.Context, q Querier, userID int, hashed []string) error {
if _, err := t.Delete(lookup.EntityUserTOTPBackupCodes).Where(Eq(lookup.BackupCodesUserID, userID)).Exec(ctx, q); err != nil {
return err
}
now := t.Now()
for _, h := range hashed {
if err := t.Insert(lookup.EntityUserTOTPBackupCodes).Set(
Set(lookup.BackupCodesUserID, userID),
Set(lookup.BackupCodesCodeHash, h),
Set(lookup.BackupCodesUsed, false),
Set(lookup.BackupCodesCreatedAt, now),
).Exec(ctx, q); err != nil {
return err
}
}
return nil
}
// Enable implements lookup.TOTPStore.
func (t *TOTP) Enable(ctx context.Context, userID int, secret string, hashedCodes []string) error {
return t.tx(ctx, func(q Querier) error {
n, err := t.Update(lookup.EntityUsers).Set(
Set(lookup.UsersTOTPSecret, secret),
Set(lookup.UsersTOTPEnabled, true),
Set(lookup.UsersTOTPEnabledAt, t.Now()),
).Where(Eq(lookup.UsersID, userID)).Exec(ctx, q)
if err != nil {
return err
}
if n == 0 {
return fmt.Errorf("user not found")
}
return t.replaceBackupCodes(ctx, q, userID, hashedCodes)
})
}
// Disable implements lookup.TOTPStore.
func (t *TOTP) Disable(ctx context.Context, userID int) error {
return t.tx(ctx, func(q Querier) error {
n, err := t.Update(lookup.EntityUsers).Set(Set(lookup.UsersTOTPSecret, nil), Set(lookup.UsersTOTPEnabled, false)).
Where(Eq(lookup.UsersID, userID)).Exec(ctx, q)
if err != nil {
return err
}
if n == 0 {
return fmt.Errorf("user not found")
}
_, err = t.Delete(lookup.EntityUserTOTPBackupCodes).Where(Eq(lookup.BackupCodesUserID, userID)).Exec(ctx, q)
return err
})
}
// Status implements lookup.TOTPStore.
func (t *TOTP) Status(ctx context.Context, userID int) (bool, error) {
var enabled bool
err := t.do(func(q Querier) error {
return t.From(lookup.EntityUsers).Cols(lookup.UsersTOTPEnabled).Where(Eq(lookup.UsersID, userID)).
QueryRow(ctx, q, t.boolDest(&enabled))
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return false, fmt.Errorf("user not found")
}
return false, fmt.Errorf("get 2FA status query failed: %w", err)
}
return enabled, nil
}
// Secret implements lookup.TOTPStore.
func (t *TOTP) Secret(ctx context.Context, userID int) (string, error) {
var secret sql.NullString
var enabled bool
err := t.do(func(q Querier) error {
return t.From(lookup.EntityUsers).Cols(lookup.UsersTOTPSecret, lookup.UsersTOTPEnabled).Where(Eq(lookup.UsersID, userID)).
QueryRow(ctx, q, &secret, t.boolDest(&enabled))
})
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return "", fmt.Errorf("user not found")
}
return "", fmt.Errorf("get 2FA secret query failed: %w", err)
}
if !enabled {
return "", fmt.Errorf("TOTP not enabled for user")
}
return secret.String, nil
}
// RegenerateBackupCodes implements lookup.TOTPStore.
func (t *TOTP) RegenerateBackupCodes(ctx context.Context, userID int, hashedCodes []string) error {
return t.tx(ctx, func(q Querier) error {
ok, err := t.From(lookup.EntityUsers).Cols(lookup.UsersID).
Where(Eq(lookup.UsersID, userID), Eq(lookup.UsersTOTPEnabled, true)).Exists(ctx, q)
if err != nil {
return err
}
if !ok {
return fmt.Errorf("user not found or TOTP not enabled")
}
return t.replaceBackupCodes(ctx, q, userID, hashedCodes)
})
}
// ValidateBackupCode implements lookup.TOTPStore. An unknown code is (false, nil); a used
// code is an error. The code is consumed with a conditional update so it cannot be spent twice.
func (t *TOTP) ValidateBackupCode(ctx context.Context, userID int, codeHash string) (bool, error) {
var valid bool
err := t.tx(ctx, func(q Querier) error {
var id int64
var used bool
err := t.From(lookup.EntityUserTOTPBackupCodes).Cols(lookup.BackupCodesID, lookup.BackupCodesUsed).
Where(Eq(lookup.BackupCodesUserID, userID), Eq(lookup.BackupCodesCodeHash, codeHash)).
QueryRow(ctx, q, &id, t.boolDest(&used))
if errors.Is(err, sql.ErrNoRows) {
return nil
}
if err != nil {
return err
}
if used {
return fmt.Errorf("backup code already used")
}
n, err := t.Update(lookup.EntityUserTOTPBackupCodes).
Set(Set(lookup.BackupCodesUsed, true), Set(lookup.BackupCodesUsedAt, t.Now())).
Where(Eq(lookup.BackupCodesID, id), Eq(lookup.BackupCodesUsed, false)).Exec(ctx, q)
if err != nil {
return err
}
if n == 0 {
return fmt.Errorf("backup code already used")
}
valid = true
return nil
})
if err != nil {
return false, err
}
return valid, nil
}