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
+69
View File
@@ -13,6 +13,8 @@ const (
HTTPAddressEnvironment = "CMROUBAO_HTTP_ADDR"
DatabasePathEnvironment = "CMROUBAO_DATABASE_PATH"
AssetDirectoryEnvironment = "CMROUBAO_ASSET_DIR"
TLSCertificateEnvironment = "CMROUBAO_TLS_CERT_FILE"
TLSPrivateKeyEnvironment = "CMROUBAO_TLS_KEY_FILE"
defaultHTTPAddress = "127.0.0.1:8080"
defaultDatabasePath = "var/cmroubao.db"
@@ -25,6 +27,8 @@ type Config struct {
HTTPAddress string
DatabasePath string
AssetDirectory string
TLSCertificate string
TLSPrivateKey string
ReadHeaderTimeout time.Duration
ReadTimeout time.Duration
WriteTimeout time.Duration
@@ -68,11 +72,39 @@ func Load(lookup LookupEnvironment) (Config, error) {
if err != nil {
return Config{}, err
}
tlsCertificate, certificateSet, err := optionalEnvironmentValue(
lookup,
TLSCertificateEnvironment,
)
if err != nil {
return Config{}, err
}
tlsPrivateKey, privateKeySet, err := optionalEnvironmentValue(
lookup,
TLSPrivateKeyEnvironment,
)
if err != nil {
return Config{}, err
}
if certificateSet != privateKeySet {
return Config{}, errors.New(
TLSCertificateEnvironment + " and " +
TLSPrivateKeyEnvironment + " must be set together",
)
}
if !certificateSet && !isLoopbackAddress(httpAddress) {
return Config{}, errors.New(
HTTPAddressEnvironment +
" must use loopback unless TLS is configured",
)
}
return Config{
HTTPAddress: httpAddress,
DatabasePath: filepath.Clean(databasePath),
AssetDirectory: assetDirectory,
TLSCertificate: cleanOptionalPath(tlsCertificate),
TLSPrivateKey: cleanOptionalPath(tlsPrivateKey),
ReadHeaderTimeout: 5 * time.Second,
ReadTimeout: 15 * time.Second,
WriteTimeout: 30 * time.Second,
@@ -82,6 +114,43 @@ func Load(lookup LookupEnvironment) (Config, error) {
}, nil
}
func cleanOptionalPath(value string) string {
if value == "" {
return ""
}
return filepath.Clean(value)
}
func optionalEnvironmentValue(
lookup LookupEnvironment,
name string,
) (string, bool, error) {
value, exists := lookup(name)
if !exists {
return "", false, nil
}
value = strings.TrimSpace(value)
if value == "" {
return "", false, errors.New(name + " must not be blank")
}
if strings.ContainsRune(value, '\x00') {
return "", false, errors.New(name + " contains an invalid character")
}
return value, true, nil
}
func isLoopbackAddress(address string) bool {
host, _, err := net.SplitHostPort(address)
if err != nil {
return false
}
if strings.EqualFold(host, "localhost") {
return true
}
ip := net.ParseIP(host)
return ip != nil && ip.IsLoopback()
}
func validatedAssetDirectory(path string) (string, error) {
if strings.ContainsRune(path, '\x00') {
return "", errors.New(
@@ -39,6 +39,8 @@ func TestLoadAcceptsExplicitConfiguration(t *testing.T) {
HTTPAddressEnvironment: "192.0.2.10:9090",
DatabasePathEnvironment: "tmp/test.db",
AssetDirectoryEnvironment: "tmp/assets",
TLSCertificateEnvironment: "tmp/server.crt",
TLSPrivateKeyEnvironment: "tmp/server.key",
}
cfg, err := Load(mapEnvironment(values))
@@ -55,6 +57,14 @@ func TestLoadAcceptsExplicitConfiguration(t *testing.T) {
if cfg.AssetDirectory != filepath.Clean(values[AssetDirectoryEnvironment]) {
t.Fatalf("AssetDirectory = %q", cfg.AssetDirectory)
}
if cfg.TLSCertificate != filepath.Clean(values[TLSCertificateEnvironment]) ||
cfg.TLSPrivateKey != filepath.Clean(values[TLSPrivateKeyEnvironment]) {
t.Fatalf(
"TLS files = %q / %q",
cfg.TLSCertificate,
cfg.TLSPrivateKey,
)
}
}
func TestLoadRejectsUnsafeOrInvalidValues(t *testing.T) {
@@ -116,6 +126,25 @@ func TestLoadRejectsUnsafeOrInvalidValues(t *testing.T) {
AssetDirectoryEnvironment: ".",
},
},
{
name: "non loopback without TLS",
values: map[string]string{
HTTPAddressEnvironment: "0.0.0.0:8080",
},
},
{
name: "TLS certificate without key",
values: map[string]string{
TLSCertificateEnvironment: "tmp/server.crt",
},
},
{
name: "blank TLS key",
values: map[string]string{
TLSCertificateEnvironment: "tmp/server.crt",
TLSPrivateKeyEnvironment: " ",
},
},
}
for _, test := range tests {
+146
View File
@@ -0,0 +1,146 @@
package domain
import (
"errors"
"fmt"
"strings"
"time"
"unicode/utf8"
)
type UserRole string
const (
UserRoleAdmin UserRole = "ADMIN"
UserRoleBuyer UserRole = "BUYER"
MaxUsernameBytes = 128
MaxDeviceNameBytes = 128
MaxVersionBytes = 128
MinPasswordBytes = 12
MaxPasswordBytes = 72
)
type User struct {
ID string
Username string
PasswordHash string
Role UserRole
IsActive bool
CreatedAt time.Time
UpdatedAt time.Time
}
type Device struct {
ID string
Name string
TokenHash string
BoundUserID *string
AppVersion *string
AndroidVersion *string
PDDVersion *string
LastSeenAt *time.Time
IsEnabled bool
CreatedAt time.Time
UpdatedAt time.Time
}
type AdminSession struct {
ID string
TokenHash string
UserID string
ExpiresAt time.Time
RevokedAt *time.Time
CreatedAt time.Time
}
type AccessToken struct {
ID string
TokenHash string
UserID string
DeviceID string
ExpiresAt time.Time
RevokedAt *time.Time
CreatedAt time.Time
}
type AuthPrincipal struct {
UserID string
Username string
Role UserRole
SessionID string
DeviceID string
ExpiresAt time.Time
}
type AuthValidationError struct {
Fields map[string]string
}
func (e *AuthValidationError) Error() string {
return "authentication input validation failed"
}
func NormalizeUsername(value string) string {
return strings.ToLower(strings.TrimSpace(value))
}
func ValidateUserInput(
username string,
password string,
role UserRole,
) error {
fields := make(map[string]string)
normalized := NormalizeUsername(username)
switch {
case normalized == "":
fields["username"] = "required"
case !utf8.ValidString(normalized):
fields["username"] = "must be valid UTF-8"
case len([]byte(normalized)) > MaxUsernameBytes:
fields["username"] = fmt.Sprintf(
"must not exceed %d UTF-8 bytes",
MaxUsernameBytes,
)
}
switch {
case password == "":
fields["password"] = "required"
case !utf8.ValidString(password):
fields["password"] = "must be valid UTF-8"
case len([]byte(password)) < MinPasswordBytes:
fields["password"] = fmt.Sprintf(
"must be at least %d UTF-8 bytes",
MinPasswordBytes,
)
case len([]byte(password)) > MaxPasswordBytes:
fields["password"] = fmt.Sprintf(
"must not exceed %d UTF-8 bytes",
MaxPasswordBytes,
)
}
if role != UserRoleAdmin && role != UserRoleBuyer {
fields["role"] = "must be ADMIN or BUYER"
}
if len(fields) > 0 {
return &AuthValidationError{Fields: fields}
}
return nil
}
func ValidateDeviceName(name string) error {
name = strings.TrimSpace(name)
switch {
case name == "":
return errors.New("device name is required")
case !utf8.ValidString(name):
return errors.New("device name must be valid UTF-8")
case len([]byte(name)) > MaxDeviceNameBytes:
return fmt.Errorf(
"device name must not exceed %d UTF-8 bytes",
MaxDeviceNameBytes,
)
default:
return nil
}
}
+63
View File
@@ -0,0 +1,63 @@
package domain
import (
"strings"
"testing"
)
func TestNormalizeUsernameAndValidateUserInput(t *testing.T) {
if got := NormalizeUsername(" Admin.User "); got != "admin.user" {
t.Fatalf("NormalizeUsername() = %q", got)
}
if err := ValidateUserInput(
"admin.user",
"strong-password",
UserRoleAdmin,
); err != nil {
t.Fatalf("ValidateUserInput(valid) error = %v", err)
}
for name, input := range map[string]struct {
username string
password string
role UserRole
}{
"blank username": {" ", "password", UserRoleAdmin},
"blank password": {"admin", "", UserRoleAdmin},
"short password": {"admin", "short", UserRoleAdmin},
"invalid password UTF-8": {
"admin",
string([]byte{0xff, 0xfe, 0xfd}),
UserRoleAdmin,
},
"oversized password": {
"admin",
strings.Repeat("x", MaxPasswordBytes+1),
UserRoleAdmin,
},
"invalid role": {"admin", "password", "AUDITOR"},
} {
t.Run(name, func(t *testing.T) {
if err := ValidateUserInput(
input.username,
input.password,
input.role,
); err == nil {
t.Fatal("ValidateUserInput() error = nil")
}
})
}
}
func TestValidateDeviceName(t *testing.T) {
if err := ValidateDeviceName("Test device"); err != nil {
t.Fatalf("ValidateDeviceName(valid) error = %v", err)
}
if err := ValidateDeviceName(" "); err == nil {
t.Fatal("ValidateDeviceName(blank) error = nil")
}
if err := ValidateDeviceName(
strings.Repeat("x", MaxDeviceNameBytes+1),
); err == nil {
t.Fatal("ValidateDeviceName(oversized) error = nil")
}
}
+23 -21
View File
@@ -33,30 +33,32 @@ const (
)
type PurchaseTask struct {
ID string
CreatorSubject string
SourceRef *string
Title string
Description string
SKU string
ImageAssetID string
Quantity int
MaxBudgetCents *int64
Currency string
Status TaskStatus
Version int64
CancelReason *string
CanceledAt *time.Time
CreatedAt time.Time
UpdatedAt time.Time
ID string
CreatorSubject string
CreatedByUserID *string
SourceRef *string
Title string
Description string
SKU string
ImageAssetID string
Quantity int
MaxBudgetCents *int64
Currency string
Status TaskStatus
Version int64
CancelReason *string
CanceledAt *time.Time
CreatedAt time.Time
UpdatedAt time.Time
}
type TaskEvent struct {
ID string
TaskID string
Type string
Message string
OccurredAt time.Time
ID string
TaskID string
ActorUserID *string
Type string
Message string
OccurredAt time.Time
}
type TaskDetail struct {
@@ -27,12 +27,13 @@ func TestRunnerSupportsUpStatusDownAndIdempotentUp(t *testing.T) {
if err != nil {
t.Fatalf("Up() error = %v", err)
}
if applied != 2 {
t.Fatalf("Up() applied = %d, want 2", applied)
if applied != 3 {
t.Fatalf("Up() applied = %d, want 3", applied)
}
assertStatuses(t, runner, map[int64]bool{
1: true,
2: true,
3: true,
})
applied, err = runner.Up(context.Background())
@@ -48,7 +49,8 @@ func TestRunnerSupportsUpStatusDownAndIdempotentUp(t *testing.T) {
}
assertStatuses(t, runner, map[int64]bool{
1: true,
2: false,
2: true,
3: false,
})
applied, err = runner.Up(context.Background())
@@ -61,6 +63,7 @@ func TestRunnerSupportsUpStatusDownAndIdempotentUp(t *testing.T) {
assertStatuses(t, runner, map[int64]bool{
1: true,
2: true,
3: true,
})
}
@@ -0,0 +1,69 @@
package password
import (
"errors"
"unicode/utf8"
"cmroubao/backend-api/internal/domain"
"golang.org/x/crypto/bcrypt"
)
var (
ErrInvalidPassword = errors.New("password is invalid")
ErrPasswordMismatch = errors.New("password does not match")
)
type Bcrypt struct {
cost int
dummyHash string
}
func NewBcrypt(cost int) (*Bcrypt, error) {
if cost < bcrypt.MinCost || cost > bcrypt.MaxCost {
return nil, errors.New("bcrypt cost is out of range")
}
dummyHash, err := bcrypt.GenerateFromPassword(
[]byte("cmroubao-dummy-password"),
cost,
)
if err != nil {
return nil, errors.New("initialize bcrypt dummy hash")
}
return &Bcrypt{cost: cost, dummyHash: string(dummyHash)}, nil
}
func (manager *Bcrypt) Hash(plain string) (string, error) {
if !utf8.ValidString(plain) ||
len([]byte(plain)) < domain.MinPasswordBytes ||
len([]byte(plain)) > domain.MaxPasswordBytes {
return "", ErrInvalidPassword
}
value, err := bcrypt.GenerateFromPassword([]byte(plain), manager.cost)
if err != nil {
return "", errors.New("hash password")
}
return string(value), nil
}
func (manager *Bcrypt) VerifyDummy(plain string) {
_ = bcrypt.CompareHashAndPassword(
[]byte(manager.dummyHash),
[]byte(plain),
)
}
func (manager *Bcrypt) Verify(encoded string, plain string) error {
if encoded == "" ||
plain == "" ||
len([]byte(plain)) > domain.MaxPasswordBytes {
return ErrPasswordMismatch
}
if err := bcrypt.CompareHashAndPassword(
[]byte(encoded),
[]byte(plain),
); err != nil {
return ErrPasswordMismatch
}
return nil
}
@@ -0,0 +1,67 @@
package password
import (
"errors"
"strings"
"testing"
"golang.org/x/crypto/bcrypt"
)
func TestBcryptHashesAndVerifiesWithoutStoringPlaintext(t *testing.T) {
manager, err := NewBcrypt(bcrypt.MinCost)
if err != nil {
t.Fatalf("NewBcrypt() error = %v", err)
}
encoded, err := manager.Hash("correct horse battery staple")
if err != nil {
t.Fatalf("Hash() error = %v", err)
}
if encoded == "correct horse battery staple" {
t.Fatal("Hash() returned plaintext")
}
if err := manager.Verify(encoded, "correct horse battery staple"); err != nil {
t.Fatalf("Verify(correct) error = %v", err)
}
if err := manager.Verify(encoded, "wrong"); !errors.Is(
err,
ErrPasswordMismatch,
) {
t.Fatalf("Verify(wrong) error = %v", err)
}
}
func TestBcryptRejectsInvalidCostAndOversizedPasswords(t *testing.T) {
if _, err := NewBcrypt(bcrypt.MinCost - 1); err == nil {
t.Fatal("NewBcrypt(invalid) error = nil")
}
manager, err := NewBcrypt(bcrypt.MinCost)
if err != nil {
t.Fatalf("NewBcrypt() error = %v", err)
}
oversized := strings.Repeat("x", 73)
if _, err := manager.Hash("too-short"); !errors.Is(
err,
ErrInvalidPassword,
) {
t.Fatalf("Hash(short) error = %v", err)
}
if _, err := manager.Hash(oversized); !errors.Is(
err,
ErrInvalidPassword,
) {
t.Fatalf("Hash(oversized) error = %v", err)
}
if _, err := manager.Hash(
string([]byte{0xff, 0xfe, 0xfd}),
); !errors.Is(err, ErrInvalidPassword) {
t.Fatalf("Hash(invalid UTF-8) error = %v", err)
}
if err := manager.Verify("$2a$04$invalid", oversized); !errors.Is(
err,
ErrPasswordMismatch,
) {
t.Fatalf("Verify(oversized) error = %v", err)
}
manager.VerifyDummy("password-123")
}
@@ -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),
@@ -0,0 +1,131 @@
package authcommon
import (
"errors"
"math"
"net"
"strings"
"sync"
"time"
)
type AttemptLimiter interface {
Allow(string) (bool, time.Duration)
Reset(string)
}
type attemptWindow struct {
count int
resetAt time.Time
}
type InMemoryAttemptLimiter struct {
mu sync.Mutex
maxAttempts int
window time.Duration
maxKeys int
now func() time.Time
attempts map[string]attemptWindow
}
func NewAttemptLimiter(
maxAttempts int,
window time.Duration,
maxKeys int,
) (*InMemoryAttemptLimiter, error) {
return newAttemptLimiter(maxAttempts, window, maxKeys, time.Now)
}
func newAttemptLimiter(
maxAttempts int,
window time.Duration,
maxKeys int,
now func() time.Time,
) (*InMemoryAttemptLimiter, error) {
if maxAttempts < 1 || window <= 0 || maxKeys < 1 || now == nil {
return nil, errors.New("attempt limiter configuration is invalid")
}
return &InMemoryAttemptLimiter{
maxAttempts: maxAttempts,
window: window,
maxKeys: maxKeys,
now: now,
attempts: make(map[string]attemptWindow),
}, nil
}
func (limiter *InMemoryAttemptLimiter) Allow(
key string,
) (bool, time.Duration) {
now := limiter.now().UTC()
key = strings.TrimSpace(key)
if key == "" {
key = "unknown"
}
limiter.mu.Lock()
defer limiter.mu.Unlock()
current, found := limiter.attempts[key]
if found && !current.resetAt.After(now) {
delete(limiter.attempts, key)
found = false
}
if !found {
if len(limiter.attempts) >= limiter.maxKeys {
limiter.removeExpired(now)
}
if len(limiter.attempts) >= limiter.maxKeys {
return false, limiter.window
}
limiter.attempts[key] = attemptWindow{
count: 1,
resetAt: now.Add(limiter.window),
}
return true, 0
}
if current.count >= limiter.maxAttempts {
return false, current.resetAt.Sub(now)
}
current.count++
limiter.attempts[key] = current
return true, 0
}
func (limiter *InMemoryAttemptLimiter) Reset(key string) {
limiter.mu.Lock()
delete(limiter.attempts, strings.TrimSpace(key))
limiter.mu.Unlock()
}
func (limiter *InMemoryAttemptLimiter) removeExpired(now time.Time) {
for key, current := range limiter.attempts {
if !current.resetAt.After(now) {
delete(limiter.attempts, key)
}
}
}
func LoginAttemptKey(scope, remoteAddress string) string {
host, _, err := net.SplitHostPort(strings.TrimSpace(remoteAddress))
if err != nil {
host = strings.TrimSpace(remoteAddress)
}
if parsed := net.ParseIP(host); parsed != nil {
host = parsed.String()
}
if host == "" {
host = "unknown"
}
return scope + ":" + host
}
func RetryAfterSeconds(wait time.Duration) int {
seconds := int(math.Ceil(wait.Seconds()))
if seconds < 1 {
return 1
}
return seconds
}
var _ AttemptLimiter = (*InMemoryAttemptLimiter)(nil)
@@ -0,0 +1,75 @@
package authcommon
import (
"testing"
"time"
)
func TestAttemptLimiterBlocksUntilWindowExpiresAndCanReset(t *testing.T) {
now := time.Date(2026, 7, 26, 1, 2, 3, 0, time.UTC)
limiter, err := newAttemptLimiter(
2,
time.Minute,
10,
func() time.Time { return now },
)
if err != nil {
t.Fatalf("newAttemptLimiter() error = %v", err)
}
if allowed, _ := limiter.Allow("admin:127.0.0.1"); !allowed {
t.Fatal("first attempt was blocked")
}
if allowed, _ := limiter.Allow("admin:127.0.0.1"); !allowed {
t.Fatal("second attempt was blocked")
}
if allowed, wait := limiter.Allow("admin:127.0.0.1"); allowed ||
wait != time.Minute {
t.Fatalf("third attempt = %v, wait = %v", allowed, wait)
}
limiter.Reset("admin:127.0.0.1")
if allowed, _ := limiter.Allow("admin:127.0.0.1"); !allowed {
t.Fatal("attempt after reset was blocked")
}
now = now.Add(time.Minute)
if allowed, _ := limiter.Allow("admin:127.0.0.1"); !allowed {
t.Fatal("attempt after window expiry was blocked")
}
}
func TestAttemptLimiterBoundsTrackedKeys(t *testing.T) {
now := time.Date(2026, 7, 26, 1, 2, 3, 0, time.UTC)
limiter, err := newAttemptLimiter(
1,
time.Minute,
1,
func() time.Time { return now },
)
if err != nil {
t.Fatalf("newAttemptLimiter() error = %v", err)
}
if allowed, _ := limiter.Allow("one"); !allowed {
t.Fatal("first key was blocked")
}
if allowed, wait := limiter.Allow("two"); allowed ||
wait != time.Minute {
t.Fatalf("second key = %v, wait = %v", allowed, wait)
}
now = now.Add(time.Minute)
if allowed, _ := limiter.Allow("two"); !allowed {
t.Fatal("expired key was not evicted")
}
}
func TestLoginAttemptKeyUsesRemoteAddressOnly(t *testing.T) {
if got := LoginAttemptKey("buyer", "127.0.0.1:1234"); got != "buyer:127.0.0.1" {
t.Fatalf("IPv4 key = %q", got)
}
if got := LoginAttemptKey("admin", "[2001:db8::1]:443"); got != "admin:2001:db8::1" {
t.Fatalf("IPv6 key = %q", got)
}
if got := RetryAfterSeconds(time.Millisecond); got != 1 {
t.Fatalf("RetryAfterSeconds() = %d", got)
}
}
@@ -0,0 +1,49 @@
package authcommon
import (
"context"
"crypto/subtle"
"encoding/base64"
"strings"
"cmroubao/backend-api/internal/domain"
)
const (
AdminSessionCookieName = "cmroubao_admin_session"
CSRFCookieName = "cmroubao_admin_csrf"
CSRFFormField = "csrf_token"
CSRFHeader = "X-CSRF-Token"
)
type principalContextKey struct{}
func WithPrincipal(
ctx context.Context,
principal domain.AuthPrincipal,
) context.Context {
return context.WithValue(ctx, principalContextKey{}, principal)
}
func Principal(ctx context.Context) (domain.AuthPrincipal, bool) {
principal, ok := ctx.Value(principalContextKey{}).(domain.AuthPrincipal)
return principal, ok
}
func ValidCSRFPair(cookieValue, presentedValue string) bool {
cookieValue = strings.TrimSpace(cookieValue)
presentedValue = strings.TrimSpace(presentedValue)
if !ValidOpaqueValue(cookieValue) ||
len(cookieValue) != len(presentedValue) {
return false
}
return subtle.ConstantTimeCompare(
[]byte(cookieValue),
[]byte(presentedValue),
) == 1
}
func ValidOpaqueValue(value string) bool {
decoded, err := base64.RawURLEncoding.DecodeString(value)
return err == nil && len(decoded) == 32
}
@@ -0,0 +1,41 @@
package authcommon
import (
"context"
"testing"
"cmroubao/backend-api/internal/domain"
)
func TestPrincipalRoundTrip(t *testing.T) {
want := domain.AuthPrincipal{
UserID: "user-1",
Username: "buyer",
Role: domain.UserRoleBuyer,
DeviceID: "device-1",
}
ctx := WithPrincipal(context.Background(), want)
got, ok := Principal(ctx)
if !ok || got != want {
t.Fatalf("Principal() = %+v, %t", got, ok)
}
}
func TestValidCSRFPairRequiresOneCanonical256BitValue(t *testing.T) {
valid := "YWFhYWFhYWFhYWFhYWFhYWFhYWFhYWFhYWFhYWFhYWE"
if !ValidCSRFPair(valid, valid) {
t.Fatal("valid CSRF pair was rejected")
}
for _, candidate := range []string{
"",
"short",
valid + "=",
valid[:len(valid)-1] + "b",
} {
if ValidCSRFPair(valid, candidate) {
t.Fatalf("invalid presented value %q was accepted", candidate)
}
}
}
@@ -1,36 +0,0 @@
package httpapi
import (
"net"
"net/http"
"github.com/gin-gonic/gin"
)
func loopbackAdminOnly() gin.HandlerFunc {
return func(ctx *gin.Context) {
host, _, err := net.SplitHostPort(ctx.Request.RemoteAddr)
if err != nil {
denyNonLocalAdmin(ctx)
return
}
address := net.ParseIP(host)
if address == nil || !address.IsLoopback() {
denyNonLocalAdmin(ctx)
return
}
ctx.Next()
}
}
func denyNonLocalAdmin(ctx *gin.Context) {
ctx.Header("Cache-Control", "no-store")
ctx.AbortWithStatusJSON(
http.StatusForbidden,
errorResponse(
ctx,
"ADMIN_SESSION_REQUIRED",
"admin session required",
),
)
}
@@ -1,64 +0,0 @@
package httpapi
import (
"net/http"
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
)
func TestLoopbackAdminOnlyAllowsLoopbackAddresses(t *testing.T) {
for _, remoteAddress := range []string{
"127.0.0.1:12345",
"[::1]:12345",
} {
t.Run(remoteAddress, func(t *testing.T) {
router := gin.New()
router.Use(requestIDMiddleware(), loopbackAdminOnly())
router.GET("/tasks", func(ctx *gin.Context) {
ctx.Status(http.StatusNoContent)
})
request := httptest.NewRequest(http.MethodGet, "/tasks", nil)
request.RemoteAddr = remoteAddress
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != http.StatusNoContent {
t.Fatalf("status = %d", response.Code)
}
})
}
}
func TestLoopbackAdminOnlyRejectsRemoteOrMalformedAddresses(t *testing.T) {
for _, remoteAddress := range []string{
"192.0.2.1:12345",
"not-an-address",
} {
t.Run(remoteAddress, func(t *testing.T) {
router := gin.New()
router.Use(requestIDMiddleware(), loopbackAdminOnly())
router.GET("/tasks", func(ctx *gin.Context) {
ctx.Status(http.StatusNoContent)
})
request := httptest.NewRequest(http.MethodGet, "/tasks", nil)
request.RemoteAddr = remoteAddress
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != http.StatusForbidden {
t.Fatalf("status = %d", response.Code)
}
assertErrorCode(t, response, "ADMIN_SESSION_REQUIRED")
if response.Header().Get("Cache-Control") != "no-store" {
t.Fatalf(
"Cache-Control = %q",
response.Header().Get("Cache-Control"),
)
}
})
}
}
@@ -11,6 +11,7 @@ import (
"time"
"cmroubao/backend-api/internal/domain"
"cmroubao/backend-api/internal/transport/authcommon"
"cmroubao/backend-api/internal/usecase"
"github.com/gin-gonic/gin"
@@ -209,6 +210,7 @@ func (h *adminHandlers) createTask(ctx *gin.Context) {
ctx.Request.Context(),
usecase.CreateTaskCommand{
CreatorSubject: localAdminSubject,
ActorUserID: adminActorUserID(ctx),
IdempotencyKey: ctx.GetHeader("Idempotency-Key"),
SourceRef: request.SourceRef,
Title: request.Title,
@@ -226,6 +228,14 @@ func (h *adminHandlers) createTask(ctx *gin.Context) {
ctx.JSON(http.StatusCreated, taskSummaryResponse(result.Task))
}
func adminActorUserID(ctx *gin.Context) string {
principal, ok := authcommon.Principal(ctx.Request.Context())
if !ok || principal.Role != domain.UserRoleAdmin {
return ""
}
return principal.UserID
}
func (h *adminHandlers) listTasks(ctx *gin.Context) {
query := usecase.ListTasksQuery{
CreatorSubject: localAdminSubject,
@@ -292,10 +302,11 @@ func (h *adminHandlers) taskDetail(ctx *gin.Context) {
events := make([]gin.H, 0, len(detail.Events))
for _, event := range detail.Events {
events = append(events, gin.H{
"id": event.ID,
"type": event.Type,
"message": event.Message,
"occurred_at": formatTime(event.OccurredAt),
"id": event.ID,
"actor_user_id": event.ActorUserID,
"type": event.Type,
"message": event.Message,
"occurred_at": formatTime(event.OccurredAt),
})
}
ctx.Header("Cache-Control", "no-store")
@@ -355,6 +366,7 @@ func (h *adminHandlers) cancelTask(ctx *gin.Context) {
ctx.Request.Context(),
usecase.CancelTaskCommand{
CreatorSubject: localAdminSubject,
ActorUserID: adminActorUserID(ctx),
TaskID: ctx.Param("id"),
Reason: request.Reason,
},
@@ -15,6 +15,7 @@ import (
"path/filepath"
"strings"
"testing"
"time"
"cmroubao/backend-api/internal/platform/assetstore"
"cmroubao/backend-api/internal/platform/database"
@@ -210,6 +211,32 @@ func TestAdminAPIAssetAndTaskLifecycle(t *testing.T) {
canceled,
)
}
canceledDetailResponse := performAdminRequest(
t,
router,
http.MethodGet,
"/api/v1/tasks/"+taskID,
"",
nil,
"",
)
var canceledDetail map[string]any
decodeResponse(t, canceledDetailResponse, &canceledDetail)
events, _ := canceledDetail["events"].([]any)
if canceledDetailResponse.Code != http.StatusOK || len(events) != 2 {
t.Fatalf(
"canceled detail status/body = %d / %#v",
canceledDetailResponse.Code,
canceledDetail,
)
}
for _, value := range events {
event, _ := value.(map[string]any)
if event["actor_user_id"] !=
"00000000-0000-4000-8000-000000000099" {
t.Fatalf("event actor = %#v", event)
}
}
secondCancel := performAdminRequest(
t,
@@ -229,7 +256,7 @@ func TestAdminAPIAssetAndTaskLifecycle(t *testing.T) {
}
}
func TestAdminRoutesRejectNonLoopbackRequests(t *testing.T) {
func TestAdminRoutesRejectRequestsWithoutAdminSession(t *testing.T) {
router := newAdminIntegrationRouter(t)
request := httptest.NewRequest(http.MethodGet, "/api/v1/tasks", nil)
request.RemoteAddr = "192.0.2.10:3210"
@@ -237,7 +264,7 @@ func TestAdminRoutesRejectNonLoopbackRequests(t *testing.T) {
router.ServeHTTP(response, request)
if response.Code != http.StatusForbidden {
if response.Code != http.StatusUnauthorized {
t.Fatalf("status = %d, body = %s", response.Code, response.Body.String())
}
var body map[string]any
@@ -275,7 +302,7 @@ func TestAdminAssetUploadRequiresIdempotencyKey(t *testing.T) {
type emptyAdminWeb struct{}
func (emptyAdminWeb) Register(gin.IRoutes) {}
func (emptyAdminWeb) RegisterProtected(gin.IRoutes) {}
func newAdminIntegrationRouter(t *testing.T) http.Handler {
t.Helper()
@@ -292,6 +319,18 @@ func newAdminIntegrationRouter(t *testing.T) http.Handler {
if _, err := runner.Up(ctx); err != nil {
t.Fatalf("migration.Up() error = %v", err)
}
now := time.Now().UTC().Format(time.RFC3339Nano)
if _, err := db.ExecContext(
ctx,
`INSERT INTO users (
id, username, password_hash, role, is_active, created_at, updated_at
) VALUES (?, 'admin', 'test-only-hash', 'ADMIN', 1, ?, ?)`,
"00000000-0000-4000-8000-000000000099",
now,
now,
); err != nil {
t.Fatalf("seed admin user: %v", err)
}
repositories, err := repository.New(db)
if err != nil {
t.Fatalf("repository.New() error = %v", err)
@@ -318,9 +357,11 @@ func newAdminIntegrationRouter(t *testing.T) http.Handler {
t.Fatalf("NewAdminRouteRegistrar() error = %v", err)
}
router, err := NewRouter(RouterDependencies{
Database: db,
RegisterAdminRoutes: registrar,
LogEvent: discardEvent,
Database: db,
RegisterPublicRoutes: discardRoutes,
RegisterAdminRoutes: registrar,
AdminSessions: allowAdminAuthenticator{},
LogEvent: discardEvent,
})
if err != nil {
t.Fatalf("NewRouter() error = %v", err)
@@ -375,7 +416,18 @@ func performAdminRequest(
) *httptest.ResponseRecorder {
t.Helper()
request := httptest.NewRequest(method, target, body)
request.RemoteAddr = "127.0.0.1:3210"
request.AddCookie(&http.Cookie{
Name: "cmroubao_admin_session",
Value: "test-session",
})
if method != http.MethodGet && method != http.MethodHead {
const csrfToken = "YWFhYWFhYWFhYWFhYWFhYWFhYWFhYWFhYWFhYWFhYWE"
request.AddCookie(&http.Cookie{
Name: "cmroubao_admin_csrf",
Value: csrfToken,
})
request.Header.Set("X-CSRF-Token", csrfToken)
}
if contentType != "" {
request.Header.Set("Content-Type", contentType)
}
@@ -0,0 +1,355 @@
package httpapi
import (
"context"
"errors"
"net/http"
"net/url"
"strconv"
"strings"
"cmroubao/backend-api/internal/domain"
"cmroubao/backend-api/internal/transport/authcommon"
"cmroubao/backend-api/internal/usecase"
"github.com/gin-gonic/gin"
)
type AdminAuthenticator interface {
AuthenticateAdmin(
context.Context,
string,
) (domain.AuthPrincipal, error)
}
type DeviceAuthenticator interface {
AuthenticateAccessToken(
context.Context,
string,
) (domain.AuthPrincipal, error)
}
type BuyerTokenService interface {
LoginBuyerDevice(
context.Context,
usecase.LoginBuyerDeviceCommand,
) (usecase.AccessTokenResult, error)
}
func NewPublicAuthRegistrar(
service BuyerTokenService,
limiter authcommon.AttemptLimiter,
) (RouteRegistrar, error) {
if service == nil {
return nil, errors.New("buyer token service is required")
}
if limiter == nil {
return nil, errors.New("buyer login limiter is required")
}
handler := &authHandlers{service: service, limiter: limiter}
return func(routes gin.IRoutes) error {
routes.POST("/api/v1/auth/token", handler.issueBuyerToken)
return nil
}, nil
}
type authHandlers struct {
service BuyerTokenService
limiter authcommon.AttemptLimiter
}
func (h *authHandlers) issueBuyerToken(ctx *gin.Context) {
if !hasMediaType(ctx, "application/json") {
writePublicError(
ctx,
http.StatusUnsupportedMediaType,
"UNSUPPORTED_MEDIA_TYPE",
"application/json is required",
false,
gin.H{},
)
return
}
var request struct {
Username string `json:"username"`
Password string `json:"password"`
DeviceID string `json:"device_id"`
DeviceToken string `json:"device_token"`
AppVersion string `json:"app_version"`
AndroidVersion string `json:"android_version"`
}
if err := decodeJSON(ctx, &request); err != nil {
writePublicError(
ctx,
http.StatusBadRequest,
"INVALID_JSON",
"request body must be valid JSON",
false,
gin.H{},
)
return
}
attemptKey := authcommon.LoginAttemptKey(
"buyer",
ctx.Request.RemoteAddr,
)
if allowed, wait := h.limiter.Allow(attemptKey); !allowed {
ctx.Header(
"Retry-After",
strconv.Itoa(authcommon.RetryAfterSeconds(wait)),
)
writePublicError(
ctx,
http.StatusTooManyRequests,
"AUTH_RATE_LIMITED",
"too many authentication attempts",
true,
gin.H{},
)
return
}
result, err := h.service.LoginBuyerDevice(
ctx.Request.Context(),
usecase.LoginBuyerDeviceCommand{
Username: request.Username,
Password: request.Password,
DeviceID: request.DeviceID,
DeviceToken: request.DeviceToken,
AppVersion: request.AppVersion,
AndroidVersion: request.AndroidVersion,
},
)
if err != nil {
writeAuthError(ctx, err)
return
}
h.limiter.Reset(attemptKey)
ctx.Header("Cache-Control", "no-store")
ctx.JSON(http.StatusOK, gin.H{
"access_token": result.Token,
"token_type": "Bearer",
"expires_in": int(usecase.AccessTokenLifetime.Seconds()),
"user": gin.H{
"id": result.User.ID,
"username": result.User.Username,
"role": result.User.Role,
},
"device": gin.H{
"id": result.Device.ID,
"enabled": result.Device.IsEnabled,
},
})
}
func requireAdminSession(
authenticator AdminAuthenticator,
) gin.HandlerFunc {
return func(ctx *gin.Context) {
cookie, err := ctx.Request.Cookie(
authcommon.AdminSessionCookieName,
)
if err != nil || cookie.Value == "" {
denyAdminSession(ctx)
return
}
principal, err := authenticator.AuthenticateAdmin(
ctx.Request.Context(),
cookie.Value,
)
if err != nil ||
principal.Role != domain.UserRoleAdmin ||
principal.DeviceID != "" {
denyAdminSession(ctx)
return
}
ctx.Request = ctx.Request.WithContext(
authcommon.WithPrincipal(
ctx.Request.Context(),
principal,
),
)
if isUnsafeAdminAPIRequest(ctx.Request) &&
!validAPIRequestCSRF(ctx.Request) {
writePublicError(
ctx,
http.StatusForbidden,
"CSRF_INVALID",
"CSRF token is invalid",
false,
gin.H{},
)
ctx.Abort()
return
}
ctx.Next()
}
}
func RequireDeviceAccess(
authenticator DeviceAuthenticator,
) gin.HandlerFunc {
return func(ctx *gin.Context) {
token, ok := bearerToken(ctx.GetHeader("Authorization"))
if !ok {
denyDeviceAccess(ctx)
return
}
principal, err := authenticator.AuthenticateAccessToken(
ctx.Request.Context(),
token,
)
if err != nil ||
principal.Role != domain.UserRoleBuyer ||
principal.DeviceID == "" {
denyDeviceAccess(ctx)
return
}
ctx.Request = ctx.Request.WithContext(
authcommon.WithPrincipal(
ctx.Request.Context(),
principal,
),
)
ctx.Next()
}
}
func denyAdminSession(ctx *gin.Context) {
ctx.Header("Cache-Control", "no-store")
if strings.HasPrefix(ctx.Request.URL.Path, "/api/") {
ctx.Abort()
writePublicError(
ctx,
http.StatusUnauthorized,
"ADMIN_SESSION_REQUIRED",
"admin session required",
false,
gin.H{},
)
return
}
next := ctx.Request.URL.RequestURI()
if next == "" ||
(next != "/tasks" && !strings.HasPrefix(next, "/tasks?") &&
!strings.HasPrefix(next, "/tasks/")) {
next = "/tasks"
}
ctx.Abort()
ctx.Redirect(
http.StatusSeeOther,
"/login?next="+url.QueryEscape(next),
)
}
func denyDeviceAccess(ctx *gin.Context) {
ctx.Abort()
writePublicError(
ctx,
http.StatusUnauthorized,
"DEVICE_ACCESS_REQUIRED",
"device access token required",
false,
gin.H{},
)
}
func validAPIRequestCSRF(request *http.Request) bool {
cookie, err := request.Cookie(authcommon.CSRFCookieName)
if err != nil {
return false
}
return authcommon.ValidCSRFPair(
cookie.Value,
request.Header.Get(authcommon.CSRFHeader),
)
}
func isUnsafeAdminAPIRequest(request *http.Request) bool {
if !strings.HasPrefix(request.URL.Path, "/api/") {
return false
}
switch request.Method {
case http.MethodGet, http.MethodHead, http.MethodOptions:
return false
default:
return true
}
}
func bearerToken(value string) (string, bool) {
parts := strings.Fields(value)
if len(parts) != 2 ||
!strings.EqualFold(parts[0], "Bearer") ||
!authcommon.ValidOpaqueValue(parts[1]) {
return "", false
}
return parts[1], true
}
func writeAuthError(ctx *gin.Context, err error) {
var typed *usecase.Error
if !errors.As(err, &typed) {
writePublicError(
ctx,
http.StatusInternalServerError,
"INTERNAL_ERROR",
"internal server error",
false,
gin.H{},
)
return
}
switch typed.Kind {
case usecase.ErrorKindUnauthorized, usecase.ErrorKindInvalid:
writePublicError(
ctx,
http.StatusUnauthorized,
"AUTH_INVALID_CREDENTIALS",
"invalid credentials",
false,
gin.H{},
)
case usecase.ErrorKindForbidden:
code := "AUTH_FORBIDDEN"
message := "authentication is not allowed"
if typed.Code == "AUTH_ACCOUNT_OR_DEVICE_DISABLED" {
code = "AUTH_ACCOUNT_OR_DEVICE_DISABLED"
message = "account or device is disabled"
}
writePublicError(
ctx,
http.StatusForbidden,
code,
message,
false,
gin.H{},
)
case usecase.ErrorKindConflict:
writePublicError(
ctx,
http.StatusConflict,
"AUTH_RESOURCE_CONFLICT",
"authentication resource conflict",
false,
gin.H{},
)
case usecase.ErrorKindUnavailable:
writePublicError(
ctx,
http.StatusServiceUnavailable,
"AUTH_UNAVAILABLE",
"authentication service is unavailable",
true,
gin.H{},
)
default:
writePublicError(
ctx,
http.StatusInternalServerError,
"INTERNAL_ERROR",
"internal server error",
false,
gin.H{},
)
}
}
@@ -0,0 +1,405 @@
package httpapi
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"cmroubao/backend-api/internal/domain"
"cmroubao/backend-api/internal/transport/authcommon"
"cmroubao/backend-api/internal/usecase"
"github.com/gin-gonic/gin"
)
const (
testOpaqueToken = "YWFhYWFhYWFhYWFhYWFhYWFhYWFhYWFhYWFhYWFhYWE"
testCSRFOpaqueToken = "YmJiYmJiYmJiYmJiYmJiYmJiYmJiYmJiYmJiYmJiYmI"
)
type fakeBuyerTokenService struct {
command usecase.LoginBuyerDeviceCommand
result usecase.AccessTokenResult
err error
calls int
}
func (service *fakeBuyerTokenService) LoginBuyerDevice(
_ context.Context,
command usecase.LoginBuyerDeviceCommand,
) (usecase.AccessTokenResult, error) {
service.calls++
service.command = command
return service.result, service.err
}
func TestIssueBuyerTokenReturnsOnlyPublicIdentityAndAccessToken(t *testing.T) {
service := &fakeBuyerTokenService{
result: usecase.AccessTokenResult{
Token: testOpaqueToken,
ExpiresAt: time.Now().Add(time.Hour),
User: domain.User{
ID: "buyer-1",
Username: "buyer01",
Role: domain.UserRoleBuyer,
},
Device: domain.Device{
ID: "device-1",
IsEnabled: true,
},
},
}
router := newPublicAuthTestRouter(t, service)
body := `{
"username":"buyer01",
"password":"private-password",
"device_id":"device-1",
"device_token":"private-device-token",
"app_version":"0.1.0",
"android_version":"15"
}`
request := httptest.NewRequest(
http.MethodPost,
"/api/v1/auth/token",
strings.NewReader(body),
)
request.Header.Set("Content-Type", "application/json")
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != http.StatusOK {
t.Fatalf("status/body = %d / %s", response.Code, response.Body)
}
if service.command.DeviceID != "device-1" ||
service.command.AppVersion != "0.1.0" ||
service.command.AndroidVersion != "15" ||
service.command.Password != "private-password" {
t.Fatalf("command = %+v", service.command)
}
var decoded map[string]any
if err := json.Unmarshal(response.Body.Bytes(), &decoded); err != nil {
t.Fatalf("decode response: %v", err)
}
if decoded["access_token"] != testOpaqueToken ||
decoded["token_type"] != "Bearer" ||
decoded["expires_in"] != float64(3600) {
t.Fatalf("response = %#v", decoded)
}
responseText := response.Body.String()
for _, secret := range []string{
"private-password",
"private-device-token",
"android_version",
"app_version",
} {
if strings.Contains(responseText, secret) {
t.Fatalf("response leaked %q: %s", secret, responseText)
}
}
if response.Header().Get("Cache-Control") != "no-store" {
t.Fatalf("Cache-Control = %q", response.Header().Get("Cache-Control"))
}
}
func TestIssueBuyerTokenUsesGenericCredentialErrors(t *testing.T) {
service := &fakeBuyerTokenService{
err: &usecase.Error{
Kind: usecase.ErrorKindInvalid,
Code: "AUTH_VALIDATION_FAILED",
Message: "private validation detail",
Fields: map[string]string{"password": "private detail"},
},
}
router := newPublicAuthTestRouter(t, service)
request := httptest.NewRequest(
http.MethodPost,
"/api/v1/auth/token",
strings.NewReader(`{
"username":"buyer",
"password":"secret-value",
"device_id":"device",
"device_token":"token",
"app_version":"0.1",
"android_version":"15"
}`),
)
request.Header.Set("Content-Type", "application/json")
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != http.StatusUnauthorized {
t.Fatalf("status/body = %d / %s", response.Code, response.Body)
}
assertErrorCode(t, response, "AUTH_INVALID_CREDENTIALS")
for _, privateValue := range []string{
"secret-value",
"private validation detail",
"private detail",
} {
if strings.Contains(response.Body.String(), privateValue) {
t.Fatalf("error leaked %q", privateValue)
}
}
}
func TestIssueBuyerTokenRateLimitSkipsAuthenticationWork(t *testing.T) {
service := &fakeBuyerTokenService{
err: &usecase.Error{
Kind: usecase.ErrorKindUnauthorized,
Code: "AUTH_INVALID_CREDENTIALS",
},
}
limiter, err := authcommon.NewAttemptLimiter(1, time.Minute, 10)
if err != nil {
t.Fatalf("NewAttemptLimiter() error = %v", err)
}
registrar, err := NewPublicAuthRegistrar(service, limiter)
if err != nil {
t.Fatalf("NewPublicAuthRegistrar() error = %v", err)
}
router := gin.New()
router.Use(requestIDMiddleware())
if err := registrar(router); err != nil {
t.Fatalf("register public auth: %v", err)
}
body := `{
"username":"buyer",
"password":"secret-value",
"device_id":"device",
"device_token":"token",
"app_version":"0.1",
"android_version":"15"
}`
requestToken := func() *httptest.ResponseRecorder {
request := httptest.NewRequest(
http.MethodPost,
"/api/v1/auth/token",
strings.NewReader(body),
)
request.Header.Set("Content-Type", "application/json")
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
return response
}
if response := requestToken(); response.Code != http.StatusUnauthorized {
t.Fatalf("first status = %d", response.Code)
}
response := requestToken()
if response.Code != http.StatusTooManyRequests ||
response.Header().Get("Retry-After") == "" ||
!strings.Contains(response.Body.String(), `"code":"AUTH_RATE_LIMITED"`) ||
service.calls != 1 {
t.Fatalf(
"second status/retry/calls/body = %d / %q / %d / %s",
response.Code,
response.Header().Get("Retry-After"),
service.calls,
response.Body,
)
}
}
func TestAdminSessionMiddlewareSeparatesWebAPIAndCSRF(t *testing.T) {
authenticator := &stubAuthenticator{
adminPrincipal: domain.AuthPrincipal{
UserID: "admin-1",
Username: "admin",
Role: domain.UserRoleAdmin,
SessionID: "session-1",
},
}
router := gin.New()
router.Use(requestIDMiddleware())
protected := router.Group("")
protected.Use(requireAdminSession(authenticator))
protected.GET("/tasks/item", func(ctx *gin.Context) {
principal, ok := authcommon.Principal(ctx.Request.Context())
if !ok {
ctx.Status(http.StatusInternalServerError)
return
}
ctx.String(http.StatusOK, principal.UserID)
})
protected.POST("/api/v1/tasks", func(ctx *gin.Context) {
ctx.Status(http.StatusNoContent)
})
webRequest := httptest.NewRequest(
http.MethodGet,
"/tasks/item?q=1",
nil,
)
webResponse := httptest.NewRecorder()
router.ServeHTTP(webResponse, webRequest)
if webResponse.Code != http.StatusSeeOther ||
webResponse.Header().Get("Location") !=
"/login?next=%2Ftasks%2Fitem%3Fq%3D1" {
t.Fatalf(
"web status/location = %d / %q",
webResponse.Code,
webResponse.Header().Get("Location"),
)
}
apiRequest := httptest.NewRequest(
http.MethodPost,
"/api/v1/tasks",
nil,
)
apiRequest.Header.Set("Authorization", "Bearer "+testOpaqueToken)
apiResponse := httptest.NewRecorder()
router.ServeHTTP(apiResponse, apiRequest)
if apiResponse.Code != http.StatusUnauthorized {
t.Fatalf("API status/body = %d / %s", apiResponse.Code, apiResponse.Body)
}
assertErrorCode(t, apiResponse, "ADMIN_SESSION_REQUIRED")
badCSRF := httptest.NewRequest(
http.MethodPost,
"/api/v1/tasks",
nil,
)
badCSRF.AddCookie(&http.Cookie{
Name: authcommon.AdminSessionCookieName,
Value: testOpaqueToken,
})
badCSRFResponse := httptest.NewRecorder()
router.ServeHTTP(badCSRFResponse, badCSRF)
if badCSRFResponse.Code != http.StatusForbidden {
t.Fatalf(
"CSRF status/body = %d / %s",
badCSRFResponse.Code,
badCSRFResponse.Body,
)
}
assertErrorCode(t, badCSRFResponse, "CSRF_INVALID")
goodRequest := httptest.NewRequest(
http.MethodPost,
"/api/v1/tasks",
nil,
)
goodRequest.AddCookie(&http.Cookie{
Name: authcommon.AdminSessionCookieName,
Value: testOpaqueToken,
})
goodRequest.AddCookie(&http.Cookie{
Name: authcommon.CSRFCookieName,
Value: testCSRFOpaqueToken,
})
goodRequest.Header.Set(authcommon.CSRFHeader, testCSRFOpaqueToken)
goodResponse := httptest.NewRecorder()
router.ServeHTTP(goodResponse, goodRequest)
if goodResponse.Code != http.StatusNoContent {
t.Fatalf("authenticated status = %d", goodResponse.Code)
}
}
func TestDeviceMiddlewareRejectsCookieAndAcceptsBuyerBearer(t *testing.T) {
authenticator := &stubAuthenticator{
devicePrincipal: domain.AuthPrincipal{
UserID: "buyer-1",
Username: "buyer",
Role: domain.UserRoleBuyer,
DeviceID: "device-1",
},
}
router := gin.New()
router.Use(requestIDMiddleware(), RequireDeviceAccess(authenticator))
router.GET("/api/v1/device-probe", func(ctx *gin.Context) {
principal, _ := authcommon.Principal(ctx.Request.Context())
ctx.String(http.StatusOK, principal.DeviceID)
})
cookieRequest := httptest.NewRequest(
http.MethodGet,
"/api/v1/device-probe",
nil,
)
cookieRequest.AddCookie(&http.Cookie{
Name: authcommon.AdminSessionCookieName,
Value: testOpaqueToken,
})
cookieResponse := httptest.NewRecorder()
router.ServeHTTP(cookieResponse, cookieRequest)
if cookieResponse.Code != http.StatusUnauthorized {
t.Fatalf("cookie status = %d", cookieResponse.Code)
}
assertErrorCode(t, cookieResponse, "DEVICE_ACCESS_REQUIRED")
bearerRequest := httptest.NewRequest(
http.MethodGet,
"/api/v1/device-probe",
nil,
)
bearerRequest.Header.Set("Authorization", "Bearer "+testOpaqueToken)
bearerResponse := httptest.NewRecorder()
router.ServeHTTP(bearerResponse, bearerRequest)
if bearerResponse.Code != http.StatusOK ||
bearerResponse.Body.String() != "device-1" {
t.Fatalf(
"bearer status/body = %d / %q",
bearerResponse.Code,
bearerResponse.Body.String(),
)
}
}
type stubAuthenticator struct {
adminPrincipal domain.AuthPrincipal
adminErr error
devicePrincipal domain.AuthPrincipal
deviceErr error
}
func (auth *stubAuthenticator) AuthenticateAdmin(
context.Context,
string,
) (domain.AuthPrincipal, error) {
return auth.adminPrincipal, auth.adminErr
}
func (auth *stubAuthenticator) AuthenticateAccessToken(
context.Context,
string,
) (domain.AuthPrincipal, error) {
return auth.devicePrincipal, auth.deviceErr
}
func newPublicAuthTestRouter(
t *testing.T,
service BuyerTokenService,
) http.Handler {
t.Helper()
limiter, err := authcommon.NewAttemptLimiter(
100,
time.Minute,
100,
)
if err != nil {
t.Fatalf("NewAttemptLimiter() error = %v", err)
}
registrar, err := NewPublicAuthRegistrar(service, limiter)
if err != nil {
t.Fatalf("NewPublicAuthRegistrar() error = %v", err)
}
router := gin.New()
router.Use(requestIDMiddleware())
if err := registrar(router); err != nil {
t.Fatalf("register public auth: %v", err)
}
return router
}
var (
_ BuyerTokenService = (*fakeBuyerTokenService)(nil)
_ AdminAuthenticator = (*stubAuthenticator)(nil)
_ DeviceAuthenticator = (*stubAuthenticator)(nil)
)
@@ -22,13 +22,15 @@ type EventLogger func(string)
type RouteRegistrar func(gin.IRoutes) error
type RouterDependencies struct {
Database DatabasePinger
RegisterAdminRoutes RouteRegistrar
LogEvent EventLogger
Database DatabasePinger
RegisterPublicRoutes RouteRegistrar
RegisterAdminRoutes RouteRegistrar
AdminSessions AdminAuthenticator
LogEvent EventLogger
}
type AdminWeb interface {
Register(gin.IRoutes)
RegisterProtected(gin.IRoutes)
}
func NewAdminRouteRegistrar(
@@ -45,7 +47,7 @@ func NewAdminRouteRegistrar(
if err := registerAdminAPI(routes, services); err != nil {
return err
}
web.Register(routes)
web.RegisterProtected(routes)
return nil
}, nil
}
@@ -57,6 +59,12 @@ func NewRouter(dependencies RouterDependencies) (http.Handler, error) {
if dependencies.RegisterAdminRoutes == nil {
return nil, errors.New("admin route registrar is required")
}
if dependencies.RegisterPublicRoutes == nil {
return nil, errors.New("public route registrar is required")
}
if dependencies.AdminSessions == nil {
return nil, errors.New("admin authenticator is required")
}
if dependencies.LogEvent == nil {
return nil, errors.New("event logger is required")
}
@@ -70,8 +78,11 @@ func NewRouter(dependencies RouterDependencies) (http.Handler, error) {
}
router.GET("/healthz", healthHandler(dependencies.Database))
if err := dependencies.RegisterPublicRoutes(router); err != nil {
return nil, err
}
adminRoutes := router.Group("")
adminRoutes.Use(loopbackAdminOnly())
adminRoutes.Use(requireAdminSession(dependencies.AdminSessions))
if err := dependencies.RegisterAdminRoutes(adminRoutes); err != nil {
return nil, err
}
@@ -9,6 +9,9 @@ import (
"regexp"
"strings"
"testing"
"time"
"cmroubao/backend-api/internal/domain"
"github.com/gin-gonic/gin"
)
@@ -114,9 +117,11 @@ func TestSafeRecoveryReturnsStableErrorWithoutLoggingRequestHeaders(t *testing.T
func TestRouterRequiresDependencies(t *testing.T) {
valid := RouterDependencies{
Database: fakePinger{},
RegisterAdminRoutes: discardRoutes,
LogEvent: discardEvent,
Database: fakePinger{},
RegisterPublicRoutes: discardRoutes,
RegisterAdminRoutes: discardRoutes,
AdminSessions: allowAdminAuthenticator{},
LogEvent: discardEvent,
}
missingDatabase := valid
missingDatabase.Database = nil
@@ -128,6 +133,16 @@ func TestRouterRequiresDependencies(t *testing.T) {
if _, err := NewRouter(missingRoutes); err == nil {
t.Fatal("NewRouter(nil routes) error = nil")
}
missingPublicRoutes := valid
missingPublicRoutes.RegisterPublicRoutes = nil
if _, err := NewRouter(missingPublicRoutes); err == nil {
t.Fatal("NewRouter(nil public routes) error = nil")
}
missingAuth := valid
missingAuth.AdminSessions = nil
if _, err := NewRouter(missingAuth); err == nil {
t.Fatal("NewRouter(nil admin auth) error = nil")
}
missingLogger := valid
missingLogger.LogEvent = nil
if _, err := NewRouter(missingLogger); err == nil {
@@ -140,9 +155,11 @@ func newTestRouter(
logEvent EventLogger,
) (http.Handler, error) {
return NewRouter(RouterDependencies{
Database: database,
RegisterAdminRoutes: discardRoutes,
LogEvent: logEvent,
Database: database,
RegisterPublicRoutes: discardRoutes,
RegisterAdminRoutes: discardRoutes,
AdminSessions: allowAdminAuthenticator{},
LogEvent: logEvent,
})
}
@@ -223,6 +240,21 @@ func discardEvent(string) {}
func discardRoutes(gin.IRoutes) error { return nil }
type allowAdminAuthenticator struct{}
func (allowAdminAuthenticator) AuthenticateAdmin(
context.Context,
string,
) (domain.AuthPrincipal, error) {
return domain.AuthPrincipal{
UserID: "00000000-0000-4000-8000-000000000099",
Username: "admin",
Role: domain.UserRoleAdmin,
SessionID: "admin-session",
ExpiresAt: time.Now().Add(time.Hour),
}, nil
}
var requestIDPattern = regexp.MustCompile(
`^[0-9a-f]{8}-[0-9a-f]{4}-4[0-9a-f]{3}-[89ab][0-9a-f]{3}-[0-9a-f]{12}$`,
)
@@ -0,0 +1,75 @@
package webui
import (
"context"
"errors"
"cmroubao/backend-api/internal/usecase"
)
type AuthUsecaseAdapter struct {
service *usecase.AuthService
}
func NewAuthUsecaseAdapter(
service *usecase.AuthService,
) (*AuthUsecaseAdapter, error) {
if service == nil {
return nil, errors.New("admin auth use case is required")
}
return &AuthUsecaseAdapter{service: service}, nil
}
func (adapter *AuthUsecaseAdapter) LoginAdmin(
ctx context.Context,
input AdminLoginInput,
) (AdminLoginResult, error) {
result, err := adapter.service.LoginAdmin(
ctx,
usecase.LoginAdminCommand{
Username: input.Username,
Password: input.Password,
},
)
if err != nil {
return AdminLoginResult{}, mapAuthUsecaseError(err)
}
return AdminLoginResult{
Token: result.Token,
ExpiresAt: result.ExpiresAt,
}, nil
}
func (adapter *AuthUsecaseAdapter) LogoutAdmin(
ctx context.Context,
token string,
) error {
if err := adapter.service.LogoutAdmin(ctx, token); err != nil {
return mapAuthUsecaseError(err)
}
return nil
}
func mapAuthUsecaseError(err error) error {
var typed *usecase.Error
if !errors.As(err, &typed) {
return err
}
var public error
switch typed.Kind {
case usecase.ErrorKindUnauthorized,
usecase.ErrorKindForbidden,
usecase.ErrorKindInvalid:
public = ErrInvalidCredentials
case usecase.ErrorKindUnavailable:
public = ErrUnavailable
default:
return err
}
return &adapterError{
public: public,
cause: err,
}
}
var _ AdminSessionService = (*AuthUsecaseAdapter)(nil)
@@ -0,0 +1,311 @@
package webui
import (
"bytes"
"context"
"errors"
"net/http"
"net/url"
"path"
"strconv"
"strings"
"time"
"cmroubao/backend-api/internal/transport/authcommon"
"github.com/gin-gonic/gin"
)
const (
AdminSessionCookieName = authcommon.AdminSessionCookieName
maxLoginFormBytes = 16 << 10
)
var (
ErrInvalidCredentials = errors.New("invalid credentials")
ErrAccountDisabled = errors.New("account disabled")
)
type AdminSessionService interface {
LoginAdmin(context.Context, AdminLoginInput) (AdminLoginResult, error)
LogoutAdmin(context.Context, string) error
}
type AdminLoginInput struct {
Username string
Password string
}
type AdminLoginResult struct {
Token string
ExpiresAt time.Time
}
type AuthHandler struct {
sessions AdminSessionService
renderer *Renderer
limiter authcommon.AttemptLimiter
}
func NewAuthHandler(
sessions AdminSessionService,
renderer *Renderer,
limiter authcommon.AttemptLimiter,
) (*AuthHandler, error) {
if sessions == nil {
return nil, errors.New("admin session service is required")
}
if renderer == nil {
return nil, errors.New("admin auth renderer is required")
}
if limiter == nil {
return nil, errors.New("admin login limiter is required")
}
return &AuthHandler{
sessions: sessions,
renderer: renderer,
limiter: limiter,
}, nil
}
func (h *AuthHandler) RegisterPublic(routes gin.IRoutes) {
routes.GET("/login", SecurityHeaders(), h.LoginPage)
routes.POST("/login", SecurityHeaders(), h.Login)
}
func (h *AuthHandler) RegisterProtected(routes gin.IRoutes) {
routes.POST("/logout", SecurityHeaders(), h.Logout)
}
func (h *AuthHandler) LoginPage(ctx *gin.Context) {
token, err := csrfToken(ctx)
if err != nil {
h.renderError(ctx, http.StatusInternalServerError)
return
}
h.render(ctx, http.StatusOK, loginPage{
Page: pageView{Title: "管理端登录"},
CSRFToken: token,
Next: safeNext(ctx.Query("next")),
})
}
func (h *AuthHandler) Login(ctx *gin.Context) {
previousSessionToken := ""
if cookie, err := ctx.Request.Cookie(
authcommon.AdminSessionCookieName,
); err == nil {
previousSessionToken = cookie.Value
}
ctx.Request.Body = http.MaxBytesReader(
ctx.Writer,
ctx.Request.Body,
maxLoginFormBytes,
)
if err := ctx.Request.ParseForm(); err != nil {
h.renderError(ctx, http.StatusBadRequest)
return
}
next := safeNext(ctx.PostForm("next"))
page := loginPage{
Page: pageView{Title: "管理端登录"},
CSRFToken: strings.TrimSpace(
ctx.PostForm(authcommon.CSRFFormField),
),
Next: next,
Username: strings.TrimSpace(ctx.PostForm("username")),
}
if !validCSRF(ctx) {
page.Message = "登录页面已失效,请刷新后重试。"
h.render(ctx, http.StatusForbidden, page)
return
}
password := ctx.PostForm("password")
if page.Username == "" {
page.UsernameError = "请输入账号。"
}
if password == "" {
page.PasswordError = "请输入密码。"
}
if page.UsernameError != "" || page.PasswordError != "" {
page.Message = "请检查并补全必填项。"
h.render(ctx, http.StatusUnprocessableEntity, page)
return
}
attemptKey := authcommon.LoginAttemptKey(
"admin",
ctx.Request.RemoteAddr,
)
if allowed, wait := h.limiter.Allow(attemptKey); !allowed {
ctx.Header(
"Retry-After",
strconv.Itoa(authcommon.RetryAfterSeconds(wait)),
)
page.Message = "登录尝试次数过多,请稍后再试。"
h.render(ctx, http.StatusTooManyRequests, page)
return
}
result, err := h.sessions.LoginAdmin(
ctx.Request.Context(),
AdminLoginInput{
Username: page.Username,
Password: password,
},
)
if err != nil {
switch {
case errors.Is(err, ErrInvalidCredentials),
errors.Is(err, ErrAccountDisabled):
page.Message = "账号或密码不正确。"
h.render(ctx, http.StatusUnauthorized, page)
case errors.Is(err, context.DeadlineExceeded),
errors.Is(err, ErrUnavailable):
page.Message = "登录服务暂时不可用,请稍后重试。"
h.render(ctx, http.StatusServiceUnavailable, page)
default:
h.renderError(ctx, http.StatusInternalServerError)
}
return
}
h.limiter.Reset(attemptKey)
if previousSessionToken != "" &&
previousSessionToken != result.Token {
if err := h.sessions.LogoutAdmin(
ctx.Request.Context(),
previousSessionToken,
); err != nil {
_ = h.sessions.LogoutAdmin(
ctx.Request.Context(),
result.Token,
)
clearAdminSessionCookie(ctx)
h.renderError(ctx, http.StatusServiceUnavailable)
return
}
}
setAdminSessionCookie(ctx, result)
if _, err := rotateCSRFToken(ctx); err != nil {
_ = h.sessions.LogoutAdmin(ctx.Request.Context(), result.Token)
clearAdminSessionCookie(ctx)
h.renderError(ctx, http.StatusInternalServerError)
return
}
ctx.Redirect(http.StatusSeeOther, next)
}
func (h *AuthHandler) Logout(ctx *gin.Context) {
if !validCSRF(ctx) {
h.renderError(ctx, http.StatusForbidden)
return
}
cookie, err := ctx.Request.Cookie(authcommon.AdminSessionCookieName)
if err == nil && cookie.Value != "" {
if err := h.sessions.LogoutAdmin(
ctx.Request.Context(),
cookie.Value,
); err != nil &&
!errors.Is(err, ErrNotFound) {
h.renderError(ctx, http.StatusServiceUnavailable)
return
}
}
clearAdminSessionCookie(ctx)
_, _ = rotateCSRFToken(ctx)
ctx.Redirect(http.StatusSeeOther, "/login")
}
func setAdminSessionCookie(
ctx *gin.Context,
result AdminLoginResult,
) {
maxAge := int(time.Until(result.ExpiresAt).Seconds())
if maxAge < 1 {
maxAge = 1
}
http.SetCookie(ctx.Writer, &http.Cookie{
Name: authcommon.AdminSessionCookieName,
Value: result.Token,
Path: "/",
Expires: result.ExpiresAt.UTC(),
MaxAge: maxAge,
HttpOnly: true,
Secure: ctx.Request.TLS != nil,
SameSite: http.SameSiteLaxMode,
})
}
func clearAdminSessionCookie(ctx *gin.Context) {
http.SetCookie(ctx.Writer, &http.Cookie{
Name: authcommon.AdminSessionCookieName,
Value: "",
Path: "/",
Expires: time.Unix(1, 0).UTC(),
MaxAge: -1,
HttpOnly: true,
Secure: ctx.Request.TLS != nil,
SameSite: http.SameSiteLaxMode,
})
}
func safeNext(value string) string {
value = strings.TrimSpace(value)
if value == "" {
return "/tasks"
}
if strings.Contains(value, `\`) ||
strings.HasPrefix(value, "//") {
return "/tasks"
}
parsed, err := url.Parse(value)
if err != nil ||
parsed.IsAbs() ||
parsed.Host != "" ||
parsed.Fragment != "" ||
path.Clean(parsed.Path) != parsed.Path {
return "/tasks"
}
if parsed.Path != "/tasks" &&
!strings.HasPrefix(parsed.Path, "/tasks/") {
return "/tasks"
}
return parsed.String()
}
func (h *AuthHandler) render(
ctx *gin.Context,
status int,
page loginPage,
) {
var output bytes.Buffer
if err := h.renderer.Execute(&output, "login", page); err != nil {
ctx.Data(
http.StatusInternalServerError,
formContentType,
[]byte("页面暂时无法显示,请稍后重试。"),
)
return
}
ctx.Data(status, formContentType, output.Bytes())
}
func (h *AuthHandler) renderError(ctx *gin.Context, status int) {
token, _ := csrfToken(ctx)
h.render(ctx, status, loginPage{
Page: pageView{Title: "管理端登录"},
CSRFToken: token,
Next: "/tasks",
Message: "操作失败,请刷新页面后重试。",
})
}
type loginPage struct {
Page pageView
CSRFToken string
Next string
Username string
Message string
UsernameError string
PasswordError string
}
@@ -0,0 +1,443 @@
package webui
import (
"context"
"crypto/tls"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"testing"
"time"
"cmroubao/backend-api/internal/transport/authcommon"
"github.com/gin-gonic/gin"
)
func TestLoginPageIssuesCSRFAndRendersSafeNext(t *testing.T) {
router, _ := newAuthTestRouter(t, &fakeAdminSessions{})
response := performRequest(
t,
router,
http.MethodGet,
"/login?next=%2Ftasks%2Fnew",
nil,
"",
)
if response.Code != http.StatusOK {
t.Fatalf("status = %d, body = %s", response.Code, response.Body)
}
cookie := csrfCookie(t, response)
body := response.Body.String()
for _, expected := range []string{
`name="csrf_token" value="` + cookie.Value + `"`,
`name="next" value="/tasks/new"`,
`autocomplete="username"`,
`autocomplete="current-password"`,
`data-password-toggle`,
} {
if !strings.Contains(body, expected) {
t.Fatalf("login body missing %q", expected)
}
}
assertSecurityHeaders(t, response)
}
func TestLoginRejectsInvalidCredentialsWithoutLeakingPassword(t *testing.T) {
sessions := &fakeAdminSessions{loginErr: ErrInvalidCredentials}
router, _ := newAuthTestRouter(t, sessions)
csrf := loginCSRF(t, router)
form := url.Values{
"csrf_token": {csrf.Value},
"username": {"admin"},
"password": {"private-password-value"},
"next": {"/tasks"},
}
request := httptest.NewRequest(
http.MethodPost,
"/login",
strings.NewReader(form.Encode()),
)
request.Header.Set("Content-Type", "application/x-www-form-urlencoded")
request.AddCookie(csrf)
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != http.StatusUnauthorized ||
!strings.Contains(response.Body.String(), "账号或密码不正确") {
t.Fatalf("status/body = %d / %s", response.Code, response.Body)
}
if strings.Contains(response.Body.String(), "private-password-value") {
t.Fatal("login response contains submitted password")
}
if sessions.loginInput.Username != "admin" ||
sessions.loginInput.Password != "private-password-value" {
t.Fatalf("login input = %+v", sessions.loginInput)
}
}
func TestLoginRateLimitBlocksBeforePasswordVerification(t *testing.T) {
sessions := &fakeAdminSessions{loginErr: ErrInvalidCredentials}
limiter, err := authcommon.NewAttemptLimiter(1, time.Minute, 10)
if err != nil {
t.Fatalf("NewAttemptLimiter() error = %v", err)
}
router, _ := newAuthTestRouterWithLimiter(t, sessions, limiter)
csrf := loginCSRF(t, router)
form := url.Values{
"csrf_token": {csrf.Value},
"username": {"admin"},
"password": {"private-password-value"},
"next": {"/tasks"},
}
requestLogin := func() *httptest.ResponseRecorder {
request := httptest.NewRequest(
http.MethodPost,
"/login",
strings.NewReader(form.Encode()),
)
request.Header.Set(
"Content-Type",
"application/x-www-form-urlencoded",
)
request.AddCookie(csrf)
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
return response
}
if response := requestLogin(); response.Code != http.StatusUnauthorized {
t.Fatalf("first status = %d", response.Code)
}
response := requestLogin()
if response.Code != http.StatusTooManyRequests ||
response.Header().Get("Retry-After") == "" ||
!strings.Contains(response.Body.String(), "尝试次数过多") ||
sessions.loginCalls != 1 {
t.Fatalf(
"second status/retry/calls = %d / %q / %d",
response.Code,
response.Header().Get("Retry-After"),
sessions.loginCalls,
)
}
}
func TestLoginRotatesSessionAndCSRFThenUsesSafeRedirect(t *testing.T) {
expiresAt := time.Now().UTC().Add(8 * time.Hour)
sessions := &fakeAdminSessions{
loginResult: AdminLoginResult{
Token: mustToken(t),
ExpiresAt: expiresAt,
},
}
router, _ := newAuthTestRouter(t, sessions)
csrf := loginCSRF(t, router)
form := url.Values{
"csrf_token": {csrf.Value},
"username": {"admin"},
"password": {"valid-password"},
"next": {"https://attacker.invalid/private"},
}
request := httptest.NewRequest(
http.MethodPost,
"/login",
strings.NewReader(form.Encode()),
)
request.Header.Set("Content-Type", "application/x-www-form-urlencoded")
request.AddCookie(csrf)
request.AddCookie(&http.Cookie{
Name: AdminSessionCookieName,
Value: mustToken(t),
})
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != http.StatusSeeOther ||
response.Header().Get("Location") != "/tasks" {
t.Fatalf(
"status/location = %d / %q",
response.Code,
response.Header().Get("Location"),
)
}
var sessionCookie, rotatedCSRF *http.Cookie
for _, cookie := range response.Result().Cookies() {
switch cookie.Name {
case AdminSessionCookieName:
sessionCookie = cookie
case csrfCookieName:
rotatedCSRF = cookie
}
}
if sessionCookie == nil ||
sessionCookie.Value != sessions.loginResult.Token ||
!sessionCookie.HttpOnly ||
sessionCookie.SameSite != http.SameSiteLaxMode ||
sessionCookie.Path != "/" ||
sessionCookie.Secure {
t.Fatalf("session cookie = %+v", sessionCookie)
}
if rotatedCSRF == nil || rotatedCSRF.Value == csrf.Value {
t.Fatalf("CSRF was not rotated: %+v", rotatedCSRF)
}
if sessions.logoutToken == "" ||
sessions.logoutToken == sessions.loginResult.Token {
t.Fatalf(
"previous session was not revoked: %q",
sessions.logoutToken,
)
}
}
func TestAdminSessionCookieIsSecureForTLSRequest(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.GET("/secure-cookie", func(ctx *gin.Context) {
setAdminSessionCookie(ctx, AdminLoginResult{
Token: mustToken(t),
ExpiresAt: time.Now().UTC().Add(time.Hour),
})
ctx.Status(http.StatusNoContent)
})
request := httptest.NewRequest(http.MethodGet, "/secure-cookie", nil)
request.TLS = &tls.ConnectionState{}
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
var sessionCookie *http.Cookie
for _, cookie := range response.Result().Cookies() {
if cookie.Name == AdminSessionCookieName {
sessionCookie = cookie
break
}
}
if sessionCookie == nil || !sessionCookie.Secure {
t.Fatalf("TLS session cookie = %+v", sessionCookie)
}
}
func TestLoginFailsClosedWhenPreviousSessionCannotBeRevoked(
t *testing.T,
) {
newToken := mustToken(t)
oldToken := mustToken(t)
for oldToken == newToken {
oldToken = mustToken(t)
}
sessions := &fakeAdminSessions{
loginResult: AdminLoginResult{
Token: newToken,
ExpiresAt: time.Now().UTC().Add(8 * time.Hour),
},
logoutErr: ErrUnavailable,
}
router, _ := newAuthTestRouter(t, sessions)
csrf := loginCSRF(t, router)
form := url.Values{
"csrf_token": {csrf.Value},
"username": {"admin"},
"password": {"valid-password"},
"next": {"/tasks"},
}
request := httptest.NewRequest(
http.MethodPost,
"/login",
strings.NewReader(form.Encode()),
)
request.Header.Set("Content-Type", "application/x-www-form-urlencoded")
request.AddCookie(csrf)
request.AddCookie(&http.Cookie{
Name: AdminSessionCookieName,
Value: oldToken,
})
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != http.StatusServiceUnavailable ||
sessions.logoutToken != newToken {
t.Fatalf(
"status/last revoked token = %d / %q",
response.Code,
sessions.logoutToken,
)
}
for _, cookie := range response.Result().Cookies() {
if cookie.Name == AdminSessionCookieName &&
cookie.Value == newToken {
t.Fatal("new session cookie was returned after revoke failure")
}
}
}
func TestSafeNextRejectsExternalAndAmbiguousPaths(t *testing.T) {
tests := []string{
"https://attacker.invalid/tasks",
"//attacker.invalid/tasks",
`/tasks\redirect`,
"/tasks/../admin",
"/healthz",
"tasks",
}
for _, candidate := range tests {
if actual := safeNext(candidate); actual != "/tasks" {
t.Fatalf("safeNext(%q) = %q", candidate, actual)
}
}
if actual := safeNext("/tasks/item?id=1"); actual != "/tasks/item?id=1" {
t.Fatalf("safeNext(valid) = %q", actual)
}
}
func TestLogoutRequiresCSRFAndRevokesSession(t *testing.T) {
sessions := &fakeAdminSessions{}
router, _ := newAuthTestRouter(t, sessions)
csrf := loginCSRF(t, router)
sessionValue := mustToken(t)
form := url.Values{"csrf_token": {csrf.Value}}
request := httptest.NewRequest(
http.MethodPost,
"/logout",
strings.NewReader(form.Encode()),
)
request.Header.Set("Content-Type", "application/x-www-form-urlencoded")
request.AddCookie(csrf)
request.AddCookie(&http.Cookie{
Name: AdminSessionCookieName,
Value: sessionValue,
})
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != http.StatusSeeOther ||
response.Header().Get("Location") != "/login" ||
sessions.logoutToken != sessionValue {
t.Fatalf(
"status/location/token = %d / %q / %q",
response.Code,
response.Header().Get("Location"),
sessions.logoutToken,
)
}
foundCleared := false
for _, cookie := range response.Result().Cookies() {
if cookie.Name == AdminSessionCookieName && cookie.MaxAge < 0 {
foundCleared = true
}
}
if !foundCleared {
t.Fatal("logout did not clear the admin session cookie")
}
}
func TestLogoutRejectsWrongCSRFWithoutRevoking(t *testing.T) {
sessions := &fakeAdminSessions{}
router, _ := newAuthTestRouter(t, sessions)
request := httptest.NewRequest(
http.MethodPost,
"/logout",
strings.NewReader("csrf_token=wrong"),
)
request.Header.Set("Content-Type", "application/x-www-form-urlencoded")
request.AddCookie(&http.Cookie{
Name: csrfCookieName,
Value: mustToken(t),
})
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != http.StatusForbidden || sessions.logoutToken != "" {
t.Fatalf(
"status/logout token = %d / %q",
response.Code,
sessions.logoutToken,
)
}
}
type fakeAdminSessions struct {
loginInput AdminLoginInput
loginResult AdminLoginResult
loginErr error
loginCalls int
logoutToken string
logoutErr error
}
func (service *fakeAdminSessions) LoginAdmin(
_ context.Context,
input AdminLoginInput,
) (AdminLoginResult, error) {
service.loginCalls++
service.loginInput = input
return service.loginResult, service.loginErr
}
func (service *fakeAdminSessions) LogoutAdmin(
_ context.Context,
token string,
) error {
service.logoutToken = token
return service.logoutErr
}
func newAuthTestRouter(
t *testing.T,
sessions AdminSessionService,
) (http.Handler, *AuthHandler) {
t.Helper()
limiter, err := authcommon.NewAttemptLimiter(
100,
time.Minute,
100,
)
if err != nil {
t.Fatalf("NewAttemptLimiter() error = %v", err)
}
return newAuthTestRouterWithLimiter(t, sessions, limiter)
}
func newAuthTestRouterWithLimiter(
t *testing.T,
sessions AdminSessionService,
limiter authcommon.AttemptLimiter,
) (http.Handler, *AuthHandler) {
t.Helper()
gin.SetMode(gin.TestMode)
renderer, err := NewRenderer()
if err != nil {
t.Fatalf("NewRenderer() error = %v", err)
}
handler, err := NewAuthHandler(sessions, renderer, limiter)
if err != nil {
t.Fatalf("NewAuthHandler() error = %v", err)
}
router := gin.New()
handler.RegisterPublic(router)
handler.RegisterProtected(router)
return router, handler
}
func loginCSRF(t *testing.T, router http.Handler) *http.Cookie {
t.Helper()
response := performRequest(
t,
router,
http.MethodGet,
"/login",
nil,
"",
)
return csrfCookie(t, response)
}
var _ AdminSessionService = (*fakeAdminSessions)(nil)
+69 -28
View File
@@ -4,7 +4,6 @@ import (
"bytes"
"context"
"crypto/rand"
"crypto/subtle"
"encoding/base64"
"errors"
"io"
@@ -14,6 +13,8 @@ import (
"time"
"unicode/utf8"
"cmroubao/backend-api/internal/transport/authcommon"
"github.com/gin-gonic/gin"
)
@@ -24,8 +25,8 @@ const (
maxTitleBytes = 2048
maxSKUBytes = 512
maxDescriptionBytes = 8192
csrfCookieName = "cmroubao_admin_csrf"
csrfFormField = "csrf_token"
csrfCookieName = authcommon.CSRFCookieName
csrfFormField = authcommon.CSRFFormField
formContentType = "text/html; charset=utf-8"
cssContentType = "text/css; charset=utf-8"
javascriptContentType = "text/javascript; charset=utf-8"
@@ -50,8 +51,16 @@ func NewHandler(service Service, renderer *Renderer) (*Handler, error) {
}
func (h *Handler) Register(routes gin.IRoutes) {
h.RegisterStatic(routes)
h.RegisterProtected(routes)
}
func (h *Handler) RegisterStatic(routes gin.IRoutes) {
routes.GET("/static/admin.css", SecurityHeaders(), h.Stylesheet)
routes.GET("/static/admin.js", SecurityHeaders(), h.Script)
}
func (h *Handler) RegisterProtected(routes gin.IRoutes) {
routes.GET("/tasks", SecurityHeaders(), h.ListTasks)
routes.GET("/tasks/new", SecurityHeaders(), h.NewTask)
routes.POST("/tasks", SecurityHeaders(), h.CreateTask)
@@ -97,6 +106,11 @@ func (h *Handler) serveStatic(
}
func (h *Handler) ListTasks(ctx *gin.Context) {
token, err := csrfToken(ctx)
if err != nil {
h.renderError(ctx, http.StatusInternalServerError, "页面暂时无法打开", "请稍后重试。")
return
}
input := ListTasksInput{
Query: strings.TrimSpace(ctx.Query("q")),
Status: strings.TrimSpace(ctx.Query("status")),
@@ -127,6 +141,7 @@ func (h *Handler) ListTasks(ctx *gin.Context) {
Page: pageView{
Title: "采购任务",
TasksCurrent: true,
CSRFToken: token,
},
Query: input.Query,
Status: input.Status,
@@ -138,7 +153,7 @@ func (h *Handler) ListTasks(ctx *gin.Context) {
}
func (h *Handler) NewTask(ctx *gin.Context) {
token, err := h.csrfToken(ctx)
token, err := csrfToken(ctx)
if err != nil {
h.renderError(ctx, http.StatusInternalServerError, "页面暂时无法打开", "请稍后重试。")
return
@@ -167,7 +182,7 @@ func (h *Handler) CreateTask(ctx *gin.Context) {
if ctx.Request.MultipartForm != nil {
defer ctx.Request.MultipartForm.RemoveAll()
}
if !h.validCSRF(ctx) {
if !validCSRF(ctx) {
h.renderError(
ctx,
http.StatusForbidden,
@@ -239,7 +254,7 @@ func (h *Handler) TaskDetail(ctx *gin.Context) {
h.renderServiceError(ctx, err, "无法加载任务详情,请稍后重试。")
return
}
token, tokenErr := h.csrfToken(ctx)
token, tokenErr := csrfToken(ctx)
if tokenErr != nil {
h.renderError(ctx, http.StatusInternalServerError, "页面暂时无法打开", "请稍后重试。")
return
@@ -253,6 +268,7 @@ func (h *Handler) TaskDetail(ctx *gin.Context) {
Page: pageView{
Title: "任务详情",
TasksCurrent: true,
CSRFToken: token,
},
Task: taskDetailViewFrom(task),
CSRFToken: token,
@@ -263,7 +279,7 @@ func (h *Handler) TaskDetail(ctx *gin.Context) {
}
func (h *Handler) CancelTask(ctx *gin.Context) {
if !h.validCSRF(ctx) {
if !validCSRF(ctx) {
h.renderError(
ctx,
http.StatusForbidden,
@@ -328,7 +344,7 @@ func (h *Handler) uploadReference(
}
func (h *Handler) createPageFromRequest(ctx *gin.Context) newTaskPageView {
token, err := h.csrfToken(ctx)
token, err := csrfToken(ctx)
if err != nil {
token = ""
}
@@ -338,6 +354,7 @@ func (h *Handler) createPageFromRequest(ctx *gin.Context) newTaskPageView {
Page: pageView{
Title: "新建采购任务",
NewCurrent: true,
CSRFToken: token,
},
CSRFToken: token,
UploadKey: strings.TrimSpace(ctx.PostForm("upload_key")),
@@ -438,19 +455,35 @@ func validBudget(value string) bool {
return units > 0 || fraction > 0
}
func (h *Handler) csrfToken(ctx *gin.Context) (string, error) {
if cookie, err := ctx.Request.Cookie(csrfCookieName); err == nil &&
validToken(cookie.Value) {
return cookie.Value, nil
func csrfToken(ctx *gin.Context) (string, error) {
cookies := csrfCookies(ctx.Request)
for index := len(cookies) - 1; index >= 0; index-- {
if validToken(cookies[index].Value) {
return cookies[index].Value, nil
}
}
return rotateCSRFToken(ctx)
}
func rotateCSRFToken(ctx *gin.Context) (string, error) {
token, err := newToken()
if err != nil {
return "", err
}
http.SetCookie(ctx.Writer, &http.Cookie{
Name: csrfCookieName,
Value: token,
Name: authcommon.CSRFCookieName,
Value: "",
Path: "/tasks",
Expires: time.Unix(1, 0).UTC(),
MaxAge: -1,
HttpOnly: true,
Secure: ctx.Request.TLS != nil,
SameSite: http.SameSiteStrictMode,
})
http.SetCookie(ctx.Writer, &http.Cookie{
Name: authcommon.CSRFCookieName,
Value: token,
Path: "/",
MaxAge: 3600,
HttpOnly: true,
Secure: ctx.Request.TLS != nil,
@@ -459,24 +492,28 @@ func (h *Handler) csrfToken(ctx *gin.Context) (string, error) {
return token, nil
}
func (h *Handler) validCSRF(ctx *gin.Context) bool {
cookie, err := ctx.Request.Cookie(csrfCookieName)
if err != nil || !validToken(cookie.Value) {
return false
func validCSRF(ctx *gin.Context) bool {
presented := ctx.PostForm(authcommon.CSRFFormField)
for _, cookie := range csrfCookies(ctx.Request) {
if authcommon.ValidCSRFPair(cookie.Value, presented) {
return true
}
}
formToken := strings.TrimSpace(ctx.PostForm(csrfFormField))
if len(cookie.Value) != len(formToken) {
return false
return false
}
func csrfCookies(request *http.Request) []*http.Cookie {
result := make([]*http.Cookie, 0, 2)
for _, cookie := range request.Cookies() {
if cookie.Name == authcommon.CSRFCookieName {
result = append(result, cookie)
}
}
return subtle.ConstantTimeCompare(
[]byte(cookie.Value),
[]byte(formToken),
) == 1
return result
}
func validToken(value string) bool {
decoded, err := base64.RawURLEncoding.DecodeString(value)
return err == nil && len(decoded) == 32
return authcommon.ValidOpaqueValue(value)
}
func newToken() (string, error) {
@@ -500,6 +537,7 @@ func newTaskPage(token string) (newTaskPageView, error) {
Page: pageView{
Title: "新建采购任务",
NewCurrent: true,
CSRFToken: token,
},
CSRFToken: token,
UploadKey: uploadKey,
@@ -561,9 +599,11 @@ func (h *Handler) renderError(
title string,
message string,
) {
token, _ := csrfToken(ctx)
h.render(ctx, status, "error", errorPage{
Page: pageView{
Title: title,
Title: title,
CSRFToken: token,
},
Heading: title,
Message: message,
@@ -622,6 +662,7 @@ type pageView struct {
Title string
TasksCurrent bool
NewCurrent bool
CSRFToken string
}
type statusOption struct {
@@ -112,7 +112,7 @@ func TestNewTaskIssuesReusableStrictCSRFCookie(t *testing.T) {
cookie := csrfCookie(t, response)
if !cookie.HttpOnly ||
cookie.SameSite != http.SameSiteStrictMode ||
cookie.Path != "/tasks" {
cookie.Path != "/" {
t.Fatalf("CSRF cookie = %+v", cookie)
}
body := response.Body.String()
@@ -128,6 +128,7 @@ func TestNewTaskIssuesReusableStrictCSRFCookie(t *testing.T) {
`name="quantity"`,
`name="max_budget"`,
`name="image"`,
`action="/logout"`,
"最高总预算",
} {
if !strings.Contains(body, required) {
@@ -141,6 +142,39 @@ func TestNewTaskIssuesReusableStrictCSRFCookie(t *testing.T) {
assertSecurityHeaders(t, response)
}
func TestNewTaskPrefersRootCSRFCookieDuringLegacyPathMigration(
t *testing.T,
) {
router := newTestRouter(t, &fakeService{})
legacy := mustToken(t)
root := mustToken(t)
for root == legacy {
root = mustToken(t)
}
request := httptest.NewRequest(http.MethodGet, "/tasks/new", nil)
request.AddCookie(&http.Cookie{
Name: csrfCookieName,
Value: legacy,
Path: "/tasks",
})
request.AddCookie(&http.Cookie{
Name: csrfCookieName,
Value: root,
Path: "/",
})
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != http.StatusOK ||
!strings.Contains(
response.Body.String(),
`name="csrf_token" value="`+root+`"`,
) {
t.Fatalf("status/body = %d / %s", response.Code, response.Body)
}
}
func TestCreateTaskRejectsCSRFBeforeCallingService(t *testing.T) {
service := &fakeService{}
router := newTestRouter(t, service)
@@ -622,7 +656,9 @@ func csrfCookie(
) *http.Cookie {
t.Helper()
for _, cookie := range response.Result().Cookies() {
if cookie.Name == csrfCookieName {
if cookie.Name == csrfCookieName &&
cookie.Path == "/" &&
cookie.Value != "" {
return cookie
}
}
@@ -144,6 +144,27 @@ a {
font-weight: 700;
}
.logout-form {
margin: 0;
}
.logout-button {
min-height: 44px;
padding: 0 12px;
border: 1px solid var(--border);
border-radius: 5px;
background: var(--surface);
color: var(--text-muted);
font: inherit;
cursor: pointer;
}
.logout-button:hover,
.logout-button:focus-visible {
border-color: var(--text-muted);
color: var(--text);
}
.page {
width: min(calc(100% - 32px), 1180px);
margin: 0 auto;
@@ -705,6 +726,65 @@ tbody tr:last-child td {
cursor: wait;
}
.login-body {
min-height: 100vh;
display: grid;
background: #eef1f0;
}
.login-page {
width: min(calc(100% - 28px), 430px);
margin: auto;
padding: 32px 0;
}
.login-brand {
min-height: 44px;
display: inline-flex;
align-items: center;
gap: 10px;
margin-bottom: 18px;
color: var(--ink);
font-weight: 800;
text-decoration: none;
}
.login-panel {
padding: 28px;
border: 1px solid var(--line);
border-radius: 8px;
background: var(--surface);
}
.login-panel h1 {
margin: 0;
}
.login-panel form {
display: grid;
gap: 16px;
margin-top: 22px;
}
.password-field {
display: grid;
grid-template-columns: minmax(0, 1fr) auto;
}
.password-field input {
border-radius: 5px 0 0 5px;
}
.password-toggle {
min-width: 64px;
border-left: 0;
border-radius: 0 5px 5px 0;
}
.login-submit {
width: 100%;
}
@media (max-width: 760px) {
.page {
width: min(calc(100% - 20px), 1180px);
@@ -797,6 +877,11 @@ tbody tr:last-child td {
font-size: 13px;
}
.logout-button {
padding-inline: 8px;
font-size: 13px;
}
.title-row,
.detail-title {
align-items: stretch;
@@ -9,6 +9,19 @@
if (summary) summary.focus();
}
const passwordToggle = document.querySelector("[data-password-toggle]");
if (passwordToggle) {
const password = document.getElementById(
passwordToggle.getAttribute("aria-controls"),
);
passwordToggle.addEventListener("click", () => {
const showing = password.type === "text";
password.type = showing ? "password" : "text";
passwordToggle.textContent = showing ? "显示" : "隐藏";
passwordToggle.setAttribute("aria-pressed", String(!showing));
});
}
document.querySelectorAll("[data-loading-form]").forEach((form) => {
form.addEventListener("submit", () => {
form.setAttribute("aria-busy", "true");
@@ -0,0 +1,55 @@
{{define "login"}}
<!doctype html>
<html lang="zh-CN">
<head>
<title>{{.Page.Title}} - 采购任务管理</title>
{{template "document-head" .}}
</head>
<body class="login-body">
<main id="main-content" class="login-page">
<a class="login-brand" href="/login" aria-label="采购任务管理登录">
<span class="brand-mark" aria-hidden="true">采</span>
<span>采购任务管理</span>
</a>
<section class="login-panel" aria-labelledby="login-title">
<h1 id="login-title">管理端登录</h1>
<p class="subtitle">使用采购管理员账号继续</p>
{{if .Message}}
<div class="notice notice-error" role="alert" tabindex="-1" data-error-summary>
{{.Message}}
</div>
{{end}}
<form method="post" action="/login" novalidate data-loading-form>
<input type="hidden" name="csrf_token" value="{{.CSRFToken}}">
<input type="hidden" name="next" value="{{.Next}}">
<div class="field">
<label for="username">账号</label>
<input id="username" name="username" value="{{.Username}}"
autocomplete="username" autofocus required
aria-describedby="username-error"
{{if .UsernameError}}aria-invalid="true" data-error-field{{end}}>
<p id="username-error" class="field-error">{{.UsernameError}}</p>
</div>
<div class="field">
<label for="login-password">密码</label>
<div class="password-field">
<input id="login-password" name="password" type="password"
autocomplete="current-password" required
aria-describedby="password-error"
{{if .PasswordError}}aria-invalid="true" data-error-field{{end}}>
<button class="button password-toggle" type="button"
aria-controls="login-password" aria-pressed="false"
data-password-toggle>显示</button>
</div>
<p id="password-error" class="field-error">{{.PasswordError}}</p>
</div>
<button class="button primary login-submit" type="submit"
data-loading-label="正在登录…">登录</button>
</form>
</section>
</main>
</body>
</html>
{{end}}
@@ -18,5 +18,11 @@
<a href="/tasks" {{if .Page.TasksCurrent}}aria-current="page"{{end}}>任务列表</a>
<a href="/tasks/new" {{if .Page.NewCurrent}}aria-current="page"{{end}}>新建任务</a>
</nav>
{{if .Page.CSRFToken}}
<form class="logout-form" method="post" action="/logout">
<input type="hidden" name="csrf_token" value="{{.Page.CSRFToken}}">
<button class="logout-button" type="submit">退出</button>
</form>
{{end}}
</header>
{{end}}
@@ -5,6 +5,7 @@ import (
"errors"
"cmroubao/backend-api/internal/domain"
"cmroubao/backend-api/internal/transport/authcommon"
"cmroubao/backend-api/internal/usecase"
)
@@ -104,6 +105,7 @@ func (adapter *UsecaseAdapter) CreateTask(
}
result, err := adapter.tasks.Create(ctx, usecase.CreateTaskCommand{
CreatorSubject: localAdminSubject,
ActorUserID: actorUserID(ctx),
IdempotencyKey: input.IdempotencyKey,
Title: input.Title,
Description: input.Description,
@@ -118,12 +120,21 @@ func (adapter *UsecaseAdapter) CreateTask(
return taskFromPurchase(result.Task), nil
}
func actorUserID(ctx context.Context) string {
principal, ok := authcommon.Principal(ctx)
if !ok || principal.Role != domain.UserRoleAdmin {
return ""
}
return principal.UserID
}
func (adapter *UsecaseAdapter) CancelPending(
ctx context.Context,
input CancelPendingInput,
) (Task, error) {
task, err := adapter.tasks.Cancel(ctx, usecase.CancelTaskCommand{
CreatorSubject: localAdminSubject,
ActorUserID: actorUserID(ctx),
TaskID: input.TaskID,
Reason: "管理员取消",
})
@@ -0,0 +1,7 @@
package usecase
import "encoding/base64"
func rawURLDecode(value string) ([]byte, error) {
return base64.RawURLEncoding.DecodeString(value)
}
@@ -0,0 +1,57 @@
package usecase
import "errors"
const (
ErrorKindUnauthorized ErrorKind = "UNAUTHORIZED"
ErrorKindForbidden ErrorKind = "FORBIDDEN"
)
func wrapAuthRepositoryError(err error) error {
switch {
case errors.Is(err, ErrAuthCredentials):
return newError(
ErrorKindUnauthorized,
"AUTH_INVALID_CREDENTIALS",
"invalid credentials",
err,
)
case errors.Is(err, ErrAuthDisabled):
return newError(
ErrorKindForbidden,
"AUTH_ACCOUNT_OR_DEVICE_DISABLED",
"account or device is disabled",
err,
)
case errors.Is(err, ErrAuthForbidden):
return newError(
ErrorKindForbidden,
"AUTH_FORBIDDEN",
"identity is not allowed to perform this operation",
err,
)
case errors.Is(err, ErrAuthExpired):
return newError(
ErrorKindUnauthorized,
"AUTH_SESSION_EXPIRED",
"session expired",
err,
)
case errors.Is(err, ErrAuthRevoked):
return newError(
ErrorKindUnauthorized,
"AUTH_INVALID_TOKEN",
"invalid token",
err,
)
case errors.Is(err, ErrAuthConflict):
return newError(
ErrorKindConflict,
"AUTH_RESOURCE_CONFLICT",
"authentication resource already exists",
err,
)
default:
return wrapRepositoryError(err)
}
}
@@ -0,0 +1,62 @@
package usecase
import (
"context"
"errors"
"time"
"cmroubao/backend-api/internal/domain"
)
type PasswordManager interface {
Hash(string) (string, error)
Verify(string, string) error
VerifyDummy(string)
}
type OpaqueTokenGenerator interface {
NewToken() (string, error)
}
type AuthRepository interface {
FindUserByUsername(context.Context, string) (domain.User, error)
ProvisionUser(context.Context, domain.User) (domain.User, error)
ProvisionDevice(context.Context, domain.Device) (domain.Device, error)
SetUserActive(context.Context, string, bool, time.Time) error
SetDeviceEnabled(context.Context, string, bool, time.Time) error
CreateAdminSession(
context.Context,
domain.AdminSession,
time.Time,
) error
AuthenticateAdminSession(
context.Context,
string,
time.Time,
) (domain.AuthPrincipal, error)
RevokeAdminSession(context.Context, string, time.Time) error
CreateAccessTokenAndBindDevice(
context.Context,
string,
string,
string,
string,
string,
domain.AccessToken,
time.Time,
) (domain.Device, error)
AuthenticateAccessToken(
context.Context,
string,
time.Time,
) (domain.AuthPrincipal, error)
}
var (
ErrAuthCredentials = errors.New("authentication credentials are invalid")
ErrAuthDisabled = errors.New("authentication subject is disabled")
ErrAuthForbidden = errors.New("authentication role is forbidden")
ErrAuthExpired = errors.New("authentication session expired")
ErrAuthRevoked = errors.New("authentication session revoked")
ErrAuthConflict = errors.New("authentication resource conflict")
)
@@ -0,0 +1,20 @@
package usecase
import (
"crypto/rand"
"encoding/base64"
"errors"
"io"
)
const opaqueTokenBytes = 32
type CryptoTokenGenerator struct{}
func (CryptoTokenGenerator) NewToken() (string, error) {
value := make([]byte, opaqueTokenBytes)
if _, err := io.ReadFull(rand.Reader, value); err != nil {
return "", errors.New("generate opaque token")
}
return base64.RawURLEncoding.EncodeToString(value), nil
}
@@ -0,0 +1,543 @@
package usecase
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"strings"
"time"
"unicode/utf8"
"cmroubao/backend-api/internal/domain"
)
const (
AdminSessionLifetime = 8 * time.Hour
AccessTokenLifetime = time.Hour
)
type AuthService struct {
repository AuthRepository
passwords PasswordManager
clock Clock
ids IDGenerator
tokens OpaqueTokenGenerator
}
type LoginAdminCommand struct {
Username string
Password string
}
type AdminSessionResult struct {
Token string
ExpiresAt time.Time
User domain.User
}
type LoginBuyerDeviceCommand struct {
Username string
Password string
DeviceID string
DeviceToken string
AppVersion string
AndroidVersion string
}
type AccessTokenResult struct {
Token string
ExpiresAt time.Time
User domain.User
Device domain.Device
}
type ProvisionUserCommand struct {
Username string
Password string
Role domain.UserRole
Active bool
}
type ProvisionDeviceCommand struct {
Name string
Enabled bool
}
type ProvisionDeviceResult struct {
Device domain.Device
DeviceToken string
}
type SetUserActiveCommand struct {
Username string
Active bool
}
type SetDeviceEnabledCommand struct {
DeviceID string
Enabled bool
}
func NewAuthService(
repository AuthRepository,
passwords PasswordManager,
clock Clock,
ids IDGenerator,
tokens OpaqueTokenGenerator,
) (*AuthService, error) {
switch {
case repository == nil:
return nil, errors.New("auth repository is required")
case passwords == nil:
return nil, errors.New("password manager is required")
case clock == nil:
return nil, errors.New("clock is required")
case ids == nil:
return nil, errors.New("ID generator is required")
case tokens == nil:
return nil, errors.New("opaque token generator is required")
default:
return &AuthService{
repository: repository,
passwords: passwords,
clock: clock,
ids: ids,
tokens: tokens,
}, nil
}
}
func (s *AuthService) LoginAdmin(
ctx context.Context,
command LoginAdminCommand,
) (AdminSessionResult, error) {
username, err := validateLoginInput(command.Username, command.Password)
if err != nil {
return AdminSessionResult{}, err
}
user, err := s.verifiedUser(ctx, username, command.Password)
if err != nil {
return AdminSessionResult{}, err
}
if !user.IsActive {
return AdminSessionResult{}, wrapAuthRepositoryError(ErrAuthDisabled)
}
if user.Role != domain.UserRoleAdmin {
return AdminSessionResult{}, wrapAuthRepositoryError(ErrAuthCredentials)
}
rawToken, session, now, err := s.newAdminSession(user.ID)
if err != nil {
return AdminSessionResult{}, err
}
if err := s.repository.CreateAdminSession(ctx, session, now); err != nil {
return AdminSessionResult{}, wrapAuthRepositoryError(err)
}
user.PasswordHash = ""
return AdminSessionResult{
Token: rawToken,
ExpiresAt: session.ExpiresAt,
User: user,
}, nil
}
func (s *AuthService) AuthenticateAdmin(
ctx context.Context,
rawToken string,
) (domain.AuthPrincipal, error) {
if !validOpaqueToken(rawToken) {
return domain.AuthPrincipal{},
wrapAuthRepositoryError(ErrAuthRevoked)
}
principal, err := s.repository.AuthenticateAdminSession(
ctx,
hashSecret(rawToken),
s.clock.Now().UTC(),
)
if err != nil {
return domain.AuthPrincipal{}, wrapAuthRepositoryError(err)
}
if principal.Role != domain.UserRoleAdmin || principal.DeviceID != "" {
return domain.AuthPrincipal{},
wrapAuthRepositoryError(ErrAuthForbidden)
}
return principal, nil
}
func (s *AuthService) LogoutAdmin(
ctx context.Context,
rawToken string,
) error {
if !validOpaqueToken(rawToken) {
return nil
}
if err := s.repository.RevokeAdminSession(
ctx,
hashSecret(rawToken),
s.clock.Now().UTC(),
); err != nil {
return wrapAuthRepositoryError(err)
}
return nil
}
func (s *AuthService) LoginBuyerDevice(
ctx context.Context,
command LoginBuyerDeviceCommand,
) (AccessTokenResult, error) {
username, err := validateLoginInput(command.Username, command.Password)
if err != nil {
return AccessTokenResult{}, err
}
if strings.TrimSpace(command.DeviceID) == "" ||
!validOpaqueToken(command.DeviceToken) {
return AccessTokenResult{},
wrapAuthRepositoryError(ErrAuthCredentials)
}
appVersion := strings.TrimSpace(command.AppVersion)
androidVersion := strings.TrimSpace(command.AndroidVersion)
versionFields := make(map[string]string)
if appVersion == "" {
versionFields["app_version"] = "required"
} else if len([]byte(appVersion)) > domain.MaxVersionBytes {
versionFields["app_version"] = "must not exceed 128 UTF-8 bytes"
}
if androidVersion == "" {
versionFields["android_version"] = "required"
} else if len([]byte(androidVersion)) > domain.MaxVersionBytes {
versionFields["android_version"] =
"must not exceed 128 UTF-8 bytes"
}
if len(versionFields) > 0 {
return AccessTokenResult{}, invalidError(
"AUTH_VALIDATION_FAILED",
"authentication input is invalid",
versionFields,
)
}
user, err := s.verifiedUser(ctx, username, command.Password)
if err != nil {
return AccessTokenResult{}, err
}
if !user.IsActive {
return AccessTokenResult{}, wrapAuthRepositoryError(ErrAuthDisabled)
}
if user.Role != domain.UserRoleBuyer {
return AccessTokenResult{},
wrapAuthRepositoryError(ErrAuthCredentials)
}
rawToken, access, now, err := s.newAccessToken(
user.ID,
strings.TrimSpace(command.DeviceID),
)
if err != nil {
return AccessTokenResult{}, err
}
device, err := s.repository.CreateAccessTokenAndBindDevice(
ctx,
user.ID,
access.DeviceID,
hashSecret(command.DeviceToken),
appVersion,
androidVersion,
access,
now,
)
if err != nil {
return AccessTokenResult{}, wrapAuthRepositoryError(err)
}
user.PasswordHash = ""
device.TokenHash = ""
return AccessTokenResult{
Token: rawToken,
ExpiresAt: access.ExpiresAt,
User: user,
Device: device,
}, nil
}
func (s *AuthService) AuthenticateAccessToken(
ctx context.Context,
rawToken string,
) (domain.AuthPrincipal, error) {
if !validOpaqueToken(rawToken) {
return domain.AuthPrincipal{},
wrapAuthRepositoryError(ErrAuthRevoked)
}
principal, err := s.repository.AuthenticateAccessToken(
ctx,
hashSecret(rawToken),
s.clock.Now().UTC(),
)
if err != nil {
return domain.AuthPrincipal{}, wrapAuthRepositoryError(err)
}
if principal.Role != domain.UserRoleBuyer || principal.DeviceID == "" {
return domain.AuthPrincipal{},
wrapAuthRepositoryError(ErrAuthForbidden)
}
return principal, nil
}
func (s *AuthService) ProvisionUser(
ctx context.Context,
command ProvisionUserCommand,
) (domain.User, error) {
if err := domain.ValidateUserInput(
command.Username,
command.Password,
command.Role,
); err != nil {
var validation *domain.AuthValidationError
if errors.As(err, &validation) {
return domain.User{}, invalidError(
"AUTH_VALIDATION_FAILED",
"authentication input is invalid",
validation.Fields,
)
}
return domain.User{}, invalidError(
"AUTH_VALIDATION_FAILED",
"authentication input is invalid",
map[string]string{},
)
}
passwordHash, err := s.passwords.Hash(command.Password)
if err != nil {
return domain.User{}, internalAuthFailure(err)
}
id, err := s.ids.NewID()
if err != nil {
return domain.User{}, internalAuthFailure(err)
}
now := s.clock.Now().UTC()
user, err := s.repository.ProvisionUser(ctx, domain.User{
ID: id,
Username: domain.NormalizeUsername(command.Username),
PasswordHash: passwordHash,
Role: command.Role,
IsActive: command.Active,
CreatedAt: now,
UpdatedAt: now,
})
if err != nil {
return domain.User{}, wrapAuthRepositoryError(err)
}
user.PasswordHash = ""
return user, nil
}
func (s *AuthService) ProvisionDevice(
ctx context.Context,
command ProvisionDeviceCommand,
) (ProvisionDeviceResult, error) {
if err := domain.ValidateDeviceName(command.Name); err != nil {
return ProvisionDeviceResult{}, invalidError(
"AUTH_VALIDATION_FAILED",
"authentication input is invalid",
map[string]string{"name": err.Error()},
)
}
id, err := s.ids.NewID()
if err != nil {
return ProvisionDeviceResult{}, internalAuthFailure(err)
}
rawToken, err := s.tokens.NewToken()
if err != nil {
return ProvisionDeviceResult{}, internalAuthFailure(err)
}
if !validOpaqueToken(rawToken) {
return ProvisionDeviceResult{},
internalAuthFailure(errors.New("token generator returned invalid token"))
}
now := s.clock.Now().UTC()
device, err := s.repository.ProvisionDevice(ctx, domain.Device{
ID: id,
Name: strings.TrimSpace(command.Name),
TokenHash: hashSecret(rawToken),
IsEnabled: command.Enabled,
CreatedAt: now,
UpdatedAt: now,
})
if err != nil {
return ProvisionDeviceResult{}, wrapAuthRepositoryError(err)
}
device.TokenHash = ""
return ProvisionDeviceResult{
Device: device,
DeviceToken: rawToken,
}, nil
}
func (s *AuthService) SetUserActive(
ctx context.Context,
command SetUserActiveCommand,
) error {
username := domain.NormalizeUsername(command.Username)
if username == "" ||
!utf8.ValidString(username) ||
len([]byte(username)) > domain.MaxUsernameBytes {
return invalidError(
"AUTH_VALIDATION_FAILED",
"authentication input is invalid",
map[string]string{"username": "invalid"},
)
}
if err := s.repository.SetUserActive(
ctx,
username,
command.Active,
s.clock.Now().UTC(),
); err != nil {
return wrapAuthRepositoryError(err)
}
return nil
}
func (s *AuthService) SetDeviceEnabled(
ctx context.Context,
command SetDeviceEnabledCommand,
) error {
deviceID := strings.TrimSpace(command.DeviceID)
if !isUUID(deviceID) {
return invalidError(
"AUTH_VALIDATION_FAILED",
"authentication input is invalid",
map[string]string{"device_id": "must be a UUID"},
)
}
if err := s.repository.SetDeviceEnabled(
ctx,
deviceID,
command.Enabled,
s.clock.Now().UTC(),
); err != nil {
return wrapAuthRepositoryError(err)
}
return nil
}
func (s *AuthService) verifiedUser(
ctx context.Context,
username string,
password string,
) (domain.User, error) {
user, err := s.repository.FindUserByUsername(ctx, username)
if err != nil {
if errors.Is(err, ErrRepositoryNotFound) {
s.passwords.VerifyDummy(password)
return domain.User{},
wrapAuthRepositoryError(ErrAuthCredentials)
}
return domain.User{}, wrapAuthRepositoryError(err)
}
if err := s.passwords.Verify(user.PasswordHash, password); err != nil {
return domain.User{}, wrapAuthRepositoryError(ErrAuthCredentials)
}
return user, nil
}
func (s *AuthService) newAdminSession(
userID string,
) (string, domain.AdminSession, time.Time, error) {
rawToken, err := s.tokens.NewToken()
if err != nil {
return "", domain.AdminSession{}, time.Time{},
internalAuthFailure(err)
}
if !validOpaqueToken(rawToken) {
return "", domain.AdminSession{}, time.Time{},
internalAuthFailure(errors.New("token generator returned invalid token"))
}
id, err := s.ids.NewID()
if err != nil {
return "", domain.AdminSession{}, time.Time{},
internalAuthFailure(err)
}
now := s.clock.Now().UTC()
return rawToken, domain.AdminSession{
ID: id,
TokenHash: hashSecret(rawToken),
UserID: userID,
ExpiresAt: now.Add(AdminSessionLifetime),
CreatedAt: now,
}, now, nil
}
func (s *AuthService) newAccessToken(
userID string,
deviceID string,
) (string, domain.AccessToken, time.Time, error) {
rawToken, err := s.tokens.NewToken()
if err != nil {
return "", domain.AccessToken{}, time.Time{},
internalAuthFailure(err)
}
if !validOpaqueToken(rawToken) {
return "", domain.AccessToken{}, time.Time{},
internalAuthFailure(errors.New("token generator returned invalid token"))
}
id, err := s.ids.NewID()
if err != nil {
return "", domain.AccessToken{}, time.Time{},
internalAuthFailure(err)
}
now := s.clock.Now().UTC()
return rawToken, domain.AccessToken{
ID: id,
TokenHash: hashSecret(rawToken),
UserID: userID,
DeviceID: deviceID,
ExpiresAt: now.Add(AccessTokenLifetime),
CreatedAt: now,
}, now, nil
}
func validateLoginInput(username, password string) (string, error) {
normalized := domain.NormalizeUsername(username)
fields := make(map[string]string)
if normalized == "" ||
len([]byte(normalized)) > domain.MaxUsernameBytes {
fields["username"] = "invalid"
}
if password == "" ||
!utf8.ValidString(password) ||
len([]byte(password)) > domain.MaxPasswordBytes {
fields["password"] = "invalid"
}
if len(fields) > 0 {
return "", invalidError(
"AUTH_VALIDATION_FAILED",
"authentication input is invalid",
fields,
)
}
return normalized, nil
}
func hashSecret(value string) string {
sum := sha256.Sum256([]byte(value))
return hex.EncodeToString(sum[:])
}
func validOpaqueToken(value string) bool {
decoded, err := decodeOpaqueToken(value)
return err == nil && len(decoded) == opaqueTokenBytes
}
func decodeOpaqueToken(value string) ([]byte, error) {
// Tokens are generated with RawURLEncoding; accepting padding would create
// multiple textual representations of the same credential.
return rawURLDecode(value)
}
func internalAuthFailure(err error) error {
return newError(
ErrorKindInternal,
"INTERNAL_ERROR",
"internal server error",
err,
)
}
@@ -0,0 +1,481 @@
package usecase
import (
"context"
"encoding/base64"
"errors"
"fmt"
"testing"
"time"
"cmroubao/backend-api/internal/domain"
)
func TestAuthServiceAdminSessionLifecycle(t *testing.T) {
fixture := newAuthServiceFixture(t)
fixture.repository.users["admin"] = domain.User{
ID: "user-admin",
Username: "admin",
PasswordHash: "hashed:password-123",
Role: domain.UserRoleAdmin,
IsActive: true,
}
result, err := fixture.service.LoginAdmin(
context.Background(),
LoginAdminCommand{Username: " ADMIN ", Password: "password-123"},
)
if err != nil {
t.Fatalf("LoginAdmin() error = %v", err)
}
if result.ExpiresAt.Sub(fixture.clock.now) != AdminSessionLifetime {
t.Fatalf("admin lifetime = %s", result.ExpiresAt.Sub(fixture.clock.now))
}
if result.User.PasswordHash != "" {
t.Fatal("LoginAdmin() exposed password hash")
}
if fixture.repository.adminSession.TokenHash == result.Token ||
len(fixture.repository.adminSession.TokenHash) != 64 {
t.Fatal("admin session was not stored as a SHA-256 hash")
}
fixture.repository.adminPrincipal = domain.AuthPrincipal{
UserID: result.User.ID,
Username: result.User.Username,
Role: domain.UserRoleAdmin,
SessionID: fixture.repository.adminSession.ID,
ExpiresAt: result.ExpiresAt,
}
principal, err := fixture.service.AuthenticateAdmin(
context.Background(),
result.Token,
)
if err != nil {
t.Fatalf("AuthenticateAdmin() error = %v", err)
}
if principal.Role != domain.UserRoleAdmin || principal.DeviceID != "" {
t.Fatalf("AuthenticateAdmin() principal = %+v", principal)
}
if err := fixture.service.LogoutAdmin(
context.Background(),
result.Token,
); err != nil {
t.Fatalf("LogoutAdmin() error = %v", err)
}
if fixture.repository.revokedHash !=
fixture.repository.adminSession.TokenHash {
t.Fatal("LogoutAdmin() did not revoke the hashed token")
}
}
func TestAuthServiceBuyerDeviceTokenLifecycle(t *testing.T) {
fixture := newAuthServiceFixture(t)
fixture.repository.users["buyer"] = domain.User{
ID: "user-buyer",
Username: "buyer",
PasswordHash: "hashed:password-123",
Role: domain.UserRoleBuyer,
IsActive: true,
}
deviceToken := validTestToken(99)
fixture.repository.device = domain.Device{
ID: "device-1",
Name: "Device",
TokenHash: hashSecret(deviceToken),
IsEnabled: true,
}
result, err := fixture.service.LoginBuyerDevice(
context.Background(),
LoginBuyerDeviceCommand{
Username: "buyer",
Password: "password-123",
DeviceID: "device-1",
DeviceToken: deviceToken,
AppVersion: "0.1.0",
AndroidVersion: "16",
},
)
if err != nil {
t.Fatalf("LoginBuyerDevice() error = %v", err)
}
if result.ExpiresAt.Sub(fixture.clock.now) != AccessTokenLifetime {
t.Fatalf("access lifetime = %s", result.ExpiresAt.Sub(fixture.clock.now))
}
if result.User.PasswordHash != "" || result.Device.TokenHash != "" {
t.Fatal("LoginBuyerDevice() exposed a credential hash")
}
if fixture.repository.access.TokenHash == result.Token ||
len(fixture.repository.access.TokenHash) != 64 {
t.Fatal("access token was not stored as a SHA-256 hash")
}
if fixture.repository.deviceTokenHash != hashSecret(deviceToken) {
t.Fatal("device credential was not passed as a hash")
}
fixture.repository.accessPrincipal = domain.AuthPrincipal{
UserID: result.User.ID,
Username: result.User.Username,
Role: domain.UserRoleBuyer,
SessionID: fixture.repository.access.ID,
DeviceID: result.Device.ID,
ExpiresAt: result.ExpiresAt,
}
principal, err := fixture.service.AuthenticateAccessToken(
context.Background(),
result.Token,
)
if err != nil {
t.Fatalf("AuthenticateAccessToken() error = %v", err)
}
if principal.DeviceID != "device-1" {
t.Fatalf("AuthenticateAccessToken() principal = %+v", principal)
}
}
func TestAuthServiceRejectsWrongRoleDisabledAndMalformedTokens(t *testing.T) {
fixture := newAuthServiceFixture(t)
fixture.repository.users["buyer"] = domain.User{
ID: "buyer",
Username: "buyer",
PasswordHash: "hashed:password-123",
Role: domain.UserRoleBuyer,
IsActive: true,
}
_, err := fixture.service.LoginAdmin(
context.Background(),
LoginAdminCommand{Username: "buyer", Password: "password-123"},
)
assertAuthCode(t, err, "AUTH_INVALID_CREDENTIALS")
fixture.repository.users["buyer"] = domain.User{
ID: "buyer",
Username: "buyer",
PasswordHash: "hashed:password-123",
Role: domain.UserRoleBuyer,
IsActive: false,
}
_, err = fixture.service.LoginBuyerDevice(
context.Background(),
LoginBuyerDeviceCommand{
Username: "buyer",
Password: "password-123",
DeviceID: "device",
DeviceToken: validTestToken(90),
AppVersion: "0.1.0",
AndroidVersion: "16",
},
)
assertAuthCode(t, err, "AUTH_ACCOUNT_OR_DEVICE_DISABLED")
_, err = fixture.service.AuthenticateAdmin(
context.Background(),
"not-a-token",
)
assertAuthCode(t, err, "AUTH_INVALID_TOKEN")
}
func TestAuthServiceProvisionsHashedCredentials(t *testing.T) {
fixture := newAuthServiceFixture(t)
user, err := fixture.service.ProvisionUser(
context.Background(),
ProvisionUserCommand{
Username: " Admin ",
Password: "password-123",
Role: domain.UserRoleAdmin,
Active: true,
},
)
if err != nil {
t.Fatalf("ProvisionUser() error = %v", err)
}
if user.Username != "admin" || user.PasswordHash != "" {
t.Fatalf("ProvisionUser() = %+v", user)
}
stored := fixture.repository.users["admin"]
if stored.PasswordHash != "hashed:password-123" {
t.Fatalf("stored password hash = %q", stored.PasswordHash)
}
device, err := fixture.service.ProvisionDevice(
context.Background(),
ProvisionDeviceCommand{Name: " Test Device ", Enabled: true},
)
if err != nil {
t.Fatalf("ProvisionDevice() error = %v", err)
}
if device.DeviceToken == "" || device.Device.TokenHash != "" {
t.Fatalf("ProvisionDevice() = %+v", device)
}
if fixture.repository.device.TokenHash == device.DeviceToken ||
fixture.repository.device.TokenHash != hashSecret(device.DeviceToken) {
t.Fatal("stored device token is not a SHA-256 hash")
}
}
func TestAuthServiceChangesUserAndDeviceStatus(t *testing.T) {
fixture := newAuthServiceFixture(t)
if err := fixture.service.SetUserActive(
context.Background(),
SetUserActiveCommand{Username: " Buyer ", Active: false},
); err != nil {
t.Fatalf("SetUserActive() error = %v", err)
}
deviceID := "00000000-0000-4000-8000-000000000099"
if err := fixture.service.SetDeviceEnabled(
context.Background(),
SetDeviceEnabledCommand{DeviceID: deviceID, Enabled: false},
); err != nil {
t.Fatalf("SetDeviceEnabled() error = %v", err)
}
if fixture.repository.userStatusName != "buyer" ||
fixture.repository.userActive ||
fixture.repository.deviceStatusID != deviceID ||
fixture.repository.deviceEnabled {
t.Fatalf(
"user/device state = %q/%v %q/%v",
fixture.repository.userStatusName,
fixture.repository.userActive,
fixture.repository.deviceStatusID,
fixture.repository.deviceEnabled,
)
}
err := fixture.service.SetDeviceEnabled(
context.Background(),
SetDeviceEnabledCommand{DeviceID: "not-a-uuid", Enabled: true},
)
assertAuthCode(t, err, "AUTH_VALIDATION_FAILED")
}
func TestAuthServiceUnknownUserUsesPrecomputedDummyVerification(t *testing.T) {
fixture := newAuthServiceFixture(t)
_, err := fixture.service.LoginAdmin(
context.Background(),
LoginAdminCommand{
Username: "missing-user",
Password: "password-123",
},
)
assertAuthCode(t, err, "AUTH_INVALID_CREDENTIALS")
if fixture.passwords.dummyVerifications != 1 {
t.Fatalf(
"dummy verifications = %d, want 1",
fixture.passwords.dummyVerifications,
)
}
if fixture.passwords.hashCalls != 0 {
t.Fatalf("Hash() calls for missing user = %d", fixture.passwords.hashCalls)
}
}
func assertAuthCode(t *testing.T, err error, code string) {
t.Helper()
var authError *Error
if !errors.As(err, &authError) || authError.Code != code {
t.Fatalf("error = %v, want code %s", err, code)
}
}
type authServiceFixture struct {
service *AuthService
repository *fakeAuthRepository
clock fixedAuthClock
passwords *fakePasswordManager
}
func newAuthServiceFixture(t *testing.T) authServiceFixture {
t.Helper()
repository := &fakeAuthRepository{
users: make(map[string]domain.User),
}
clock := fixedAuthClock{
now: time.Date(2026, 7, 26, 10, 0, 0, 0, time.UTC),
}
passwords := &fakePasswordManager{}
service, err := NewAuthService(
repository,
passwords,
clock,
&sequenceIDGenerator{},
&sequenceTokenGenerator{},
)
if err != nil {
t.Fatalf("NewAuthService() error = %v", err)
}
return authServiceFixture{
service: service,
repository: repository,
clock: clock,
passwords: passwords,
}
}
type fixedAuthClock struct {
now time.Time
}
func (clock fixedAuthClock) Now() time.Time {
return clock.now
}
type fakePasswordManager struct {
hashCalls int
dummyVerifications int
}
func (manager *fakePasswordManager) Hash(value string) (string, error) {
manager.hashCalls++
return "hashed:" + value, nil
}
func (*fakePasswordManager) Verify(encoded, plain string) error {
if encoded != "hashed:"+plain {
return errors.New("mismatch")
}
return nil
}
func (manager *fakePasswordManager) VerifyDummy(string) {
manager.dummyVerifications++
}
type sequenceIDGenerator struct {
next int
}
func (generator *sequenceIDGenerator) NewID() (string, error) {
generator.next++
return fmt.Sprintf("id-%d", generator.next), nil
}
type sequenceTokenGenerator struct {
next byte
}
func (generator *sequenceTokenGenerator) NewToken() (string, error) {
generator.next++
return validTestToken(generator.next), nil
}
func validTestToken(seed byte) string {
value := make([]byte, opaqueTokenBytes)
for index := range value {
value[index] = seed + byte(index)
}
return base64.RawURLEncoding.EncodeToString(value)
}
type fakeAuthRepository struct {
users map[string]domain.User
device domain.Device
adminSession domain.AdminSession
adminPrincipal domain.AuthPrincipal
access domain.AccessToken
accessPrincipal domain.AuthPrincipal
deviceTokenHash string
revokedHash string
authenticationErr error
userActive bool
userStatusName string
deviceEnabled bool
deviceStatusID string
}
func (repository *fakeAuthRepository) FindUserByUsername(
_ context.Context,
username string,
) (domain.User, error) {
user, found := repository.users[username]
if !found {
return domain.User{}, ErrRepositoryNotFound
}
return user, nil
}
func (repository *fakeAuthRepository) ProvisionUser(
_ context.Context,
user domain.User,
) (domain.User, error) {
repository.users[user.Username] = user
return user, nil
}
func (repository *fakeAuthRepository) ProvisionDevice(
_ context.Context,
device domain.Device,
) (domain.Device, error) {
repository.device = device
return device, nil
}
func (repository *fakeAuthRepository) SetUserActive(
_ context.Context,
username string,
active bool,
_ time.Time,
) error {
repository.userStatusName = username
repository.userActive = active
return nil
}
func (repository *fakeAuthRepository) SetDeviceEnabled(
_ context.Context,
deviceID string,
enabled bool,
_ time.Time,
) error {
repository.deviceStatusID = deviceID
repository.deviceEnabled = enabled
return nil
}
func (repository *fakeAuthRepository) CreateAdminSession(
_ context.Context,
session domain.AdminSession,
_ time.Time,
) error {
repository.adminSession = session
return nil
}
func (repository *fakeAuthRepository) AuthenticateAdminSession(
_ context.Context,
_ string,
_ time.Time,
) (domain.AuthPrincipal, error) {
return repository.adminPrincipal, repository.authenticationErr
}
func (repository *fakeAuthRepository) RevokeAdminSession(
_ context.Context,
tokenHash string,
_ time.Time,
) error {
repository.revokedHash = tokenHash
return nil
}
func (repository *fakeAuthRepository) CreateAccessTokenAndBindDevice(
_ context.Context,
_ string,
_ string,
deviceTokenHash string,
_ string,
_ string,
access domain.AccessToken,
_ time.Time,
) (domain.Device, error) {
repository.deviceTokenHash = deviceTokenHash
repository.access = access
return repository.device, nil
}
func (repository *fakeAuthRepository) AuthenticateAccessToken(
_ context.Context,
_ string,
_ time.Time,
) (domain.AuthPrincipal, error) {
return repository.accessPrincipal, repository.authenticationErr
}
+45 -24
View File
@@ -29,6 +29,7 @@ type TaskService struct {
type CreateTaskCommand struct {
CreatorSubject string
ActorUserID string
IdempotencyKey string
SourceRef *string
Title string
@@ -61,6 +62,7 @@ type TaskPage struct {
type CancelTaskCommand struct {
CreatorSubject string
ActorUserID string
TaskID string
Reason string
}
@@ -91,6 +93,7 @@ func (s *TaskService) Create(
return CreateTaskResult{}, err
}
command.CreatorSubject = strings.TrimSpace(command.CreatorSubject)
command.ActorUserID = strings.TrimSpace(command.ActorUserID)
command.Title = strings.TrimSpace(command.Title)
command.SKU = strings.TrimSpace(command.SKU)
command.ImageAssetID = strings.TrimSpace(command.ImageAssetID)
@@ -105,6 +108,13 @@ func (s *TaskService) Create(
map[string]string{"image_asset_id": "must be a UUID"},
)
}
if !isUUID(command.ActorUserID) {
return CreateTaskResult{}, invalidError(
"TASK_VALIDATION_FAILED",
"task validation failed",
map[string]string{"actor_user_id": "must be a UUID"},
)
}
if err := domain.ValidateTaskInput(
command.CreatorSubject,
command.SourceRef,
@@ -156,28 +166,31 @@ func (s *TaskService) Create(
)
}
now := s.clock.Now().UTC()
actorUserID := command.ActorUserID
task := domain.PurchaseTask{
ID: taskID,
CreatorSubject: command.CreatorSubject,
SourceRef: command.SourceRef,
Title: command.Title,
Description: command.Description,
SKU: command.SKU,
ImageAssetID: command.ImageAssetID,
Quantity: command.Quantity,
MaxBudgetCents: budget,
Currency: domain.CurrencyCNY,
Status: domain.TaskStatusPending,
Version: 1,
CreatedAt: now,
UpdatedAt: now,
ID: taskID,
CreatorSubject: command.CreatorSubject,
CreatedByUserID: &actorUserID,
SourceRef: command.SourceRef,
Title: command.Title,
Description: command.Description,
SKU: command.SKU,
ImageAssetID: command.ImageAssetID,
Quantity: command.Quantity,
MaxBudgetCents: budget,
Currency: domain.CurrencyCNY,
Status: domain.TaskStatusPending,
Version: 1,
CreatedAt: now,
UpdatedAt: now,
}
event := domain.TaskEvent{
ID: eventID,
TaskID: taskID,
Type: "TASK_CREATED",
Message: "task created",
OccurredAt: now,
ID: eventID,
TaskID: taskID,
ActorUserID: &actorUserID,
Type: "TASK_CREATED",
Message: "task created",
OccurredAt: now,
}
requestHash, err := hashCreateTaskCommand(command, budget)
if err != nil {
@@ -323,6 +336,7 @@ func (s *TaskService) Cancel(
command CancelTaskCommand,
) (domain.PurchaseTask, error) {
command.CreatorSubject = strings.TrimSpace(command.CreatorSubject)
command.ActorUserID = strings.TrimSpace(command.ActorUserID)
command.TaskID = strings.TrimSpace(command.TaskID)
command.Reason = strings.TrimSpace(command.Reason)
fields := make(map[string]string)
@@ -332,6 +346,9 @@ func (s *TaskService) Cancel(
if !isUUID(command.TaskID) {
fields["task_id"] = "must be a UUID"
}
if !isUUID(command.ActorUserID) {
fields["actor_user_id"] = "must be a UUID"
}
if len([]byte(command.Reason)) > domain.MaxCancelReasonBytes {
fields["reason"] = "too long"
}
@@ -352,12 +369,14 @@ func (s *TaskService) Cancel(
)
}
now := s.clock.Now().UTC()
actorUserID := command.ActorUserID
event := domain.TaskEvent{
ID: eventID,
TaskID: command.TaskID,
Type: "TASK_CANCELED",
Message: "task canceled",
OccurredAt: now,
ID: eventID,
TaskID: command.TaskID,
ActorUserID: &actorUserID,
Type: "TASK_CANCELED",
Message: "task canceled",
OccurredAt: now,
}
task, err := s.repository.CancelPendingTask(
ctx,
@@ -378,6 +397,7 @@ func hashCreateTaskCommand(
budget *int64,
) (string, error) {
payload := struct {
ActorUserID string `json:"actor_user_id"`
SourceRef *string `json:"source_ref"`
Title string `json:"title"`
Description string `json:"description"`
@@ -386,6 +406,7 @@ func hashCreateTaskCommand(
Quantity int `json:"quantity"`
MaxBudgetCents *int64 `json:"max_budget_cents"`
}{
ActorUserID: command.ActorUserID,
SourceRef: command.SourceRef,
Title: command.Title,
Description: command.Description,
@@ -17,6 +17,7 @@ func TestTaskServiceCreateNormalizesAndHashesForIdempotency(t *testing.T) {
result, err := service.Create(context.Background(), CreateTaskCommand{
CreatorSubject: " local-admin ",
ActorUserID: "00000000-0000-4000-8000-000000000099",
IdempotencyKey: " create-1 ",
SourceRef: &sourceRef,
Title: " Demo title ",
@@ -34,6 +35,9 @@ func TestTaskServiceCreateNormalizesAndHashesForIdempotency(t *testing.T) {
result.Task.SKU != "SKU-1" ||
result.Task.SourceRef == nil ||
*result.Task.SourceRef != "source-1" ||
result.Task.CreatedByUserID == nil ||
*result.Task.CreatedByUserID !=
"00000000-0000-4000-8000-000000000099" ||
result.Task.MaxBudgetCents == nil ||
*result.Task.MaxBudgetCents != 2000 {
t.Fatalf("created task = %+v", result.Task)
@@ -46,7 +50,10 @@ func TestTaskServiceCreateNormalizesAndHashesForIdempotency(t *testing.T) {
)
}
if repository.event.Type != "TASK_CREATED" ||
repository.event.TaskID != result.Task.ID {
repository.event.TaskID != result.Task.ID ||
repository.event.ActorUserID == nil ||
*repository.event.ActorUserID !=
"00000000-0000-4000-8000-000000000099" {
t.Fatalf("event = %+v", repository.event)
}
}
@@ -56,6 +63,7 @@ func TestTaskServiceCreateMapsValidationAndRepositoryErrors(t *testing.T) {
service := mustTaskService(t, repository)
_, err := service.Create(context.Background(), CreateTaskCommand{
CreatorSubject: "local-admin",
ActorUserID: "00000000-0000-4000-8000-000000000099",
IdempotencyKey: "create-1",
Title: "title",
SKU: "sku",
@@ -66,6 +74,7 @@ func TestTaskServiceCreateMapsValidationAndRepositoryErrors(t *testing.T) {
_, err = service.Create(context.Background(), CreateTaskCommand{
CreatorSubject: "local-admin",
ActorUserID: "00000000-0000-4000-8000-000000000099",
IdempotencyKey: "create-2",
Title: "",
SKU: "",
@@ -112,10 +121,16 @@ func TestTaskServiceCancelMapsStateConflict(t *testing.T) {
service := mustTaskService(t, repository)
_, err := service.Cancel(context.Background(), CancelTaskCommand{
CreatorSubject: "local-admin",
ActorUserID: "00000000-0000-4000-8000-000000000099",
TaskID: "00000000-0000-4000-8000-000000000001",
Reason: "no longer needed",
})
assertUsecaseError(t, err, ErrorKindConflict, "TASK_STATE_CONFLICT")
if repository.event.ActorUserID == nil ||
*repository.event.ActorUserID !=
"00000000-0000-4000-8000-000000000099" {
t.Fatalf("cancel event = %+v", repository.event)
}
}
type fakeClock struct{}
@@ -194,8 +209,9 @@ func (repository *fakeTaskRepository) CancelPendingTask(
_ string,
_ string,
_ time.Time,
_ domain.TaskEvent,
event domain.TaskEvent,
) (domain.PurchaseTask, error) {
repository.event = event
return domain.PurchaseTask{}, repository.cancelErr
}