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

570 lines
13 KiB
Go
Raw Normal View History

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)
}
2026-07-29 08:43:35 +08:00
func (s *Store) ResetUserPassword(
ctx context.Context,
username string,
passwordHash string,
updatedAt time.Time,
) error {
result, err := s.db.ExecContext(
ctx,
`UPDATE users
SET password_hash = ?, updated_at = ?
WHERE username = ?`,
passwordHash,
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 readinessAt 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, readiness_reported_at,
accessibility_enabled, pdd_installed, is_enabled,
created_at, updated_at
FROM devices
WHERE id = ?`,
deviceID,
).Scan(
&device.ID,
&device.Name,
&device.TokenHash,
&boundUserID,
&appVersion,
&androidVersion,
&pddVersion,
&lastSeenAt,
&readinessAt,
&device.AccessibilityEnabled,
&device.PDDInstalled,
&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
}
if readinessAt.Valid {
parsed, err := parseTimestamp(readinessAt.String)
if err != nil {
return domain.Device{}, err
}
device.ReadinessAt = &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)