feat(tasks): implement atomic claims and leases

This commit is contained in:
QiuSW
2026-07-26 16:16:36 +08:00
parent 49db5b8305
commit ce875af889
50 changed files with 7589 additions and 207 deletions
@@ -476,13 +476,16 @@ func getDeviceByID(
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, is_enabled, created_at, updated_at
pdd_version, last_seen_at, readiness_reported_at,
accessibility_enabled, pdd_installed, is_enabled,
created_at, updated_at
FROM devices
WHERE id = ?`,
deviceID,
@@ -495,6 +498,9 @@ func getDeviceByID(
&androidVersion,
&pddVersion,
&lastSeenAt,
&readinessAt,
&device.AccessibilityEnabled,
&device.PDDInstalled,
&device.IsEnabled,
&createdAt,
&updatedAt,
@@ -524,6 +530,13 @@ func getDeviceByID(
}
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
@@ -383,6 +383,9 @@ func TestAuthMigrationCanRollbackWithoutRebuildingPurchaseTasks(
if err != nil {
t.Fatalf("migration.New() error = %v", err)
}
if err := runner.Down(context.Background()); err != nil {
t.Fatalf("Down(v4) error = %v", err)
}
if err := runner.Down(context.Background()); err != nil {
t.Fatalf("Down(v3) error = %v", err)
}
@@ -399,9 +402,9 @@ func TestAuthMigrationCanRollbackWithoutRebuildingPurchaseTasks(
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)
t.Fatalf("Up(v3-v4) error = %v", err)
} else if applied != 2 {
t.Fatalf("Up(v3-v4) applied = %d, want 2", applied)
}
}
@@ -14,7 +14,7 @@ import (
sqlite3 "github.com/mattn/go-sqlite3"
)
const timestampLayout = time.RFC3339Nano
const storageTimestampLayout = "2006-01-02T15:04:05.000000000Z"
type queryRower interface {
QueryRowContext(context.Context, string, ...any) *sql.Row
@@ -52,7 +52,14 @@ func scanTask(scanner rowScanner) (domain.PurchaseTask, error) {
var createdByUserID sql.NullString
var sourceRef sql.NullString
var maxBudget sql.NullInt64
var claimedByUserID sql.NullString
var claimedByDeviceID sql.NullString
var claimTokenHash sql.NullString
var claimIssuedAt sql.NullString
var claimExpiresAt sql.NullString
var cancelReason sql.NullString
var cancelRequestedAt sql.NullString
var cancelRequestedByUserID sql.NullString
var canceledAt sql.NullString
var createdAt string
var updatedAt string
@@ -70,7 +77,15 @@ func scanTask(scanner rowScanner) (domain.PurchaseTask, error) {
&task.Currency,
&task.Status,
&task.Version,
&claimedByUserID,
&claimedByDeviceID,
&task.ClaimGeneration,
&claimTokenHash,
&claimIssuedAt,
&claimExpiresAt,
&cancelReason,
&cancelRequestedAt,
&cancelRequestedByUserID,
&canceledAt,
&createdAt,
&updatedAt,
@@ -87,15 +102,36 @@ func scanTask(scanner rowScanner) (domain.PurchaseTask, error) {
if maxBudget.Valid {
task.MaxBudgetCents = &maxBudget.Int64
}
if claimedByUserID.Valid {
task.ClaimedByUserID = &claimedByUserID.String
}
if claimedByDeviceID.Valid {
task.ClaimedByDeviceID = &claimedByDeviceID.String
}
if claimTokenHash.Valid {
task.ClaimTokenHash = &claimTokenHash.String
}
task.ClaimIssuedAt, err = parseNullableTimestamp(claimIssuedAt)
if err != nil {
return domain.PurchaseTask{}, err
}
task.ClaimExpiresAt, err = parseNullableTimestamp(claimExpiresAt)
if err != nil {
return domain.PurchaseTask{}, err
}
if cancelReason.Valid {
task.CancelReason = &cancelReason.String
}
if canceledAt.Valid {
value, err := parseTimestamp(canceledAt.String)
if err != nil {
return domain.PurchaseTask{}, err
}
task.CanceledAt = &value
task.CancelRequestedAt, err = parseNullableTimestamp(cancelRequestedAt)
if err != nil {
return domain.PurchaseTask{}, err
}
if cancelRequestedByUserID.Valid {
task.CancelRequestedByUserID = &cancelRequestedByUserID.String
}
task.CanceledAt, err = parseNullableTimestamp(canceledAt)
if err != nil {
return domain.PurchaseTask{}, err
}
task.CreatedAt, err = parseTimestamp(createdAt)
if err != nil {
@@ -108,6 +144,42 @@ func scanTask(scanner rowScanner) (domain.PurchaseTask, error) {
return task, nil
}
func scanExecution(scanner rowScanner) (domain.TaskExecution, error) {
var execution domain.TaskExecution
var lastHeartbeatAt sql.NullString
var finishedAt sql.NullString
var startedAt string
err := scanner.Scan(
&execution.ID,
&execution.TaskID,
&execution.AttemptNo,
&execution.ClaimGeneration,
&execution.UserID,
&execution.DeviceID,
&execution.CurrentStep,
&execution.OrderSubmitted,
&startedAt,
&lastHeartbeatAt,
&finishedAt,
)
if err != nil {
return domain.TaskExecution{}, err
}
execution.StartedAt, err = parseTimestamp(startedAt)
if err != nil {
return domain.TaskExecution{}, err
}
execution.LastHeartbeatAt, err = parseNullableTimestamp(lastHeartbeatAt)
if err != nil {
return domain.TaskExecution{}, err
}
execution.FinishedAt, err = parseNullableTimestamp(finishedAt)
if err != nil {
return domain.TaskExecution{}, err
}
return execution, nil
}
func getAssetByID(
ctx context.Context,
queryer queryRower,
@@ -144,7 +216,10 @@ func getTaskByID(
`SELECT
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
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 id = ?`,
creatorSubject,
@@ -218,11 +293,11 @@ func insertIdempotency(
}
func formatTimestamp(value time.Time) string {
return value.UTC().Format(timestampLayout)
return value.UTC().Format(storageTimestampLayout)
}
func parseTimestamp(value string) (time.Time, error) {
parsed, err := time.Parse(timestampLayout, value)
parsed, err := time.Parse(time.RFC3339Nano, value)
if err != nil {
return time.Time{}, fmt.Errorf(
"%w: invalid stored timestamp",
@@ -232,6 +307,17 @@ func parseTimestamp(value string) (time.Time, error) {
return parsed.UTC(), nil
}
func parseNullableTimestamp(value sql.NullString) (*time.Time, error) {
if !value.Valid {
return nil, nil
}
parsed, err := parseTimestamp(value.String)
if err != nil {
return nil, err
}
return &parsed, nil
}
func nullableString(value *string) any {
if value == nil {
return nil
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -190,7 +190,7 @@ func TestStoreAssetAndTaskLifecycleIsTransactionalAndIdempotent(
t.Fatalf("detail = %+v", detail)
}
canceled, err := store.CancelPendingTask(
canceled, err := store.CancelTask(
ctx,
"local-admin",
task.ID,
@@ -199,7 +199,7 @@ func TestStoreAssetAndTaskLifecycleIsTransactionalAndIdempotent(
testEvent(3, task.ID, "TASK_CANCELED", now.Add(time.Minute)),
)
if err != nil {
t.Fatalf("CancelPendingTask() error = %v", err)
t.Fatalf("CancelTask() error = %v", err)
}
if canceled.Status != domain.TaskStatusCanceled ||
canceled.Version != 2 ||
@@ -207,7 +207,7 @@ func TestStoreAssetAndTaskLifecycleIsTransactionalAndIdempotent(
*canceled.CancelReason != "no longer needed" {
t.Fatalf("canceled task = %+v", canceled)
}
_, err = store.CancelPendingTask(
_, err = store.CancelTask(
ctx,
"local-admin",
task.ID,
@@ -3,6 +3,7 @@ package sqlite
import (
"context"
"database/sql"
"errors"
"strings"
"time"
@@ -145,7 +146,10 @@ func (s *Store) ListTasks(
query.WriteString(`SELECT
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
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 = ?`)
arguments := []any{filter.CreatorSubject}
@@ -242,7 +246,9 @@ func (s *Store) GetTaskDetail(
}
rows, err := tx.QueryContext(
ctx,
`SELECT id, task_id, actor_user_id, event_type, message, occurred_at
`SELECT
id, task_id, actor_user_id, actor_device_id,
event_type, message, occurred_at
FROM task_events
WHERE task_id = ?
ORDER BY occurred_at ASC, id ASC`,
@@ -256,11 +262,13 @@ func (s *Store) GetTaskDetail(
for rows.Next() {
var event domain.TaskEvent
var actorUserID sql.NullString
var actorDeviceID sql.NullString
var occurredAt string
if err := rows.Scan(
&event.ID,
&event.TaskID,
&actorUserID,
&actorDeviceID,
&event.Type,
&event.Message,
&occurredAt,
@@ -270,6 +278,9 @@ func (s *Store) GetTaskDetail(
if actorUserID.Valid {
event.ActorUserID = &actorUserID.String
}
if actorDeviceID.Valid {
event.ActorDeviceID = &actorDeviceID.String
}
event.OccurredAt, err = parseTimestamp(occurredAt)
if err != nil {
return domain.TaskDetail{}, err
@@ -282,10 +293,31 @@ func (s *Store) GetTaskDetail(
if err := rows.Close(); err != nil {
return domain.TaskDetail{}, repositoryFailure(err)
}
execution, err := scanExecution(tx.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 task_id = ?
ORDER BY attempt_no DESC
LIMIT 1`,
taskID,
))
var executionPointer *domain.TaskExecution
if errors.Is(err, sql.ErrNoRows) {
executionPointer = nil
} else if err != nil {
return domain.TaskDetail{}, repositoryFailure(err)
} else {
executionPointer = &execution
}
detail := domain.TaskDetail{
Task: task,
Asset: asset,
Events: events,
Task: task,
Asset: asset,
Execution: executionPointer,
Events: events,
}
if err := tx.Commit(); err != nil {
return domain.TaskDetail{}, repositoryFailure(err)
@@ -293,7 +325,7 @@ func (s *Store) GetTaskDetail(
return detail, nil
}
func (s *Store) CancelPendingTask(
func (s *Store) CancelTask(
ctx context.Context,
creatorSubject string,
taskID string,
@@ -313,25 +345,66 @@ func (s *Store) CancelPendingTask(
if !domain.CanCancel(task.Status) {
return domain.PurchaseTask{}, usecase.ErrTaskStateConflict
}
result, err := tx.ExecContext(
ctx,
`UPDATE purchase_tasks
SET status = 'CANCELED',
version = version + 1,
cancel_reason = NULLIF(?, ''),
canceled_at = ?,
updated_at = ?
WHERE id = ?
AND creator_subject = ?
AND status = 'PENDING'
AND version = ?`,
reason,
formatTimestamp(canceledAt),
formatTimestamp(canceledAt),
taskID,
creatorSubject,
task.Version,
)
if domain.CanRequestCancel(task.Status) &&
task.CancelRequestedAt != nil {
if err := tx.Commit(); err != nil {
return domain.PurchaseTask{}, repositoryFailure(err)
}
return task, nil
}
var result sql.Result
if domain.CanAdminCancelImmediately(task.Status) {
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,
cancel_reason = NULLIF(?, ''),
cancel_requested_at = NULL,
cancel_requested_by_user_id = NULL,
canceled_at = ?,
updated_at = ?
WHERE id = ?
AND creator_subject = ?
AND status IN ('PENDING', 'CLAIMED')
AND version = ?`,
reason,
formatTimestamp(canceledAt),
formatTimestamp(canceledAt),
taskID,
creatorSubject,
task.Version,
)
} else {
event.Type = "TASK_CANCEL_REQUESTED"
event.Message = "task cancellation requested"
result, err = tx.ExecContext(
ctx,
`UPDATE purchase_tasks
SET version = version + 1,
cancel_reason = NULLIF(?, ''),
cancel_requested_at = ?,
cancel_requested_by_user_id = ?,
updated_at = ?
WHERE id = ?
AND creator_subject = ?
AND status IN ('RUNNING', 'WAITING_CONFIRMATION')
AND cancel_requested_at IS NULL
AND version = ?`,
reason,
formatTimestamp(canceledAt),
nullableString(event.ActorUserID),
formatTimestamp(canceledAt),
taskID,
creatorSubject,
task.Version,
)
}
if err != nil {
return domain.PurchaseTask{}, repositoryFailure(err)
}
@@ -363,11 +436,13 @@ func insertTaskEvent(
_, err := tx.ExecContext(
ctx,
`INSERT INTO task_events (
id, task_id, actor_user_id, event_type, message, occurred_at
) VALUES (?, ?, ?, ?, ?, ?)`,
id, task_id, actor_user_id, actor_device_id,
event_type, message, occurred_at
) VALUES (?, ?, ?, ?, ?, ?, ?)`,
event.ID,
event.TaskID,
nullableString(event.ActorUserID),
nullableString(event.ActorDeviceID),
event.Type,
event.Message,
formatTimestamp(event.OccurredAt),