221 lines
8.3 KiB
Go
221 lines
8.3 KiB
Go
package repository
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"time"
|
|
|
|
"github.com/google/uuid"
|
|
"gorm.io/gorm"
|
|
)
|
|
|
|
// ErrPinCustomerNotFound means the customer does not exist.
|
|
var ErrPinCustomerNotFound = errors.New("pin: customer not found")
|
|
|
|
// CustomerPinState is a customer's PIN and what guards it. It lives in the customers
|
|
// table but is read and written only here, never through the Customer entity, so the
|
|
// hash cannot end up in a customer response.
|
|
type CustomerPinState struct {
|
|
CustomerID uuid.UUID
|
|
OrganizationID uuid.UUID
|
|
PhoneNumber *string
|
|
BirthDate *time.Time
|
|
PinHash *string
|
|
PinSetAt *time.Time
|
|
FailedAttempts int
|
|
LockedUntil *time.Time
|
|
TransferBlockedUntil *time.Time
|
|
}
|
|
|
|
// CustomerSecurityEvent is one row of the PIN security log.
|
|
type CustomerSecurityEvent struct {
|
|
ID uuid.UUID
|
|
CustomerID uuid.UUID
|
|
Event string
|
|
ActorUser *uuid.UUID
|
|
Reason *string
|
|
IPAddress *string
|
|
UserAgent *string
|
|
CreatedAt time.Time
|
|
}
|
|
|
|
// CustomerPinRepository stores customer PINs and their security log
|
|
// (docs/prd-point-coin.md F11).
|
|
type CustomerPinRepository interface {
|
|
GetState(ctx context.Context, customerID uuid.UUID) (*CustomerPinState, error)
|
|
// SetPin stores a new PIN hash, clears the failure counter and any lock, and sets
|
|
// or clears the transfer hold.
|
|
SetPin(ctx context.Context, customerID uuid.UUID, hash string, transferBlockedUntil *time.Time) error
|
|
// RemovePin deletes the PIN, so the customer has to create a new one through OTP.
|
|
RemovePin(ctx context.Context, customerID uuid.UUID) error
|
|
// RecordFailure adds one wrong attempt in a single statement, so wrong attempts
|
|
// made at the same time all count. A lock that has already run out starts the
|
|
// count again. When the count reaches maxAttempts the PIN is locked until
|
|
// lockUntil. It returns the count and lock after the update.
|
|
RecordFailure(ctx context.Context, customerID uuid.UUID, maxAttempts int, now, lockUntil time.Time) (int, *time.Time, error)
|
|
ClearFailures(ctx context.Context, customerID uuid.UUID) error
|
|
|
|
InsertEvent(ctx context.Context, event CustomerSecurityEvent) error
|
|
// ListEvents returns a page of the customer's log, newest first, and the total.
|
|
ListEvents(ctx context.Context, customerID uuid.UUID, offset, limit int) ([]CustomerSecurityEvent, int64, error)
|
|
}
|
|
|
|
type customerPinRepository struct {
|
|
db *gorm.DB
|
|
}
|
|
|
|
func NewCustomerPinRepository(db *gorm.DB) CustomerPinRepository {
|
|
return &customerPinRepository{db: db}
|
|
}
|
|
|
|
func (r *customerPinRepository) GetState(ctx context.Context, customerID uuid.UUID) (*CustomerPinState, error) {
|
|
var rows []struct {
|
|
CustomerID string
|
|
OrganizationID string
|
|
PhoneNumber *string
|
|
BirthDate *time.Time
|
|
PinHash *string
|
|
PinSetAt *time.Time
|
|
PinFailedAttempts int
|
|
PinLockedUntil *time.Time
|
|
TransferBlockedUntil *time.Time
|
|
}
|
|
err := DBFromContext(ctx, r.db).WithContext(ctx).Raw(`
|
|
SELECT id::text AS customer_id, organization_id::text AS organization_id,
|
|
COALESCE(phone_number, phone) AS phone_number, birth_date,
|
|
pin_hash, pin_set_at, pin_failed_attempts, pin_locked_until, transfer_blocked_until
|
|
FROM customers WHERE id = ? LIMIT 1`, customerID).Scan(&rows).Error
|
|
if err != nil {
|
|
return nil, fmt.Errorf("failed to read customer PIN: %w", err)
|
|
}
|
|
if len(rows) == 0 {
|
|
return nil, ErrPinCustomerNotFound
|
|
}
|
|
row := rows[0]
|
|
state := &CustomerPinState{
|
|
PhoneNumber: row.PhoneNumber,
|
|
BirthDate: row.BirthDate,
|
|
PinHash: row.PinHash,
|
|
PinSetAt: row.PinSetAt,
|
|
FailedAttempts: row.PinFailedAttempts,
|
|
LockedUntil: row.PinLockedUntil,
|
|
TransferBlockedUntil: row.TransferBlockedUntil,
|
|
}
|
|
state.CustomerID, _ = uuid.Parse(row.CustomerID)
|
|
state.OrganizationID, _ = uuid.Parse(row.OrganizationID)
|
|
return state, nil
|
|
}
|
|
|
|
func (r *customerPinRepository) SetPin(ctx context.Context, customerID uuid.UUID, hash string, transferBlockedUntil *time.Time) error {
|
|
result := DBFromContext(ctx, r.db).WithContext(ctx).Exec(`
|
|
UPDATE customers SET pin_hash = ?, pin_set_at = NOW(), pin_failed_attempts = 0,
|
|
pin_locked_until = NULL, transfer_blocked_until = ?, updated_at = NOW()
|
|
WHERE id = ?`, hash, transferBlockedUntil, customerID)
|
|
if result.Error != nil {
|
|
return fmt.Errorf("failed to store customer PIN: %w", result.Error)
|
|
}
|
|
if result.RowsAffected == 0 {
|
|
return ErrPinCustomerNotFound
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (r *customerPinRepository) RemovePin(ctx context.Context, customerID uuid.UUID) error {
|
|
result := DBFromContext(ctx, r.db).WithContext(ctx).Exec(`
|
|
UPDATE customers SET pin_hash = NULL, pin_set_at = NULL, pin_failed_attempts = 0,
|
|
pin_locked_until = NULL, updated_at = NOW()
|
|
WHERE id = ?`, customerID)
|
|
if result.Error != nil {
|
|
return fmt.Errorf("failed to remove customer PIN: %w", result.Error)
|
|
}
|
|
if result.RowsAffected == 0 {
|
|
return ErrPinCustomerNotFound
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (r *customerPinRepository) RecordFailure(ctx context.Context, customerID uuid.UUID, maxAttempts int, now, lockUntil time.Time) (int, *time.Time, error) {
|
|
var rows []struct {
|
|
PinFailedAttempts int
|
|
PinLockedUntil *time.Time
|
|
}
|
|
// When an earlier lock has run out, this attempt is the first of a new series.
|
|
err := DBFromContext(ctx, r.db).WithContext(ctx).Raw(`
|
|
UPDATE customers SET
|
|
pin_failed_attempts = CASE
|
|
WHEN pin_locked_until IS NOT NULL AND pin_locked_until <= @now THEN 1
|
|
ELSE pin_failed_attempts + 1 END,
|
|
pin_locked_until = CASE
|
|
WHEN pin_locked_until IS NOT NULL AND pin_locked_until <= @now THEN NULL
|
|
WHEN pin_failed_attempts + 1 >= @max THEN @lock
|
|
ELSE pin_locked_until END
|
|
WHERE id = @id
|
|
RETURNING pin_failed_attempts, pin_locked_until`,
|
|
map[string]interface{}{"now": now, "max": maxAttempts, "lock": lockUntil, "id": customerID}).
|
|
Scan(&rows).Error
|
|
if err != nil {
|
|
return 0, nil, fmt.Errorf("failed to record a wrong PIN: %w", err)
|
|
}
|
|
if len(rows) == 0 {
|
|
return 0, nil, ErrPinCustomerNotFound
|
|
}
|
|
return rows[0].PinFailedAttempts, rows[0].PinLockedUntil, nil
|
|
}
|
|
|
|
func (r *customerPinRepository) ClearFailures(ctx context.Context, customerID uuid.UUID) error {
|
|
return DBFromContext(ctx, r.db).WithContext(ctx).Exec(`
|
|
UPDATE customers SET pin_failed_attempts = 0, pin_locked_until = NULL
|
|
WHERE id = ? AND (pin_failed_attempts <> 0 OR pin_locked_until IS NOT NULL)`, customerID).Error
|
|
}
|
|
|
|
func (r *customerPinRepository) InsertEvent(ctx context.Context, event CustomerSecurityEvent) error {
|
|
err := DBFromContext(ctx, r.db).WithContext(ctx).Exec(`
|
|
INSERT INTO customer_security_events (customer_id, event, actor_user, reason, ip_address, user_agent)
|
|
VALUES (?, ?, ?, ?, ?, ?)`,
|
|
event.CustomerID, event.Event, event.ActorUser, event.Reason, event.IPAddress, event.UserAgent).Error
|
|
if err != nil {
|
|
return fmt.Errorf("failed to record security event: %w", err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (r *customerPinRepository) ListEvents(ctx context.Context, customerID uuid.UUID, offset, limit int) ([]CustomerSecurityEvent, int64, error) {
|
|
db := DBFromContext(ctx, r.db).WithContext(ctx)
|
|
var total int64
|
|
if err := db.Table("customer_security_events").Where("customer_id = ?", customerID).Count(&total).Error; err != nil {
|
|
return nil, 0, fmt.Errorf("failed to count security events: %w", err)
|
|
}
|
|
var rows []struct {
|
|
ID string
|
|
CustomerID string
|
|
Event string
|
|
ActorUser *string
|
|
Reason *string
|
|
IPAddress *string
|
|
UserAgent *string
|
|
CreatedAt time.Time
|
|
}
|
|
err := db.Raw(`
|
|
SELECT id::text AS id, customer_id::text AS customer_id, event, actor_user::text AS actor_user,
|
|
reason, ip_address, user_agent, created_at
|
|
FROM customer_security_events WHERE customer_id = ?
|
|
ORDER BY created_at DESC, id DESC OFFSET ? LIMIT ?`, customerID, offset, limit).Scan(&rows).Error
|
|
if err != nil {
|
|
return nil, 0, fmt.Errorf("failed to list security events: %w", err)
|
|
}
|
|
events := make([]CustomerSecurityEvent, 0, len(rows))
|
|
for _, row := range rows {
|
|
e := CustomerSecurityEvent{Event: row.Event, Reason: row.Reason, IPAddress: row.IPAddress, UserAgent: row.UserAgent, CreatedAt: row.CreatedAt}
|
|
e.ID, _ = uuid.Parse(row.ID)
|
|
e.CustomerID, _ = uuid.Parse(row.CustomerID)
|
|
if row.ActorUser != nil {
|
|
if id, err := uuid.Parse(*row.ActorUser); err == nil {
|
|
e.ActorUser = &id
|
|
}
|
|
}
|
|
events = append(events, e)
|
|
}
|
|
return events, total, nil
|
|
}
|