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