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(v16) error = %v", err) } if err := runner.Down(context.Background()); err != nil { t.Fatalf("Down(v15) error = %v", err) } if err := runner.Down(context.Background()); err != nil { t.Fatalf("Down(v14) error = %v", err) } if err := runner.Down(context.Background()); err != nil { t.Fatalf("Down(v13) 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-v16) error = %v", err) } else if applied != 14 { t.Fatalf("Up(v3-v16) applied = %d, want 14", 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 }