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

1270 lines
31 KiB
Go
Raw Normal View History

package sqlite
import (
"context"
"database/sql"
"errors"
"time"
"cmroubao/backend-api/internal/domain"
"cmroubao/backend-api/internal/usecase"
)
const (
lifecycleOperationClaimNext = "CLAIM_NEXT"
lifecycleOperationStart = "START"
lifecycleOperationRelease = "RELEASE"
lifecycleOperationCancelAck = "CANCEL_ACK"
)
type lifecycleRequestRecord struct {
RequestHash string
ResultKind string
TaskID *string
ClaimGeneration *int64
ExecutionID *string
}
func (s *Store) RecordDeviceHeartbeat(
ctx context.Context,
update usecase.DeviceHeartbeatUpdate,
) (usecase.DeviceHeartbeatRecord, error) {
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return usecase.DeviceHeartbeatRecord{}, repositoryFailure(err)
}
defer func() { _ = tx.Rollback() }()
device, err := getDeviceByID(ctx, tx, update.DeviceID)
if err != nil {
if errors.Is(err, usecase.ErrRepositoryNotFound) {
return usecase.DeviceHeartbeatRecord{}, usecase.ErrClaimInvalid
}
return usecase.DeviceHeartbeatRecord{}, err
}
if !device.IsEnabled ||
device.BoundUserID == nil ||
*device.BoundUserID != update.UserID {
return usecase.DeviceHeartbeatRecord{}, usecase.ErrClaimInvalid
}
_, err = tx.ExecContext(
ctx,
`UPDATE devices
SET app_version = ?,
android_version = ?,
pdd_version = ?,
last_seen_at = ?,
readiness_reported_at = ?,
accessibility_enabled = ?,
pdd_installed = ?,
updated_at = ?
WHERE id = ?`,
update.AppVersion,
update.AndroidVersion,
update.PDDVersion,
formatTimestamp(update.ReportedAt),
formatTimestamp(update.ReportedAt),
update.AccessibilityEnabled,
update.PDDInstalled,
formatTimestamp(update.ReportedAt),
update.DeviceID,
)
if err != nil {
return usecase.DeviceHeartbeatRecord{}, repositoryFailure(err)
}
device, err = getDeviceByID(ctx, tx, update.DeviceID)
if err != nil {
return usecase.DeviceHeartbeatRecord{}, err
}
var activeTaskID string
err = tx.QueryRowContext(
ctx,
`SELECT id
FROM purchase_tasks
WHERE claimed_by_device_id = ?
AND (
status IN ('RUNNING', 'WAITING_CONFIRMATION')
OR (
status = 'CLAIMED'
AND claim_expires_at > ?
)
)
ORDER BY updated_at DESC, id DESC
LIMIT 1`,
update.DeviceID,
formatTimestamp(update.ReportedAt),
).Scan(&activeTaskID)
var activeTaskIDPointer *string
if errors.Is(err, sql.ErrNoRows) {
activeTaskIDPointer = nil
} else if err != nil {
return usecase.DeviceHeartbeatRecord{}, repositoryFailure(err)
} else {
activeTaskIDPointer = &activeTaskID
}
if err := tx.Commit(); err != nil {
return usecase.DeviceHeartbeatRecord{}, repositoryFailure(err)
}
return usecase.DeviceHeartbeatRecord{
Device: device,
ActiveTaskID: activeTaskIDPointer,
}, nil
}
func (s *Store) ClaimNext(
ctx context.Context,
request usecase.ClaimNextRepositoryRequest,
) (usecase.ClaimNextRepositoryResult, error) {
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return usecase.ClaimNextRepositoryResult{}, repositoryFailure(err)
}
defer func() { _ = tx.Rollback() }()
record, found, err := lookupLifecycleRequest(
ctx,
tx,
request.UserID,
request.DeviceID,
lifecycleOperationClaimNext,
request.IdempotencyKey,
)
if err != nil {
return usecase.ClaimNextRepositoryResult{}, err
}
if found {
return replayClaimNext(ctx, tx, request, record)
}
if err := ensureDeviceReady(ctx, tx, request); err != nil {
return usecase.ClaimNextRepositoryResult{}, err
}
active, err := hasActiveDeviceTask(
ctx,
tx,
request.DeviceID,
request.Now,
)
if err != nil {
return usecase.ClaimNextRepositoryResult{}, err
}
if active {
return usecase.ClaimNextRepositoryResult{},
usecase.ErrDeviceActiveTask
}
candidate, err := scanTask(tx.QueryRowContext(
ctx,
`SELECT
id, creator_subject, created_by_user_id, source_ref, title, description, sku,
image_asset_id, quantity, max_budget_cents, currency, status,
version, claimed_by_user_id, claimed_by_device_id, claim_generation,
claim_token_hash, claim_issued_at, claim_expires_at, cancel_reason,
cancel_requested_at, cancel_requested_by_user_id, canceled_at,
created_at, updated_at
FROM purchase_tasks
WHERE creator_subject = ?
AND (
status = 'PENDING'
OR (
status = 'CLAIMED'
AND claim_expires_at <= ?
)
)
ORDER BY
CASE
WHEN status = 'CLAIMED'
AND claimed_by_device_id = ?
THEN 0
ELSE 1
END,
created_at ASC,
id ASC
LIMIT 1`,
request.CreatorSubject,
formatTimestamp(request.Now),
request.DeviceID,
))
if errors.Is(err, sql.ErrNoRows) {
if err := insertLifecycleRequest(
ctx,
tx,
request.UserID,
request.DeviceID,
lifecycleOperationClaimNext,
request.IdempotencyKey,
request.RequestHash,
"NO_TASK",
nil,
nil,
nil,
request.Now,
); err != nil {
return usecase.ClaimNextRepositoryResult{}, err
}
if err := tx.Commit(); err != nil {
return usecase.ClaimNextRepositoryResult{},
repositoryFailure(err)
}
return usecase.ClaimNextRepositoryResult{}, nil
}
if err != nil {
return usecase.ClaimNextRepositoryResult{}, repositoryFailure(err)
}
event := request.Event
event.TaskID = candidate.ID
if candidate.Status == domain.TaskStatusClaimed {
event.Type = "TASK_RECLAIMED"
event.Message = "expired task claim reclaimed"
}
result, err := tx.ExecContext(
ctx,
`UPDATE purchase_tasks
SET status = 'CLAIMED',
version = version + 1,
claimed_by_user_id = ?,
claimed_by_device_id = ?,
claim_generation = claim_generation + 1,
claim_token_hash = ?,
claim_issued_at = ?,
claim_expires_at = ?,
cancel_reason = NULL,
cancel_requested_at = NULL,
cancel_requested_by_user_id = NULL,
canceled_at = NULL,
updated_at = ?
WHERE id = ?
AND version = ?
AND (
status = 'PENDING'
OR (
status = 'CLAIMED'
AND claim_expires_at <= ?
)
)`,
request.UserID,
request.DeviceID,
request.ClaimTokenHash,
formatTimestamp(request.Now),
formatTimestamp(request.ExpiresAt),
formatTimestamp(request.Now),
candidate.ID,
candidate.Version,
formatTimestamp(request.Now),
)
if err != nil {
if isUniqueConstraint(err, "purchase_tasks.claimed_by_device_id") {
return usecase.ClaimNextRepositoryResult{},
usecase.ErrDeviceActiveTask
}
return usecase.ClaimNextRepositoryResult{}, repositoryFailure(err)
}
affected, err := result.RowsAffected()
if err != nil {
return usecase.ClaimNextRepositoryResult{}, repositoryFailure(err)
}
if affected != 1 {
return usecase.ClaimNextRepositoryResult{},
usecase.ErrTaskStateConflict
}
if err := insertTaskEvent(ctx, tx, event); err != nil {
return usecase.ClaimNextRepositoryResult{}, err
}
claimed, err := getTaskByID(
ctx,
tx,
request.CreatorSubject,
candidate.ID,
)
if err != nil {
return usecase.ClaimNextRepositoryResult{}, err
}
taskID := claimed.ID
generation := claimed.ClaimGeneration
if err := insertLifecycleRequest(
ctx,
tx,
request.UserID,
request.DeviceID,
lifecycleOperationClaimNext,
request.IdempotencyKey,
request.RequestHash,
"TASK",
&taskID,
&generation,
nil,
request.Now,
); err != nil {
return usecase.ClaimNextRepositoryResult{}, err
}
if err := tx.Commit(); err != nil {
return usecase.ClaimNextRepositoryResult{}, repositoryFailure(err)
}
return usecase.ClaimNextRepositoryResult{Task: &claimed}, nil
}
func (s *Store) StartTask(
ctx context.Context,
request usecase.StartTaskRepositoryRequest,
) (usecase.StartTaskRepositoryResult, error) {
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return usecase.StartTaskRepositoryResult{}, repositoryFailure(err)
}
defer func() { _ = tx.Rollback() }()
record, found, err := lookupLifecycleRequest(
ctx,
tx,
request.UserID,
request.DeviceID,
lifecycleOperationStart,
request.IdempotencyKey,
)
if err != nil {
return usecase.StartTaskRepositoryResult{}, err
}
if found {
if record.RequestHash != request.RequestHash {
return usecase.StartTaskRepositoryResult{},
usecase.ErrIdempotencyConflict
}
if record.ResultKind != "EXECUTION" ||
record.TaskID == nil ||
record.ExecutionID == nil ||
*record.TaskID != request.TaskID {
return usecase.StartTaskRepositoryResult{},
usecase.ErrRepositoryInvariant
}
task, err := getLifecycleTask(ctx, tx, request.TaskID)
if err != nil {
return usecase.StartTaskRepositoryResult{}, err
}
execution, err := getExecutionByID(
ctx,
tx,
*record.ExecutionID,
)
if err != nil {
return usecase.StartTaskRepositoryResult{}, err
}
if err := tx.Commit(); err != nil {
return usecase.StartTaskRepositoryResult{},
repositoryFailure(err)
}
return usecase.StartTaskRepositoryResult{
Task: task,
Execution: execution,
Replayed: true,
}, nil
}
task, err := getClaimProtectedTask(ctx, tx, request.TaskID)
if err != nil {
return usecase.StartTaskRepositoryResult{}, err
}
if err := validateClaim(
task,
request.UserID,
request.DeviceID,
request.ClaimGeneration,
request.ClaimTokenHash,
request.Now,
); err != nil {
return usecase.StartTaskRepositoryResult{}, err
}
if !domain.CanStart(task.Status) {
return usecase.StartTaskRepositoryResult{},
usecase.ErrTaskStateConflict
}
if task.Version != request.ExpectedVersion {
return usecase.StartTaskRepositoryResult{},
usecase.ErrTaskVersionConflict
}
err = tx.QueryRowContext(
ctx,
`SELECT COALESCE(MAX(attempt_no), 0) + 1
FROM task_executions
WHERE task_id = ?`,
request.TaskID,
).Scan(&request.Execution.AttemptNo)
if err != nil {
return usecase.StartTaskRepositoryResult{}, repositoryFailure(err)
}
request.Execution.LastHeartbeatAt = &request.Now
_, err = tx.ExecContext(
ctx,
`INSERT INTO task_executions (
id, task_id, attempt_no, claim_generation, user_id, device_id,
current_step, last_heartbeat_at, order_submitted, started_at,
finished_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, 0, ?, NULL)`,
request.Execution.ID,
request.Execution.TaskID,
request.Execution.AttemptNo,
request.Execution.ClaimGeneration,
request.Execution.UserID,
request.Execution.DeviceID,
request.Execution.CurrentStep,
formatTimestamp(request.Now),
formatTimestamp(request.Execution.StartedAt),
)
if err != nil {
return usecase.StartTaskRepositoryResult{}, repositoryFailure(err)
}
result, err := tx.ExecContext(
ctx,
`UPDATE purchase_tasks
SET status = 'RUNNING',
version = version + 1,
claim_expires_at = ?,
updated_at = ?
WHERE id = ?
AND status = 'CLAIMED'
AND version = ?
AND claimed_by_user_id = ?
AND claimed_by_device_id = ?
AND claim_generation = ?
AND claim_token_hash = ?
AND claim_expires_at > ?`,
formatTimestamp(request.ExpiresAt),
formatTimestamp(request.Now),
request.TaskID,
request.ExpectedVersion,
request.UserID,
request.DeviceID,
request.ClaimGeneration,
request.ClaimTokenHash,
formatTimestamp(request.Now),
)
if err != nil {
return usecase.StartTaskRepositoryResult{}, repositoryFailure(err)
}
affected, err := result.RowsAffected()
if err != nil {
return usecase.StartTaskRepositoryResult{}, repositoryFailure(err)
}
if affected != 1 {
return usecase.StartTaskRepositoryResult{},
usecase.ErrTaskVersionConflict
}
if err := insertTaskEvent(ctx, tx, request.Event); err != nil {
return usecase.StartTaskRepositoryResult{}, err
}
executionID := request.Execution.ID
taskID := request.TaskID
generation := request.ClaimGeneration
if err := insertLifecycleRequest(
ctx,
tx,
request.UserID,
request.DeviceID,
lifecycleOperationStart,
request.IdempotencyKey,
request.RequestHash,
"EXECUTION",
&taskID,
&generation,
&executionID,
request.Now,
); err != nil {
return usecase.StartTaskRepositoryResult{}, err
}
task, err = getLifecycleTask(ctx, tx, request.TaskID)
if err != nil {
return usecase.StartTaskRepositoryResult{}, err
}
if err := tx.Commit(); err != nil {
return usecase.StartTaskRepositoryResult{}, repositoryFailure(err)
}
return usecase.StartTaskRepositoryResult{
Task: task,
Execution: request.Execution,
}, nil
}
func (s *Store) HeartbeatTask(
ctx context.Context,
request usecase.TaskHeartbeatRepositoryRequest,
) (usecase.TaskHeartbeatRepositoryResult, error) {
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return usecase.TaskHeartbeatRepositoryResult{}, repositoryFailure(err)
}
defer func() { _ = tx.Rollback() }()
task, err := getClaimProtectedTask(ctx, tx, request.TaskID)
if err != nil {
return usecase.TaskHeartbeatRepositoryResult{}, err
}
if err := validateClaimOwner(
task,
request.UserID,
request.DeviceID,
request.ClaimGeneration,
request.ClaimTokenHash,
); err != nil {
return usecase.TaskHeartbeatRepositoryResult{}, err
}
if !domain.CanHeartbeat(task.Status) {
return usecase.TaskHeartbeatRepositoryResult{},
usecase.ErrTaskStateConflict
}
expired := !task.ClaimExpiresAt.After(request.Now)
if expired && request.Step != "SAFE_STOPPED" {
return usecase.TaskHeartbeatRepositoryResult{},
usecase.ErrClaimExpired
}
execution, err := getExecutionByID(ctx, tx, request.ExecutionID)
if err != nil {
return usecase.TaskHeartbeatRepositoryResult{}, err
}
if execution.TaskID != request.TaskID ||
execution.UserID != request.UserID ||
execution.DeviceID != request.DeviceID ||
execution.ClaimGeneration != request.ClaimGeneration ||
execution.FinishedAt != nil {
return usecase.TaskHeartbeatRepositoryResult{},
usecase.ErrExecutionMismatch
}
expiresAt := *task.ClaimExpiresAt
if !expired &&
request.MinimumExpiry.After(expiresAt) {
expiresAt = request.MinimumExpiry
}
result, err := tx.ExecContext(
ctx,
`UPDATE task_executions
SET current_step = ?,
last_heartbeat_at = ?
WHERE id = ?
AND finished_at IS NULL`,
request.Step,
formatTimestamp(request.Now),
request.ExecutionID,
)
if err != nil {
return usecase.TaskHeartbeatRepositoryResult{}, repositoryFailure(err)
}
affected, err := result.RowsAffected()
if err != nil {
return usecase.TaskHeartbeatRepositoryResult{}, repositoryFailure(err)
}
if affected != 1 {
return usecase.TaskHeartbeatRepositoryResult{},
usecase.ErrExecutionMismatch
}
result, err = tx.ExecContext(
ctx,
`UPDATE purchase_tasks
SET version = version + 1,
claim_expires_at = ?,
updated_at = ?
WHERE id = ?
AND status IN ('RUNNING', 'WAITING_CONFIRMATION')
AND claimed_by_user_id = ?
AND claimed_by_device_id = ?
AND claim_generation = ?
AND claim_token_hash = ?`,
formatTimestamp(expiresAt),
formatTimestamp(request.Now),
request.TaskID,
request.UserID,
request.DeviceID,
request.ClaimGeneration,
request.ClaimTokenHash,
)
if err != nil {
return usecase.TaskHeartbeatRepositoryResult{}, repositoryFailure(err)
}
affected, err = result.RowsAffected()
if err != nil {
return usecase.TaskHeartbeatRepositoryResult{}, repositoryFailure(err)
}
if affected != 1 {
return usecase.TaskHeartbeatRepositoryResult{},
usecase.ErrTaskVersionConflict
}
_, err = tx.ExecContext(
ctx,
`UPDATE devices
SET last_seen_at = ?,
updated_at = ?
WHERE id = ?`,
formatTimestamp(request.Now),
formatTimestamp(request.Now),
request.DeviceID,
)
if err != nil {
return usecase.TaskHeartbeatRepositoryResult{}, repositoryFailure(err)
}
task, err = getLifecycleTask(ctx, tx, request.TaskID)
if err != nil {
return usecase.TaskHeartbeatRepositoryResult{}, err
}
execution, err = getExecutionByID(ctx, tx, request.ExecutionID)
if err != nil {
return usecase.TaskHeartbeatRepositoryResult{}, err
}
if err := tx.Commit(); err != nil {
return usecase.TaskHeartbeatRepositoryResult{}, repositoryFailure(err)
}
return usecase.TaskHeartbeatRepositoryResult{
Task: task,
Execution: execution,
CancelRequested: task.CancelRequestedAt != nil,
}, nil
}
func (s *Store) GetActiveClaimTask(
ctx context.Context,
request usecase.TaskClaimRepositoryRequest,
) (domain.PurchaseTask, error) {
task, err := getClaimProtectedTask(ctx, s.db, request.TaskID)
if err != nil {
return domain.PurchaseTask{}, err
}
if err := validateClaim(
task,
request.UserID,
request.DeviceID,
request.ClaimGeneration,
request.ClaimTokenHash,
request.Now,
); err != nil {
return domain.PurchaseTask{}, err
}
if !domain.IsActiveTaskStatus(task.Status) {
return domain.PurchaseTask{}, usecase.ErrTaskStateConflict
}
return task, nil
}
func (s *Store) ReleaseTask(
ctx context.Context,
request usecase.ReleaseTaskRepositoryRequest,
) (domain.PurchaseTask, bool, error) {
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return domain.PurchaseTask{}, false, repositoryFailure(err)
}
defer func() { _ = tx.Rollback() }()
record, found, err := lookupLifecycleRequest(
ctx,
tx,
request.UserID,
request.DeviceID,
lifecycleOperationRelease,
request.IdempotencyKey,
)
if err != nil {
return domain.PurchaseTask{}, false, err
}
if found {
return replayTerminalTransition(
ctx,
tx,
request.RequestHash,
request.TaskID,
request.ClaimGeneration,
domain.TaskStatusPending,
record,
)
}
task, err := getClaimProtectedTask(ctx, tx, request.TaskID)
if err != nil {
return domain.PurchaseTask{}, false, err
}
if err := validateClaim(
task,
request.UserID,
request.DeviceID,
request.ClaimGeneration,
request.ClaimTokenHash,
request.Now,
); err != nil {
return domain.PurchaseTask{}, false, err
}
if !domain.CanRelease(task.Status) {
return domain.PurchaseTask{}, false, usecase.ErrTaskStateConflict
}
if task.Version != request.ExpectedVersion {
return domain.PurchaseTask{}, false,
usecase.ErrTaskVersionConflict
}
result, err := tx.ExecContext(
ctx,
`UPDATE purchase_tasks
SET status = 'PENDING',
version = version + 1,
claimed_by_user_id = NULL,
claimed_by_device_id = NULL,
claim_token_hash = NULL,
claim_issued_at = NULL,
claim_expires_at = NULL,
updated_at = ?
WHERE id = ?
AND status = 'CLAIMED'
AND version = ?
AND claimed_by_user_id = ?
AND claimed_by_device_id = ?
AND claim_generation = ?
AND claim_token_hash = ?
AND claim_expires_at > ?`,
formatTimestamp(request.Now),
request.TaskID,
request.ExpectedVersion,
request.UserID,
request.DeviceID,
request.ClaimGeneration,
request.ClaimTokenHash,
formatTimestamp(request.Now),
)
if err != nil {
return domain.PurchaseTask{}, false, repositoryFailure(err)
}
affected, err := result.RowsAffected()
if err != nil {
return domain.PurchaseTask{}, false, repositoryFailure(err)
}
if affected != 1 {
return domain.PurchaseTask{}, false,
usecase.ErrTaskVersionConflict
}
if err := insertTaskEvent(ctx, tx, request.Event); err != nil {
return domain.PurchaseTask{}, false, err
}
taskID := request.TaskID
generation := request.ClaimGeneration
if err := insertLifecycleRequest(
ctx,
tx,
request.UserID,
request.DeviceID,
lifecycleOperationRelease,
request.IdempotencyKey,
request.RequestHash,
"TASK",
&taskID,
&generation,
nil,
request.Now,
); err != nil {
return domain.PurchaseTask{}, false, err
}
task, err = getLifecycleTask(ctx, tx, request.TaskID)
if err != nil {
return domain.PurchaseTask{}, false, err
}
if err := tx.Commit(); err != nil {
return domain.PurchaseTask{}, false, repositoryFailure(err)
}
return task, false, nil
}
func (s *Store) AcknowledgeTaskCancellation(
ctx context.Context,
request usecase.CancelAcknowledgementRepositoryRequest,
) (domain.PurchaseTask, bool, error) {
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return domain.PurchaseTask{}, false, repositoryFailure(err)
}
defer func() { _ = tx.Rollback() }()
record, found, err := lookupLifecycleRequest(
ctx,
tx,
request.UserID,
request.DeviceID,
lifecycleOperationCancelAck,
request.IdempotencyKey,
)
if err != nil {
return domain.PurchaseTask{}, false, err
}
if found {
return replayTerminalTransition(
ctx,
tx,
request.RequestHash,
request.TaskID,
request.ClaimGeneration,
domain.TaskStatusCanceled,
record,
)
}
task, err := getClaimProtectedTask(ctx, tx, request.TaskID)
if err != nil {
return domain.PurchaseTask{}, false, err
}
if err := validateClaimOwner(
task,
request.UserID,
request.DeviceID,
request.ClaimGeneration,
request.ClaimTokenHash,
); err != nil {
return domain.PurchaseTask{}, false, err
}
if !domain.CanAcknowledgeCancel(task.Status) ||
task.CancelRequestedAt == nil {
return domain.PurchaseTask{}, false, usecase.ErrTaskStateConflict
}
if task.Version != request.ExpectedVersion {
return domain.PurchaseTask{}, false,
usecase.ErrTaskVersionConflict
}
execution, err := getExecutionByID(ctx, tx, request.ExecutionID)
if err != nil {
return domain.PurchaseTask{}, false, err
}
if execution.TaskID != request.TaskID ||
execution.UserID != request.UserID ||
execution.DeviceID != request.DeviceID ||
execution.ClaimGeneration != request.ClaimGeneration ||
execution.FinishedAt != nil {
return domain.PurchaseTask{}, false,
usecase.ErrExecutionMismatch
}
result, err := tx.ExecContext(
ctx,
`UPDATE purchase_tasks
SET status = 'CANCELED',
version = version + 1,
claimed_by_user_id = NULL,
claimed_by_device_id = NULL,
claim_token_hash = NULL,
claim_issued_at = NULL,
claim_expires_at = NULL,
canceled_at = ?,
updated_at = ?
WHERE id = ?
AND status IN ('RUNNING', 'WAITING_CONFIRMATION')
AND version = ?
AND cancel_requested_at IS NOT NULL
AND claimed_by_user_id = ?
AND claimed_by_device_id = ?
AND claim_generation = ?
AND claim_token_hash = ?`,
formatTimestamp(request.Now),
formatTimestamp(request.Now),
request.TaskID,
request.ExpectedVersion,
request.UserID,
request.DeviceID,
request.ClaimGeneration,
request.ClaimTokenHash,
)
if err != nil {
return domain.PurchaseTask{}, false, repositoryFailure(err)
}
affected, err := result.RowsAffected()
if err != nil {
return domain.PurchaseTask{}, false, repositoryFailure(err)
}
if affected != 1 {
return domain.PurchaseTask{}, false,
usecase.ErrTaskVersionConflict
}
result, err = tx.ExecContext(
ctx,
`UPDATE task_executions
SET finished_at = ?,
last_heartbeat_at = ?
WHERE id = ?
AND finished_at IS NULL`,
formatTimestamp(request.Now),
formatTimestamp(request.Now),
request.ExecutionID,
)
if err != nil {
return domain.PurchaseTask{}, false, repositoryFailure(err)
}
affected, err = result.RowsAffected()
if err != nil {
return domain.PurchaseTask{}, false, repositoryFailure(err)
}
if affected != 1 {
return domain.PurchaseTask{}, false,
usecase.ErrExecutionMismatch
}
if err := insertTaskEvent(ctx, tx, request.Event); err != nil {
return domain.PurchaseTask{}, false, err
}
taskID := request.TaskID
generation := request.ClaimGeneration
executionID := request.ExecutionID
if err := insertLifecycleRequest(
ctx,
tx,
request.UserID,
request.DeviceID,
lifecycleOperationCancelAck,
request.IdempotencyKey,
request.RequestHash,
"TASK",
&taskID,
&generation,
&executionID,
request.Now,
); err != nil {
return domain.PurchaseTask{}, false, err
}
task, err = getLifecycleTask(ctx, tx, request.TaskID)
if err != nil {
return domain.PurchaseTask{}, false, err
}
if err := tx.Commit(); err != nil {
return domain.PurchaseTask{}, false, repositoryFailure(err)
}
return task, false, nil
}
func replayClaimNext(
ctx context.Context,
tx *sql.Tx,
request usecase.ClaimNextRepositoryRequest,
record lifecycleRequestRecord,
) (usecase.ClaimNextRepositoryResult, error) {
if record.RequestHash != request.RequestHash {
return usecase.ClaimNextRepositoryResult{},
usecase.ErrIdempotencyConflict
}
if record.ResultKind == "NO_TASK" {
if err := tx.Commit(); err != nil {
return usecase.ClaimNextRepositoryResult{},
repositoryFailure(err)
}
return usecase.ClaimNextRepositoryResult{Replayed: true}, nil
}
if record.ResultKind != "TASK" ||
record.TaskID == nil ||
record.ClaimGeneration == nil {
return usecase.ClaimNextRepositoryResult{},
usecase.ErrRepositoryInvariant
}
task, err := getTaskByID(
ctx,
tx,
request.CreatorSubject,
*record.TaskID,
)
if err != nil {
return usecase.ClaimNextRepositoryResult{}, err
}
if !domain.IsActiveTaskStatus(task.Status) ||
task.ClaimedByUserID == nil ||
*task.ClaimedByUserID != request.UserID ||
task.ClaimedByDeviceID == nil ||
*task.ClaimedByDeviceID != request.DeviceID ||
task.ClaimGeneration != *record.ClaimGeneration ||
task.ClaimTokenHash == nil ||
*task.ClaimTokenHash != request.ClaimTokenHash ||
task.ClaimExpiresAt == nil ||
!task.ClaimExpiresAt.After(request.Now) {
return usecase.ClaimNextRepositoryResult{},
usecase.ErrClaimReplayExpired
}
if err := tx.Commit(); err != nil {
return usecase.ClaimNextRepositoryResult{},
repositoryFailure(err)
}
return usecase.ClaimNextRepositoryResult{
Task: &task,
Replayed: true,
}, nil
}
func ensureDeviceReady(
ctx context.Context,
tx *sql.Tx,
request usecase.ClaimNextRepositoryRequest,
) error {
device, err := getDeviceByID(ctx, tx, request.DeviceID)
if err != nil {
if errors.Is(err, usecase.ErrRepositoryNotFound) {
return usecase.ErrClaimInvalid
}
return err
}
if !device.IsEnabled ||
device.BoundUserID == nil ||
*device.BoundUserID != request.UserID {
return usecase.ErrClaimInvalid
}
if device.ReadinessAt == nil ||
device.ReadinessAt.Before(request.ReadinessAfter) ||
!device.AccessibilityEnabled ||
!device.PDDInstalled {
return usecase.ErrDeviceNotReady
}
return nil
}
func hasActiveDeviceTask(
ctx context.Context,
tx *sql.Tx,
deviceID string,
now time.Time,
) (bool, error) {
var active int
err := tx.QueryRowContext(
ctx,
`SELECT EXISTS (
SELECT 1
FROM purchase_tasks
WHERE claimed_by_device_id = ?
AND (
status IN ('RUNNING', 'WAITING_CONFIRMATION')
OR (
status = 'CLAIMED'
AND claim_expires_at > ?
)
)
)`,
deviceID,
formatTimestamp(now),
).Scan(&active)
if err != nil {
return false, repositoryFailure(err)
}
return active == 1, nil
}
func validateClaim(
task domain.PurchaseTask,
userID string,
deviceID string,
generation int64,
tokenHash string,
now time.Time,
) error {
if err := validateClaimOwner(
task,
userID,
deviceID,
generation,
tokenHash,
); err != nil {
return err
}
if task.ClaimExpiresAt == nil || !task.ClaimExpiresAt.After(now) {
return usecase.ErrClaimExpired
}
return nil
}
func validateClaimOwner(
task domain.PurchaseTask,
userID string,
deviceID string,
generation int64,
tokenHash string,
) error {
if task.ClaimedByUserID == nil ||
*task.ClaimedByUserID != userID ||
task.ClaimedByDeviceID == nil ||
*task.ClaimedByDeviceID != deviceID ||
task.ClaimGeneration != generation ||
task.ClaimTokenHash == nil ||
*task.ClaimTokenHash != tokenHash {
return usecase.ErrClaimInvalid
}
if task.ClaimExpiresAt == nil {
return usecase.ErrClaimExpired
}
return nil
}
func replayTerminalTransition(
ctx context.Context,
tx *sql.Tx,
requestHash string,
taskID string,
generation int64,
expectedStatus domain.TaskStatus,
record lifecycleRequestRecord,
) (domain.PurchaseTask, bool, error) {
if record.RequestHash != requestHash {
return domain.PurchaseTask{}, false,
usecase.ErrIdempotencyConflict
}
if record.ResultKind != "TASK" ||
record.TaskID == nil ||
record.ClaimGeneration == nil ||
*record.TaskID != taskID ||
*record.ClaimGeneration != generation {
return domain.PurchaseTask{}, false,
usecase.ErrRepositoryInvariant
}
task, err := getLifecycleTask(ctx, tx, taskID)
if err != nil {
return domain.PurchaseTask{}, false, err
}
if task.Status != expectedStatus ||
task.ClaimGeneration != generation {
return domain.PurchaseTask{}, false,
usecase.ErrTaskStateConflict
}
if err := tx.Commit(); err != nil {
return domain.PurchaseTask{}, false, repositoryFailure(err)
}
return task, true, nil
}
func lookupLifecycleRequest(
ctx context.Context,
tx *sql.Tx,
userID string,
deviceID string,
operation string,
idempotencyKey string,
) (lifecycleRequestRecord, bool, error) {
var record lifecycleRequestRecord
var taskID sql.NullString
var generation sql.NullInt64
var executionID sql.NullString
err := tx.QueryRowContext(
ctx,
`SELECT
request_sha256, result_kind, task_id,
claim_generation, execution_id
FROM lifecycle_requests
WHERE user_id = ?
AND device_id = ?
AND operation = ?
AND idempotency_key = ?`,
userID,
deviceID,
operation,
idempotencyKey,
).Scan(
&record.RequestHash,
&record.ResultKind,
&taskID,
&generation,
&executionID,
)
if errors.Is(err, sql.ErrNoRows) {
return lifecycleRequestRecord{}, false, nil
}
if err != nil {
return lifecycleRequestRecord{}, false, repositoryFailure(err)
}
if taskID.Valid {
record.TaskID = &taskID.String
}
if generation.Valid {
record.ClaimGeneration = &generation.Int64
}
if executionID.Valid {
record.ExecutionID = &executionID.String
}
return record, true, nil
}
func insertLifecycleRequest(
ctx context.Context,
tx *sql.Tx,
userID string,
deviceID string,
operation string,
idempotencyKey string,
requestHash string,
resultKind string,
taskID *string,
generation *int64,
executionID *string,
now time.Time,
) error {
_, err := tx.ExecContext(
ctx,
`INSERT INTO lifecycle_requests (
user_id, device_id, operation, idempotency_key, request_sha256,
result_kind, task_id, claim_generation, execution_id, created_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
userID,
deviceID,
operation,
idempotencyKey,
requestHash,
resultKind,
nullableString(taskID),
nullableInt64(generation),
nullableString(executionID),
formatTimestamp(now),
)
if err != nil {
return repositoryFailure(err)
}
return nil
}
func getLifecycleTask(
ctx context.Context,
queryer queryRower,
taskID string,
) (domain.PurchaseTask, error) {
task, err := scanTask(queryer.QueryRowContext(
ctx,
`SELECT
id, creator_subject, created_by_user_id, source_ref, title, description, sku,
image_asset_id, quantity, max_budget_cents, currency, status,
version, claimed_by_user_id, claimed_by_device_id, claim_generation,
claim_token_hash, claim_issued_at, claim_expires_at, cancel_reason,
cancel_requested_at, cancel_requested_by_user_id, canceled_at,
created_at, updated_at
FROM purchase_tasks
WHERE id = ?`,
taskID,
))
if errors.Is(err, sql.ErrNoRows) {
return domain.PurchaseTask{}, usecase.ErrRepositoryNotFound
}
if err != nil {
return domain.PurchaseTask{}, repositoryFailure(err)
}
return task, nil
}
func getClaimProtectedTask(
ctx context.Context,
queryer queryRower,
taskID string,
) (domain.PurchaseTask, error) {
task, err := getLifecycleTask(ctx, queryer, taskID)
if errors.Is(err, usecase.ErrRepositoryNotFound) {
return domain.PurchaseTask{}, usecase.ErrClaimInvalid
}
return task, err
}
func getExecutionByID(
ctx context.Context,
queryer queryRower,
executionID string,
) (domain.TaskExecution, error) {
execution, err := scanExecution(queryer.QueryRowContext(
ctx,
`SELECT
id, task_id, attempt_no, claim_generation, user_id, device_id,
current_step, order_submitted, started_at, last_heartbeat_at,
finished_at
FROM task_executions
WHERE id = ?`,
executionID,
))
if errors.Is(err, sql.ErrNoRows) {
return domain.TaskExecution{}, usecase.ErrExecutionMismatch
}
if err != nil {
return domain.TaskExecution{}, repositoryFailure(err)
}
return execution, nil
}
var _ usecase.LifecycleRepository = (*Store)(nil)