feat(auth): implement user and device authentication
This commit is contained in:
@@ -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,
|
||||
¬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
|
||||
}
|
||||
Reference in New Issue
Block a user