feat(auth): implement user and device authentication

This commit is contained in:
QiuSW
2026-07-26 15:18:48 +08:00
parent c5d3b215ff
commit 49db5b8305
66 changed files with 6216 additions and 271 deletions
@@ -0,0 +1,538 @@
package sqlite
import (
"context"
"crypto/subtle"
"database/sql"
"errors"
"time"
"cmroubao/backend-api/internal/domain"
"cmroubao/backend-api/internal/usecase"
)
func (s *Store) FindUserByUsername(
ctx context.Context,
username string,
) (domain.User, error) {
return scanUser(s.db.QueryRowContext(
ctx,
`SELECT
id, username, password_hash, role, is_active, created_at, updated_at
FROM users
WHERE username = ?`,
username,
))
}
func (s *Store) ProvisionUser(
ctx context.Context,
candidate domain.User,
) (domain.User, error) {
_, err := s.db.ExecContext(
ctx,
`INSERT INTO users (
id, username, password_hash, role, is_active, created_at, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?)`,
candidate.ID,
candidate.Username,
candidate.PasswordHash,
candidate.Role,
candidate.IsActive,
formatTimestamp(candidate.CreatedAt),
formatTimestamp(candidate.UpdatedAt),
)
if err != nil {
if isUniqueConstraint(err, "users.username") {
return domain.User{}, usecase.ErrAuthConflict
}
return domain.User{}, repositoryFailure(err)
}
return candidate, nil
}
func (s *Store) ProvisionDevice(
ctx context.Context,
candidate domain.Device,
) (domain.Device, error) {
_, err := s.db.ExecContext(
ctx,
`INSERT INTO devices (
id, name, token_hash, bound_user_id, app_version, pdd_version,
android_version, last_seen_at, is_enabled, created_at, updated_at
) VALUES (?, ?, ?, NULL, NULL, NULL, NULL, NULL, ?, ?, ?)`,
candidate.ID,
candidate.Name,
candidate.TokenHash,
candidate.IsEnabled,
formatTimestamp(candidate.CreatedAt),
formatTimestamp(candidate.UpdatedAt),
)
if err != nil {
if isUniqueConstraint(err, "devices.token_hash") {
return domain.Device{}, usecase.ErrAuthConflict
}
return domain.Device{}, repositoryFailure(err)
}
return candidate, nil
}
func (s *Store) SetUserActive(
ctx context.Context,
username string,
active bool,
updatedAt time.Time,
) error {
result, err := s.db.ExecContext(
ctx,
`UPDATE users
SET is_active = ?, updated_at = ?
WHERE username = ?`,
active,
formatTimestamp(updatedAt),
username,
)
return requireAffectedAuthResource(result, err)
}
func (s *Store) SetDeviceEnabled(
ctx context.Context,
deviceID string,
enabled bool,
updatedAt time.Time,
) error {
result, err := s.db.ExecContext(
ctx,
`UPDATE devices
SET is_enabled = ?, updated_at = ?
WHERE id = ?`,
enabled,
formatTimestamp(updatedAt),
deviceID,
)
return requireAffectedAuthResource(result, err)
}
func requireAffectedAuthResource(
result sql.Result,
err error,
) error {
if err != nil {
return repositoryFailure(err)
}
affected, err := result.RowsAffected()
if err != nil {
return repositoryFailure(err)
}
if affected != 1 {
return usecase.ErrRepositoryNotFound
}
return nil
}
func (s *Store) CreateAdminSession(
ctx context.Context,
session domain.AdminSession,
now time.Time,
) error {
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return repositoryFailure(err)
}
defer func() { _ = tx.Rollback() }()
var role domain.UserRole
var active bool
err = tx.QueryRowContext(
ctx,
"SELECT role, is_active FROM users WHERE id = ?",
session.UserID,
).Scan(&role, &active)
if errors.Is(err, sql.ErrNoRows) {
return usecase.ErrAuthCredentials
}
if err != nil {
return repositoryFailure(err)
}
if !active {
return usecase.ErrAuthDisabled
}
if role != domain.UserRoleAdmin {
return usecase.ErrAuthForbidden
}
if !session.ExpiresAt.After(now) {
return usecase.ErrAuthExpired
}
_, err = tx.ExecContext(
ctx,
`INSERT INTO admin_sessions (
id, token_hash, user_id, expires_at, revoked_at, created_at
) VALUES (?, ?, ?, ?, NULL, ?)`,
session.ID,
session.TokenHash,
session.UserID,
formatTimestamp(session.ExpiresAt),
formatTimestamp(session.CreatedAt),
)
if err != nil {
return repositoryFailure(err)
}
if err := tx.Commit(); err != nil {
return repositoryFailure(err)
}
return nil
}
func (s *Store) AuthenticateAdminSession(
ctx context.Context,
tokenHash string,
now time.Time,
) (domain.AuthPrincipal, error) {
var principal domain.AuthPrincipal
var active bool
var expiresAt string
var revokedAt sql.NullString
err := s.db.QueryRowContext(
ctx,
`SELECT
user.id, user.username, user.role, session.id,
session.expires_at, session.revoked_at, user.is_active
FROM admin_sessions AS session
JOIN users AS user ON user.id = session.user_id
WHERE session.token_hash = ?`,
tokenHash,
).Scan(
&principal.UserID,
&principal.Username,
&principal.Role,
&principal.SessionID,
&expiresAt,
&revokedAt,
&active,
)
if errors.Is(err, sql.ErrNoRows) {
return domain.AuthPrincipal{}, usecase.ErrAuthRevoked
}
if err != nil {
return domain.AuthPrincipal{}, repositoryFailure(err)
}
principal.ExpiresAt, err = parseTimestamp(expiresAt)
if err != nil {
return domain.AuthPrincipal{}, err
}
switch {
case revokedAt.Valid:
return domain.AuthPrincipal{}, usecase.ErrAuthRevoked
case !principal.ExpiresAt.After(now):
return domain.AuthPrincipal{}, usecase.ErrAuthExpired
case !active:
return domain.AuthPrincipal{}, usecase.ErrAuthDisabled
case principal.Role != domain.UserRoleAdmin:
return domain.AuthPrincipal{}, usecase.ErrAuthForbidden
default:
return principal, nil
}
}
func (s *Store) RevokeAdminSession(
ctx context.Context,
tokenHash string,
revokedAt time.Time,
) error {
_, err := s.db.ExecContext(
ctx,
`UPDATE admin_sessions
SET revoked_at = COALESCE(revoked_at, ?)
WHERE token_hash = ?`,
formatTimestamp(revokedAt),
tokenHash,
)
return repositoryFailure(err)
}
func (s *Store) CreateAccessTokenAndBindDevice(
ctx context.Context,
userID string,
deviceID string,
deviceTokenHash string,
appVersion string,
androidVersion string,
access domain.AccessToken,
now time.Time,
) (domain.Device, error) {
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return domain.Device{}, repositoryFailure(err)
}
defer func() { _ = tx.Rollback() }()
var role domain.UserRole
var active bool
err = tx.QueryRowContext(
ctx,
"SELECT role, is_active FROM users WHERE id = ?",
userID,
).Scan(&role, &active)
if errors.Is(err, sql.ErrNoRows) {
return domain.Device{}, usecase.ErrAuthCredentials
}
if err != nil {
return domain.Device{}, repositoryFailure(err)
}
if !active {
return domain.Device{}, usecase.ErrAuthDisabled
}
if role != domain.UserRoleBuyer {
return domain.Device{}, usecase.ErrAuthForbidden
}
device, err := getDeviceByID(ctx, tx, deviceID)
if errors.Is(err, usecase.ErrRepositoryNotFound) {
return domain.Device{}, usecase.ErrAuthCredentials
}
if err != nil {
return domain.Device{}, err
}
if subtle.ConstantTimeCompare(
[]byte(device.TokenHash),
[]byte(deviceTokenHash),
) != 1 {
return domain.Device{}, usecase.ErrAuthCredentials
}
if !device.IsEnabled {
return domain.Device{}, usecase.ErrAuthDisabled
}
if device.BoundUserID == nil {
result, err := tx.ExecContext(
ctx,
`UPDATE devices
SET bound_user_id = ?, app_version = ?, android_version = ?,
updated_at = ?
WHERE id = ? AND bound_user_id IS NULL`,
userID,
appVersion,
androidVersion,
formatTimestamp(now),
device.ID,
)
if err != nil {
return domain.Device{}, repositoryFailure(err)
}
affected, err := result.RowsAffected()
if err != nil {
return domain.Device{}, repositoryFailure(err)
}
if affected != 1 {
return domain.Device{}, usecase.ErrAuthConflict
}
boundUserID := userID
device.BoundUserID = &boundUserID
} else if *device.BoundUserID != userID {
return domain.Device{}, usecase.ErrAuthForbidden
} else {
_, err := tx.ExecContext(
ctx,
`UPDATE devices
SET app_version = ?, android_version = ?, updated_at = ?
WHERE id = ? AND bound_user_id = ?`,
appVersion,
androidVersion,
formatTimestamp(now),
device.ID,
userID,
)
if err != nil {
return domain.Device{}, repositoryFailure(err)
}
}
device.AppVersion = &appVersion
device.AndroidVersion = &androidVersion
device.UpdatedAt = now
if access.UserID != userID || access.DeviceID != deviceID {
return domain.Device{}, usecase.ErrRepositoryInvariant
}
if !access.ExpiresAt.After(now) {
return domain.Device{}, usecase.ErrAuthExpired
}
_, err = tx.ExecContext(
ctx,
`INSERT INTO access_tokens (
id, token_hash, user_id, device_id, expires_at, revoked_at, created_at
) VALUES (?, ?, ?, ?, ?, NULL, ?)`,
access.ID,
access.TokenHash,
access.UserID,
access.DeviceID,
formatTimestamp(access.ExpiresAt),
formatTimestamp(access.CreatedAt),
)
if err != nil {
return domain.Device{}, repositoryFailure(err)
}
if err := tx.Commit(); err != nil {
return domain.Device{}, repositoryFailure(err)
}
return device, nil
}
func (s *Store) AuthenticateAccessToken(
ctx context.Context,
tokenHash string,
now time.Time,
) (domain.AuthPrincipal, error) {
var principal domain.AuthPrincipal
var userActive bool
var deviceEnabled bool
var boundUserID sql.NullString
var expiresAt string
var revokedAt sql.NullString
err := s.db.QueryRowContext(
ctx,
`SELECT
user.id, user.username, user.role, access.id, device.id,
access.expires_at, access.revoked_at, user.is_active,
device.is_enabled, device.bound_user_id
FROM access_tokens AS access
JOIN users AS user ON user.id = access.user_id
JOIN devices AS device ON device.id = access.device_id
WHERE access.token_hash = ?`,
tokenHash,
).Scan(
&principal.UserID,
&principal.Username,
&principal.Role,
&principal.SessionID,
&principal.DeviceID,
&expiresAt,
&revokedAt,
&userActive,
&deviceEnabled,
&boundUserID,
)
if errors.Is(err, sql.ErrNoRows) {
return domain.AuthPrincipal{}, usecase.ErrAuthRevoked
}
if err != nil {
return domain.AuthPrincipal{}, repositoryFailure(err)
}
principal.ExpiresAt, err = parseTimestamp(expiresAt)
if err != nil {
return domain.AuthPrincipal{}, err
}
switch {
case revokedAt.Valid:
return domain.AuthPrincipal{}, usecase.ErrAuthRevoked
case !principal.ExpiresAt.After(now):
return domain.AuthPrincipal{}, usecase.ErrAuthExpired
case !userActive || !deviceEnabled:
return domain.AuthPrincipal{}, usecase.ErrAuthDisabled
case principal.Role != domain.UserRoleBuyer:
return domain.AuthPrincipal{}, usecase.ErrAuthForbidden
case !boundUserID.Valid || boundUserID.String != principal.UserID:
return domain.AuthPrincipal{}, usecase.ErrAuthForbidden
default:
return principal, nil
}
}
func scanUser(scanner rowScanner) (domain.User, error) {
var user domain.User
var createdAt string
var updatedAt string
err := scanner.Scan(
&user.ID,
&user.Username,
&user.PasswordHash,
&user.Role,
&user.IsActive,
&createdAt,
&updatedAt,
)
if errors.Is(err, sql.ErrNoRows) {
return domain.User{}, usecase.ErrRepositoryNotFound
}
if err != nil {
return domain.User{}, repositoryFailure(err)
}
user.CreatedAt, err = parseTimestamp(createdAt)
if err != nil {
return domain.User{}, err
}
user.UpdatedAt, err = parseTimestamp(updatedAt)
if err != nil {
return domain.User{}, err
}
return user, nil
}
func getDeviceByID(
ctx context.Context,
queryer queryRower,
deviceID string,
) (domain.Device, error) {
var device domain.Device
var boundUserID sql.NullString
var appVersion sql.NullString
var androidVersion sql.NullString
var pddVersion sql.NullString
var lastSeenAt sql.NullString
var createdAt string
var updatedAt string
err := queryer.QueryRowContext(
ctx,
`SELECT
id, name, token_hash, bound_user_id, app_version, android_version,
pdd_version, last_seen_at, is_enabled, created_at, updated_at
FROM devices
WHERE id = ?`,
deviceID,
).Scan(
&device.ID,
&device.Name,
&device.TokenHash,
&boundUserID,
&appVersion,
&androidVersion,
&pddVersion,
&lastSeenAt,
&device.IsEnabled,
&createdAt,
&updatedAt,
)
if errors.Is(err, sql.ErrNoRows) {
return domain.Device{}, usecase.ErrRepositoryNotFound
}
if err != nil {
return domain.Device{}, repositoryFailure(err)
}
if boundUserID.Valid {
device.BoundUserID = &boundUserID.String
}
if appVersion.Valid {
device.AppVersion = &appVersion.String
}
if androidVersion.Valid {
device.AndroidVersion = &androidVersion.String
}
if pddVersion.Valid {
device.PDDVersion = &pddVersion.String
}
if lastSeenAt.Valid {
parsed, err := parseTimestamp(lastSeenAt.String)
if err != nil {
return domain.Device{}, err
}
device.LastSeenAt = &parsed
}
device.CreatedAt, err = parseTimestamp(createdAt)
if err != nil {
return domain.Device{}, err
}
device.UpdatedAt, err = parseTimestamp(updatedAt)
if err != nil {
return domain.Device{}, err
}
return device, nil
}
var _ usecase.AuthRepository = (*Store)(nil)
@@ -0,0 +1,533 @@
package sqlite_test
import (
"context"
"database/sql"
"encoding/base64"
"errors"
"sync"
"testing"
"time"
"cmroubao/backend-api/internal/domain"
"cmroubao/backend-api/internal/platform/migration"
"cmroubao/backend-api/internal/platform/password"
repository "cmroubao/backend-api/internal/repository/sqlite"
"cmroubao/backend-api/internal/usecase"
"golang.org/x/crypto/bcrypt"
)
func TestAuthRepositoryPersistsOnlyHashesAndRevokesAdminSession(
t *testing.T,
) {
db := openDatabase(t)
service := newRepositoryAuthService(t, db)
ctx := context.Background()
admin, err := service.ProvisionUser(ctx, usecase.ProvisionUserCommand{
Username: "admin",
Password: "admin-password",
Role: domain.UserRoleAdmin,
Active: true,
})
if err != nil {
t.Fatalf("ProvisionUser() error = %v", err)
}
var passwordHash string
if err := db.QueryRow(
"SELECT password_hash FROM users WHERE id = ?",
admin.ID,
).Scan(&passwordHash); err != nil {
t.Fatalf("query password hash: %v", err)
}
if passwordHash == "admin-password" || passwordHash == "" {
t.Fatal("database contains missing or plaintext password")
}
login, err := service.LoginAdmin(ctx, usecase.LoginAdminCommand{
Username: "admin",
Password: "admin-password",
})
if err != nil {
t.Fatalf("LoginAdmin() error = %v", err)
}
var storedTokenHash string
if err := db.QueryRow(
"SELECT token_hash FROM admin_sessions WHERE user_id = ?",
admin.ID,
).Scan(&storedTokenHash); err != nil {
t.Fatalf("query token hash: %v", err)
}
if storedTokenHash == login.Token || len(storedTokenHash) != 64 {
t.Fatal("database contains missing or plaintext admin token")
}
principal, err := service.AuthenticateAdmin(ctx, login.Token)
if err != nil {
t.Fatalf("AuthenticateAdmin() error = %v", err)
}
if principal.UserID != admin.ID ||
principal.Role != domain.UserRoleAdmin ||
principal.DeviceID != "" {
t.Fatalf("AuthenticateAdmin() principal = %+v", principal)
}
if err := service.SetUserActive(
ctx,
usecase.SetUserActiveCommand{
Username: admin.Username,
Active: false,
},
); err != nil {
t.Fatalf("disable admin: %v", err)
}
_, err = service.AuthenticateAdmin(ctx, login.Token)
assertRepositoryAuthCode(
t,
err,
"AUTH_ACCOUNT_OR_DEVICE_DISABLED",
)
if err := service.SetUserActive(
ctx,
usecase.SetUserActiveCommand{
Username: admin.Username,
Active: true,
},
); err != nil {
t.Fatalf("enable admin: %v", err)
}
if _, err := db.Exec(
"UPDATE admin_sessions SET expires_at = ? WHERE user_id = ?",
time.Date(2000, 1, 1, 0, 0, 0, 0, time.UTC).Format(
time.RFC3339Nano,
),
admin.ID,
); err != nil {
t.Fatalf("expire admin session: %v", err)
}
_, err = service.AuthenticateAdmin(ctx, login.Token)
assertRepositoryAuthCode(t, err, "AUTH_SESSION_EXPIRED")
if _, err := db.Exec(
"UPDATE admin_sessions SET expires_at = ? WHERE user_id = ?",
login.ExpiresAt.Format(time.RFC3339Nano),
admin.ID,
); err != nil {
t.Fatalf("restore admin expiry: %v", err)
}
if err := service.LogoutAdmin(ctx, login.Token); err != nil {
t.Fatalf("LogoutAdmin() error = %v", err)
}
_, err = service.AuthenticateAdmin(ctx, login.Token)
assertRepositoryAuthCode(t, err, "AUTH_INVALID_TOKEN")
}
func TestAuthRepositoryFirstBuyerBindingIsAtomicAndPersistent(
t *testing.T,
) {
db := openDatabase(t)
service := newRepositoryAuthService(t, db)
ctx := context.Background()
for _, username := range []string{"buyer-a", "buyer-b"} {
if _, err := service.ProvisionUser(
ctx,
usecase.ProvisionUserCommand{
Username: username,
Password: "buyer-password",
Role: domain.UserRoleBuyer,
Active: true,
},
); err != nil {
t.Fatalf("ProvisionUser(%s) error = %v", username, err)
}
}
device, err := service.ProvisionDevice(
ctx,
usecase.ProvisionDeviceCommand{
Name: "PDD test device",
Enabled: true,
},
)
if err != nil {
t.Fatalf("ProvisionDevice() error = %v", err)
}
var storedDeviceTokenHash string
if err := db.QueryRow(
"SELECT token_hash FROM devices WHERE id = ?",
device.Device.ID,
).Scan(&storedDeviceTokenHash); err != nil {
t.Fatalf("query device token hash: %v", err)
}
if storedDeviceTokenHash == device.DeviceToken ||
len(storedDeviceTokenHash) != 64 {
t.Fatal("database contains missing or plaintext device token")
}
start := make(chan struct{})
results := make(chan usecase.AccessTokenResult, 2)
failures := make(chan error, 2)
var wait sync.WaitGroup
for _, username := range []string{"buyer-a", "buyer-b"} {
username := username
wait.Add(1)
go func() {
defer wait.Done()
<-start
result, err := service.LoginBuyerDevice(
ctx,
usecase.LoginBuyerDeviceCommand{
Username: username,
Password: "buyer-password",
DeviceID: device.Device.ID,
DeviceToken: device.DeviceToken,
AppVersion: "0.1.0",
AndroidVersion: "16",
},
)
if err != nil {
failures <- err
return
}
results <- result
}()
}
close(start)
wait.Wait()
close(results)
close(failures)
if len(results) != 1 || len(failures) != 1 {
t.Fatalf(
"concurrent binding successes=%d failures=%d, want 1/1",
len(results),
len(failures),
)
}
success := <-results
failure := <-failures
assertRepositoryAuthCode(t, failure, "AUTH_FORBIDDEN")
principal, err := service.AuthenticateAccessToken(ctx, success.Token)
if err != nil {
t.Fatalf("AuthenticateAccessToken() error = %v", err)
}
if principal.DeviceID != device.Device.ID ||
principal.UserID != success.User.ID {
t.Fatalf("access principal = %+v", principal)
}
if success.Device.AppVersion == nil ||
*success.Device.AppVersion != "0.1.0" ||
success.Device.AndroidVersion == nil ||
*success.Device.AndroidVersion != "16" {
t.Fatalf("device versions = %+v", success.Device)
}
var appVersion string
var androidVersion string
var pddVersion sql.NullString
if err := db.QueryRow(
`SELECT app_version, android_version, pdd_version
FROM devices WHERE id = ?`,
device.Device.ID,
).Scan(&appVersion, &androidVersion, &pddVersion); err != nil {
t.Fatalf("query device versions: %v", err)
}
if appVersion != "0.1.0" ||
androidVersion != "16" ||
pddVersion.Valid {
t.Fatalf(
"stored versions app=%q android=%q pdd=%v",
appVersion,
androidVersion,
pddVersion,
)
}
var boundUserID string
if err := db.QueryRow(
"SELECT bound_user_id FROM devices WHERE id = ?",
device.Device.ID,
).Scan(&boundUserID); err != nil {
t.Fatalf("query bound user: %v", err)
}
if boundUserID != success.User.ID {
t.Fatalf("bound user = %s, want %s", boundUserID, success.User.ID)
}
var storedTokenHash string
if err := db.QueryRow(
"SELECT token_hash FROM access_tokens WHERE user_id = ?",
success.User.ID,
).Scan(&storedTokenHash); err != nil {
t.Fatalf("query access token hash: %v", err)
}
if storedTokenHash == success.Token || len(storedTokenHash) != 64 {
t.Fatal("database contains missing or plaintext access token")
}
}
func TestAuthRepositoryRechecksCurrentUserDeviceAndExpiryState(
t *testing.T,
) {
db := openDatabase(t)
service := newRepositoryAuthService(t, db)
ctx := context.Background()
user, err := service.ProvisionUser(ctx, usecase.ProvisionUserCommand{
Username: "buyer",
Password: "buyer-password",
Role: domain.UserRoleBuyer,
Active: true,
})
if err != nil {
t.Fatalf("ProvisionUser() error = %v", err)
}
device, err := service.ProvisionDevice(
ctx,
usecase.ProvisionDeviceCommand{Name: "Device", Enabled: true},
)
if err != nil {
t.Fatalf("ProvisionDevice() error = %v", err)
}
login, err := service.LoginBuyerDevice(
ctx,
usecase.LoginBuyerDeviceCommand{
Username: user.Username,
Password: "buyer-password",
DeviceID: device.Device.ID,
DeviceToken: device.DeviceToken,
AppVersion: "0.1.0",
AndroidVersion: "16",
},
)
if err != nil {
t.Fatalf("LoginBuyerDevice() error = %v", err)
}
if err := service.SetUserActive(
ctx,
usecase.SetUserActiveCommand{
Username: user.Username,
Active: false,
},
); err != nil {
t.Fatalf("disable user: %v", err)
}
_, err = service.AuthenticateAccessToken(ctx, login.Token)
assertRepositoryAuthCode(
t,
err,
"AUTH_ACCOUNT_OR_DEVICE_DISABLED",
)
if err := service.SetUserActive(
ctx,
usecase.SetUserActiveCommand{
Username: user.Username,
Active: true,
},
); err != nil {
t.Fatalf("enable user: %v", err)
}
if err := service.SetDeviceEnabled(
ctx,
usecase.SetDeviceEnabledCommand{
DeviceID: device.Device.ID,
Enabled: false,
},
); err != nil {
t.Fatalf("disable device: %v", err)
}
_, err = service.AuthenticateAccessToken(ctx, login.Token)
assertRepositoryAuthCode(
t,
err,
"AUTH_ACCOUNT_OR_DEVICE_DISABLED",
)
if err := service.SetDeviceEnabled(
ctx,
usecase.SetDeviceEnabledCommand{
DeviceID: device.Device.ID,
Enabled: true,
},
); err != nil {
t.Fatalf("enable device: %v", err)
}
if _, err := db.Exec(
"UPDATE access_tokens SET expires_at = ? WHERE user_id = ?",
time.Date(2000, 1, 1, 0, 0, 0, 0, time.UTC).Format(
time.RFC3339Nano,
),
user.ID,
); err != nil {
t.Fatalf("expire access token: %v", err)
}
_, err = service.AuthenticateAccessToken(ctx, login.Token)
assertRepositoryAuthCode(t, err, "AUTH_SESSION_EXPIRED")
if _, err := db.Exec(
`UPDATE access_tokens
SET expires_at = ?, revoked_at = ?
WHERE user_id = ?`,
login.ExpiresAt.Format(time.RFC3339Nano),
time.Date(2026, 7, 26, 12, 30, 0, 0, time.UTC).Format(
time.RFC3339Nano,
),
user.ID,
); err != nil {
t.Fatalf("revoke access token: %v", err)
}
_, err = service.AuthenticateAccessToken(ctx, login.Token)
assertRepositoryAuthCode(t, err, "AUTH_INVALID_TOKEN")
}
func TestAuthMigrationCanRollbackWithoutRebuildingPurchaseTasks(
t *testing.T,
) {
db := openDatabase(t)
runner, err := migration.New(db)
if err != nil {
t.Fatalf("migration.New() error = %v", err)
}
if err := runner.Down(context.Background()); err != nil {
t.Fatalf("Down(v3) error = %v", err)
}
if tableExists(t, db, "users") {
t.Fatal("users table exists after auth migration rollback")
}
if columnExists(t, db, "purchase_tasks", "created_by_user_id") {
t.Fatal("created_by_user_id exists after auth migration rollback")
}
if columnExists(t, db, "task_events", "actor_user_id") {
t.Fatal("actor_user_id exists after auth migration rollback")
}
if !tableExists(t, db, "purchase_tasks") {
t.Fatal("purchase_tasks was lost during auth migration rollback")
}
if applied, err := runner.Up(context.Background()); err != nil {
t.Fatalf("Up(v3) error = %v", err)
} else if applied != 1 {
t.Fatalf("Up(v3) applied = %d, want 1", applied)
}
}
func newRepositoryAuthService(
t *testing.T,
db *sql.DB,
) *usecase.AuthService {
t.Helper()
store, err := repository.New(db)
if err != nil {
t.Fatalf("repository.New() error = %v", err)
}
passwords, err := password.NewBcrypt(bcrypt.MinCost)
if err != nil {
t.Fatalf("password.NewBcrypt() error = %v", err)
}
service, err := usecase.NewAuthService(
store,
passwords,
repositoryAuthClock{
now: time.Date(2026, 7, 26, 12, 0, 0, 0, time.UTC),
},
&repositoryAuthIDGenerator{},
&repositoryAuthTokenGenerator{},
)
if err != nil {
t.Fatalf("usecase.NewAuthService() error = %v", err)
}
return service
}
func assertRepositoryAuthCode(t *testing.T, err error, code string) {
t.Helper()
var authError *usecase.Error
if !errors.As(err, &authError) || authError.Code != code {
t.Fatalf("error = %v, want code %s", err, code)
}
}
func tableExists(
t *testing.T,
db *sql.DB,
name string,
) bool {
t.Helper()
var count int
if err := db.QueryRow(
"SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = ?",
name,
).Scan(&count); err != nil {
t.Fatalf("query table %s: %v", name, err)
}
return count == 1
}
func columnExists(
t *testing.T,
db *sql.DB,
table string,
column string,
) bool {
t.Helper()
rows, err := db.Query("PRAGMA table_info(" + table + ")")
if err != nil {
t.Fatalf("PRAGMA table_info(%s): %v", table, err)
}
defer rows.Close()
for rows.Next() {
var cid int
var name string
var dataType string
var notNull int
var defaultValue any
var primaryKey int
if err := rows.Scan(
&cid,
&name,
&dataType,
&notNull,
&defaultValue,
&primaryKey,
); err != nil {
t.Fatalf("scan table_info(%s): %v", table, err)
}
if name == column {
return true
}
}
if err := rows.Err(); err != nil {
t.Fatalf("iterate table_info(%s): %v", table, err)
}
return false
}
type repositoryAuthClock struct {
now time.Time
}
func (clock repositoryAuthClock) Now() time.Time {
return clock.now
}
type repositoryAuthIDGenerator struct {
mu sync.Mutex
next int
}
func (generator *repositoryAuthIDGenerator) NewID() (string, error) {
generator.mu.Lock()
defer generator.mu.Unlock()
generator.next++
return uuid(600 + generator.next), nil
}
type repositoryAuthTokenGenerator struct {
mu sync.Mutex
next byte
}
func (generator *repositoryAuthTokenGenerator) NewToken() (string, error) {
generator.mu.Lock()
defer generator.mu.Unlock()
generator.next++
value := make([]byte, 32)
for index := range value {
value[index] = generator.next + byte(index)
}
return base64.RawURLEncoding.EncodeToString(value), nil
}
@@ -49,6 +49,7 @@ func scanAsset(scanner rowScanner) (domain.Asset, error) {
func scanTask(scanner rowScanner) (domain.PurchaseTask, error) {
var task domain.PurchaseTask
var createdByUserID sql.NullString
var sourceRef sql.NullString
var maxBudget sql.NullInt64
var cancelReason sql.NullString
@@ -58,6 +59,7 @@ func scanTask(scanner rowScanner) (domain.PurchaseTask, error) {
err := scanner.Scan(
&task.ID,
&task.CreatorSubject,
&createdByUserID,
&sourceRef,
&task.Title,
&task.Description,
@@ -79,6 +81,9 @@ func scanTask(scanner rowScanner) (domain.PurchaseTask, error) {
if sourceRef.Valid {
task.SourceRef = &sourceRef.String
}
if createdByUserID.Valid {
task.CreatedByUserID = &createdByUserID.String
}
if maxBudget.Valid {
task.MaxBudgetCents = &maxBudget.Int64
}
@@ -137,7 +142,7 @@ func getTaskByID(
task, err := scanTask(queryer.QueryRowContext(
ctx,
`SELECT
id, creator_subject, source_ref, title, description, sku,
id, creator_subject, created_by_user_id, source_ref, title, description, sku,
image_asset_id, quantity, max_budget_cents, currency, status,
version, cancel_reason, canceled_at, created_at, updated_at
FROM purchase_tasks
@@ -321,6 +321,18 @@ func TestStoreEnforcesAssetOwnershipSourceReferenceAndStableCursor(
func openStore(t *testing.T) *repository.Store {
t.Helper()
db := openDatabase(t)
now := time.Now().UTC().Format(time.RFC3339Nano)
if _, err := db.ExecContext(
context.Background(),
`INSERT INTO users (
id, username, password_hash, role, is_active, created_at, updated_at
) VALUES (?, 'task-audit-admin', 'test-only-hash', 'ADMIN', 1, ?, ?)`,
uuid(999),
now,
now,
); err != nil {
t.Fatalf("seed admin user: %v", err)
}
store, err := repository.New(db)
if err != nil {
t.Fatalf("repository.New() error = %v", err)
@@ -368,21 +380,23 @@ func testTask(
createdAt time.Time,
) domain.PurchaseTask {
budget := int64(2000 + index)
actorUserID := uuid(999)
return domain.PurchaseTask{
ID: uuid(index),
CreatorSubject: "local-admin",
SourceRef: &source,
Title: "title " + string(rune('0'+index)),
Description: "description",
SKU: "SKU-" + string(rune('0'+index)),
ImageAssetID: assetID,
Quantity: index,
MaxBudgetCents: &budget,
Currency: domain.CurrencyCNY,
Status: domain.TaskStatusPending,
Version: 1,
CreatedAt: createdAt,
UpdatedAt: createdAt,
ID: uuid(index),
CreatorSubject: "local-admin",
CreatedByUserID: &actorUserID,
SourceRef: &source,
Title: "title " + string(rune('0'+index)),
Description: "description",
SKU: "SKU-" + string(rune('0'+index)),
ImageAssetID: assetID,
Quantity: index,
MaxBudgetCents: &budget,
Currency: domain.CurrencyCNY,
Status: domain.TaskStatusPending,
Version: 1,
CreatedAt: createdAt,
UpdatedAt: createdAt,
}
}
@@ -82,12 +82,13 @@ func (s *Store) CreateTaskIdempotent(
_, err = tx.ExecContext(
ctx,
`INSERT INTO purchase_tasks (
id, creator_subject, source_ref, title, description, sku,
id, creator_subject, created_by_user_id, source_ref, title, description, sku,
image_asset_id, quantity, max_budget_cents, currency, status,
version, cancel_reason, canceled_at, created_at, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, NULL, ?, ?)`,
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, NULL, ?, ?)`,
candidate.ID,
candidate.CreatorSubject,
nullableString(candidate.CreatedByUserID),
nullableString(candidate.SourceRef),
candidate.Title,
candidate.Description,
@@ -142,7 +143,7 @@ func (s *Store) ListTasks(
) ([]domain.PurchaseTask, error) {
var query strings.Builder
query.WriteString(`SELECT
id, creator_subject, source_ref, title, description, sku,
id, creator_subject, created_by_user_id, source_ref, title, description, sku,
image_asset_id, quantity, max_budget_cents, currency, status,
version, cancel_reason, canceled_at, created_at, updated_at
FROM purchase_tasks
@@ -241,7 +242,7 @@ func (s *Store) GetTaskDetail(
}
rows, err := tx.QueryContext(
ctx,
`SELECT id, task_id, event_type, message, occurred_at
`SELECT id, task_id, actor_user_id, event_type, message, occurred_at
FROM task_events
WHERE task_id = ?
ORDER BY occurred_at ASC, id ASC`,
@@ -254,16 +255,21 @@ func (s *Store) GetTaskDetail(
events := make([]domain.TaskEvent, 0)
for rows.Next() {
var event domain.TaskEvent
var actorUserID sql.NullString
var occurredAt string
if err := rows.Scan(
&event.ID,
&event.TaskID,
&actorUserID,
&event.Type,
&event.Message,
&occurredAt,
); err != nil {
return domain.TaskDetail{}, repositoryFailure(err)
}
if actorUserID.Valid {
event.ActorUserID = &actorUserID.String
}
event.OccurredAt, err = parseTimestamp(occurredAt)
if err != nil {
return domain.TaskDetail{}, err
@@ -357,10 +363,11 @@ func insertTaskEvent(
_, err := tx.ExecContext(
ctx,
`INSERT INTO task_events (
id, task_id, event_type, message, occurred_at
) VALUES (?, ?, ?, ?, ?)`,
id, task_id, actor_user_id, event_type, message, occurred_at
) VALUES (?, ?, ?, ?, ?, ?)`,
event.ID,
event.TaskID,
nullableString(event.ActorUserID),
event.Type,
event.Message,
formatTimestamp(event.OccurredAt),