Files
cmroubao/backend-api/internal/repository/sqlite/auth_repository_test.go
T

540 lines
13 KiB
Go
Raw Normal View History

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)
}
2026-07-27 12:23:05 +08:00
if err := runner.Down(context.Background()); err != nil {
t.Fatalf("Down(v5) error = %v", err)
}
if err := runner.Down(context.Background()); err != nil {
t.Fatalf("Down(v4) 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-v4) error = %v", err)
2026-07-27 12:23:05 +08:00
} else if applied != 3 {
t.Fatalf("Up(v3-v5) applied = %d, want 3", 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
}