561 lines
14 KiB
Go
561 lines
14 KiB
Go
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(v12) error = %v", err)
|
|
}
|
|
if err := runner.Down(context.Background()); err != nil {
|
|
t.Fatalf("Down(v11) error = %v", err)
|
|
}
|
|
if err := runner.Down(context.Background()); err != nil {
|
|
t.Fatalf("Down(v10) error = %v", err)
|
|
}
|
|
if err := runner.Down(context.Background()); err != nil {
|
|
t.Fatalf("Down(v9) error = %v", err)
|
|
}
|
|
if err := runner.Down(context.Background()); err != nil {
|
|
t.Fatalf("Down(v8) error = %v", err)
|
|
}
|
|
if err := runner.Down(context.Background()); err != nil {
|
|
t.Fatalf("Down(v7) error = %v", err)
|
|
}
|
|
if err := runner.Down(context.Background()); err != nil {
|
|
t.Fatalf("Down(v6) error = %v", err)
|
|
}
|
|
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-v12) error = %v", err)
|
|
} else if applied != 10 {
|
|
t.Fatalf("Up(v3-v12) applied = %d, want 10", 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,
|
|
¬Null,
|
|
&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
|
|
}
|