feat(admin): add atomic task claim leases

This commit is contained in:
QiuSW
2026-08-04 22:33:16 +08:00
parent 526af1eb31
commit 5f060f60ee
20 changed files with 2923 additions and 49 deletions
+830
View File
@@ -0,0 +1,830 @@
package taskclaim
import (
"context"
"crypto/rand"
"crypto/sha256"
"database/sql"
"encoding/hex"
"errors"
"fmt"
"io"
"math"
"math/big"
"strings"
"sync"
"time"
"cmbuyer/admin/internal/deviceauth"
)
const writeTimeout = 2 * time.Second
type Store struct {
database *sql.DB
secret []byte
leaseTTL time.Duration
now func() time.Time
random io.Reader
randomMu sync.Mutex
writeGate chan struct{}
// The unexported linearization hooks let package tests coordinate real SQLite
// transactions at the first write. Production construction always leaves them nil.
beforeLinearization func()
afterLinearization func()
}
func NewStore(database *sql.DB, secret []byte, leaseTTL time.Duration) (*Store, error) {
if database == nil {
return nil, errors.New("task claim database is required")
}
if len(secret) != sha256.Size {
return nil, errors.New("task claim secret must be 32 bytes")
}
if leaseTTL <= 0 {
return nil, errors.New("task claim lease TTL must be positive")
}
if _, err := database.Exec("SELECT attempt_id, claim_nonce, claim_token_sha256 FROM purchase_attempt_claims LIMIT 1"); err != nil {
return nil, errors.New("task claim migration is not available")
}
store := &Store{
database: database, secret: append([]byte(nil), secret...), leaseTTL: leaseTTL,
now: time.Now, random: rand.Reader, writeGate: make(chan struct{}, 1),
}
if err := store.validateSecretIsolation(); err != nil {
return nil, err
}
if err := store.validateStoredClaims(context.Background()); err != nil {
return nil, err
}
return store, nil
}
// validateSecretIsolation ensures the HMAC key cannot also authenticate a device. The session
// secret comparison is performed while parsing configuration, before either secret is discarded.
func (store *Store) validateSecretIsolation() error {
digest := sha256.Sum256(store.secret)
var count int
if err := store.database.QueryRow(`SELECT COUNT(*) FROM device_credentials WHERE token_sha256 = ?`, digest[:]).Scan(&count); err != nil {
return errors.New("validate task claim secret isolation")
}
if count != 0 {
return errors.New("task claim secret must be isolated from device credentials")
}
return nil
}
// validateStoredClaims covers open and closed claims. Replacing the secret must fail startup;
// silently signing a new token would destroy idempotent recovery and the ownership audit chain.
func (store *Store) validateStoredClaims(ctx context.Context) error {
rows, err := store.database.QueryContext(ctx, `SELECT claims.claimed_by_device_id, claims.task_id, claims.authorization_id,
claims.attempt_id, claims.claim_generation, claims.claim_nonce, typeof(claims.claim_nonce), length(claims.claim_nonce),
claims.claim_token_sha256, typeof(claims.claim_token_sha256), length(claims.claim_token_sha256),
claims.authorization_task_version, claims.goods_id, claims.sku_color, claims.sku_size,
claims.quantity, claims.total_price_cap, claims.authorization_expires_at, claims.closed_at,
attempts.claim_generation, attempts.status, authorizations.status, tasks.status
FROM purchase_attempt_claims AS claims
LEFT JOIN purchase_attempts AS attempts ON attempts.id = claims.attempt_id
LEFT JOIN order_authorizations AS authorizations ON authorizations.id = claims.authorization_id
LEFT JOIN tasks ON tasks.id = claims.task_id
ORDER BY claims.attempt_id`)
if err != nil {
return errors.New("validate stored task claims")
}
defer rows.Close()
for rows.Next() {
var deviceID, taskID, authorizationID, attemptID string
var generation, authorizationTaskVersion, quantity int
var nonce, storedHash []byte
var nonceType, hashType, goodsID, color, size, price, expires string
var nonceLength, hashLength int
var closed, attemptStatus, authorizationStatus, taskStatus sql.NullString
var attemptGeneration sql.NullInt64
if err := rows.Scan(&deviceID, &taskID, &authorizationID, &attemptID, &generation,
&nonce, &nonceType, &nonceLength, &storedHash, &hashType, &hashLength,
&authorizationTaskVersion, &goodsID, &color, &size, &quantity, &price, &expires, &closed,
&attemptGeneration, &attemptStatus, &authorizationStatus, &taskStatus); err != nil {
return errors.New("validate stored task claims")
}
if !deviceauth.ValidDeviceID(deviceID) || !validUUID(taskID) || !validUUID(authorizationID) || !validUUID(attemptID) ||
generation <= 0 || nonceType != "blob" || nonceLength != sha256.Size || len(nonce) != sha256.Size ||
hashType != "blob" || hashLength != sha256.Size || len(storedHash) != sha256.Size ||
authorizationTaskVersion <= 0 || !digitsOnly(goodsID) || color == "" || size == "" || quantity <= 0 ||
!canonicalMoney(price) || !validCanonicalTime(expires) || (closed.Valid && !validCanonicalTime(closed.String)) ||
!attemptGeneration.Valid || attemptGeneration.Int64 != int64(generation) ||
!validAttemptStatus(attemptStatus) || !validAuthorizationStatus(authorizationStatus) || !validTaskStatus(taskStatus) {
return errors.New("stored task claim metadata is invalid")
}
token := deriveToken(store.secret, deviceID, taskID, authorizationID, attemptID, generation, nonce)
if !matchingHash(tokenHash(token), storedHash) {
return errors.New("task claim secret does not match stored claims")
}
}
if err := rows.Err(); err != nil {
return errors.New("validate stored task claims")
}
return nil
}
func (store *Store) ClaimNext(ctx context.Context, deviceID string, command ClaimCommand) (ClaimResponse, bool, error) {
if !deviceauth.ValidDeviceID(deviceID) || !validUUID(command.SessionID) || !validUUID(command.ClaimRequestID) {
return ClaimResponse{}, false, ErrInvalid
}
writeCtx, cancel := context.WithTimeout(ctx, writeTimeout)
defer cancel()
select {
case store.writeGate <- struct{}{}:
defer func() { <-store.writeGate }()
case <-writeCtx.Done():
return ClaimResponse{}, false, writeCtx.Err()
}
transaction, err := store.database.BeginTx(writeCtx, nil)
if err != nil {
return ClaimResponse{}, false, err
}
defer transaction.Rollback()
// This must be the transaction's first database statement. The no-op conditional UPDATE takes
// SQLite's write position and linearizes a concurrent credential revocation before any replay,
// EMPTY response, conflict response, candidate read, or other business write is possible.
if store.beforeLinearization != nil {
store.beforeLinearization()
}
active, err := transaction.ExecContext(writeCtx, `UPDATE device_credentials SET status = status
WHERE device_id = ? AND status = 'ACTIVE' AND revoked_at IS NULL`, deviceID)
if err != nil {
return ClaimResponse{}, false, err
}
if ok, err := exactlyOne(active); err != nil {
return ClaimResponse{}, false, err
} else if !ok {
return ClaimResponse{}, false, ErrDeviceInactive
}
if store.afterLinearization != nil {
store.afterLinearization()
}
now, err := store.serverNow()
if err != nil {
return ClaimResponse{}, false, err
}
request, found, err := findClaimRequest(writeCtx, transaction, command.ClaimRequestID)
if err != nil {
return ClaimResponse{}, false, err
}
if found {
if request.DeviceID != deviceID || request.SessionID != command.SessionID {
return ClaimResponse{}, false, ErrIdempotencyConflict
}
switch request.Outcome {
case "EMPTY":
if err := transaction.Commit(); err != nil {
return ClaimResponse{}, false, err
}
return ClaimResponse{}, false, nil
case "BLOCKED":
if err := transaction.Commit(); err != nil {
return ClaimResponse{}, false, err
}
return ClaimResponse{}, false, ErrRequiresManual
case "CLAIMED":
record, found, err := store.loadClaimByAttempt(writeCtx, transaction, request.AttemptID)
if err != nil || !found {
if err == nil {
err = errors.New("stored claim request has no claim")
}
return ClaimResponse{}, false, err
}
response, err := store.responseFor(record, request.ResponseLeaseExpiresAt)
if err != nil {
return ClaimResponse{}, false, err
}
if err := transaction.Commit(); err != nil {
return ClaimResponse{}, false, err
}
return response, true, nil
default:
return ClaimResponse{}, false, errors.New("stored claim request outcome is invalid")
}
}
existing, found, err := store.loadOpenClaimByDevice(writeCtx, transaction, deviceID)
if err != nil {
return ClaimResponse{}, false, err
}
if found {
current := existing.SessionID == command.SessionID && existing.ClosedAt == "" &&
existing.LeaseExpiresAt.After(now) && existing.AuthorizationExpiresAt.After(now) &&
existing.CurrentAuthorizationExpiresAt.After(now) && existing.AuthorizationStatus == "CLAIMED" &&
existing.authorizationConsistent() && existing.recoverableBusinessState()
if !current {
if err := insertClaimRequest(writeCtx, transaction, command.ClaimRequestID, deviceID, command.SessionID, "BLOCKED", "", "", "manual_recovery_required", now); err != nil {
return ClaimResponse{}, false, err
}
if err := transaction.Commit(); err != nil {
return ClaimResponse{}, false, err
}
return ClaimResponse{}, false, ErrRequiresManual
}
response, err := store.responseFor(existing, existing.LeaseExpiresText)
if err != nil {
return ClaimResponse{}, false, err
}
if err := insertClaimRequest(writeCtx, transaction, command.ClaimRequestID, deviceID, command.SessionID, "CLAIMED", existing.AttemptID, existing.LeaseExpiresText, "", now); err != nil {
return ClaimResponse{}, false, err
}
if err := transaction.Commit(); err != nil {
return ClaimResponse{}, false, err
}
return response, true, nil
}
candidate, found, err := findCandidate(writeCtx, transaction, now)
if err != nil {
return ClaimResponse{}, false, err
}
if !found {
if err := insertClaimRequest(writeCtx, transaction, command.ClaimRequestID, deviceID, command.SessionID, "EMPTY", "", "", "", now); err != nil {
return ClaimResponse{}, false, err
}
if err := transaction.Commit(); err != nil {
return ClaimResponse{}, false, err
}
return ClaimResponse{}, false, nil
}
generation, err := nextGeneration(writeCtx, transaction, candidate.TaskID)
if err != nil {
return ClaimResponse{}, false, err
}
attemptID, err := store.newUUID()
if err != nil {
return ClaimResponse{}, false, err
}
nonce, err := store.randomBytes(sha256.Size)
if err != nil {
return ClaimResponse{}, false, err
}
token := deriveToken(store.secret, deviceID, candidate.TaskID, candidate.AuthorizationID, attemptID, generation, nonce)
storedTokenHash := tokenHash(token)
leaseExpires := now.Add(store.leaseTTL)
if candidate.AuthorizationExpiresAt.Before(leaseExpires) {
leaseExpires = candidate.AuthorizationExpiresAt
}
leaseText := formatTime(leaseExpires)
nowText := formatTime(now)
authorizationUpdate, err := transaction.ExecContext(writeCtx, `UPDATE order_authorizations SET status = 'CLAIMED'
WHERE id = ? AND task_id = ? AND status = 'ACTIVE' AND task_version = ?
AND goods_id = ? AND sku_color = ? AND sku_size = ? AND quantity = ?
AND total_price_cap = ? AND expires_at = ?`,
candidate.AuthorizationID, candidate.TaskID, candidate.TaskVersion, candidate.GoodsID,
candidate.SKUColor, candidate.SKUSize, candidate.Quantity, candidate.TotalPriceCap,
candidate.AuthorizationExpiresText)
if err != nil {
return ClaimResponse{}, false, err
}
if ok, err := exactlyOne(authorizationUpdate); err != nil || !ok {
if err == nil {
err = errors.New("authorization changed during claim")
}
return ClaimResponse{}, false, err
}
taskUpdate, err := transaction.ExecContext(writeCtx, `UPDATE tasks SET status = 'CLAIMED', version = version + 1, updated_at = ?
WHERE id = ? AND status = 'PENDING' AND version = ? AND title = ? AND goods_id = ?
AND sku_color = ? AND sku_size = ? AND quantity = ? AND max_total_price = ?`,
nowText, candidate.TaskID, candidate.TaskVersion, candidate.Title, candidate.GoodsID,
candidate.SKUColor, candidate.SKUSize, candidate.Quantity, candidate.TotalPriceCap)
if err != nil {
return ClaimResponse{}, false, err
}
if ok, err := exactlyOne(taskUpdate); err != nil || !ok {
if err == nil {
err = errors.New("task changed during claim")
}
return ClaimResponse{}, false, err
}
if _, err := transaction.ExecContext(writeCtx, `INSERT INTO purchase_attempts
(id, task_id, authorization_id, claim_generation, status, started_at)
VALUES (?, ?, ?, ?, 'CLAIMED', ?)`, attemptID, candidate.TaskID, candidate.AuthorizationID, generation, nowText); err != nil {
return ClaimResponse{}, false, err
}
if _, err := transaction.ExecContext(writeCtx, `INSERT INTO purchase_attempt_claims
(attempt_id, task_id, authorization_id, claimed_by_device_id, session_id, claim_generation,
task_version, task_title, authorization_task_version, goods_id, sku_color, sku_size, quantity,
total_price_cap, authorization_expires_at, claim_nonce, claim_token_sha256,
lease_expires_at, claimed_at, closed_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL)`,
attemptID, candidate.TaskID, candidate.AuthorizationID, deviceID, command.SessionID, generation,
candidate.TaskVersion+1, candidate.Title, candidate.TaskVersion, candidate.GoodsID,
candidate.SKUColor, candidate.SKUSize, candidate.Quantity, candidate.TotalPriceCap,
candidate.AuthorizationExpiresText, nonce, storedTokenHash, leaseText, nowText); err != nil {
return ClaimResponse{}, false, err
}
if err := insertClaimRequest(writeCtx, transaction, command.ClaimRequestID, deviceID, command.SessionID, "CLAIMED", attemptID, leaseText, "", now); err != nil {
return ClaimResponse{}, false, err
}
response := ClaimResponse{
Task: ClaimedTask{ID: candidate.TaskID, Version: candidate.TaskVersion + 1, Title: candidate.Title,
ProductURL: productURL(candidate.GoodsID), GoodsID: candidate.GoodsID, SKUColor: candidate.SKUColor,
SKUSize: candidate.SKUSize, Quantity: candidate.Quantity, MaxTotalPrice: candidate.TotalPriceCap},
Authorization: ClaimedAuthorization{ID: candidate.AuthorizationID, TaskVersion: candidate.TaskVersion, ExpiresAt: candidate.AuthorizationExpiresText},
Attempt: ClaimedAttempt{ID: attemptID, ClaimToken: hex.EncodeToString(token), ClaimGeneration: generation, LeaseExpiresAt: leaseText},
}
if err := transaction.Commit(); err != nil {
return ClaimResponse{}, false, err
}
return response, true, nil
}
func (store *Store) Renew(ctx context.Context, deviceID string, command RenewCommand) (RenewResponse, error) {
providedToken, tokenOK := decodeToken(command.ClaimToken)
if !deviceauth.ValidDeviceID(deviceID) || !validUUID(command.TaskID) || !validUUID(command.RenewRequestID) ||
!validUUID(command.SessionID) || !validUUID(command.AttemptID) || command.ClaimGeneration <= 0 ||
!tokenOK || !validCanonicalTime(command.ExpectedLeaseExpiresAt) {
return RenewResponse{}, ErrInvalid
}
providedHash := tokenHash(providedToken)
writeCtx, cancel := context.WithTimeout(ctx, writeTimeout)
defer cancel()
select {
case store.writeGate <- struct{}{}:
defer func() { <-store.writeGate }()
case <-writeCtx.Done():
return RenewResponse{}, writeCtx.Err()
}
transaction, err := store.database.BeginTx(writeCtx, nil)
if err != nil {
return RenewResponse{}, err
}
defer transaction.Rollback()
// As in ClaimNext, this is deliberately the first database statement in the transaction.
if store.beforeLinearization != nil {
store.beforeLinearization()
}
active, err := transaction.ExecContext(writeCtx, `UPDATE device_credentials SET status = status
WHERE device_id = ? AND status = 'ACTIVE' AND revoked_at IS NULL`, deviceID)
if err != nil {
return RenewResponse{}, err
}
if ok, err := exactlyOne(active); err != nil {
return RenewResponse{}, err
} else if !ok {
return RenewResponse{}, ErrDeviceInactive
}
if store.afterLinearization != nil {
store.afterLinearization()
}
renewal, found, err := findRenewal(writeCtx, transaction, command.RenewRequestID)
if err != nil {
return RenewResponse{}, err
}
if found {
if renewal.TaskID != command.TaskID || renewal.AttemptID != command.AttemptID || renewal.DeviceID != deviceID ||
renewal.SessionID != command.SessionID || renewal.Generation != command.ClaimGeneration ||
renewal.ExpectedLeaseExpiresAt != command.ExpectedLeaseExpiresAt || !matchingHash(renewal.TokenHash, providedHash) {
return RenewResponse{}, ErrIdempotencyConflict
}
response := RenewResponse{TaskID: renewal.TaskID, AttemptID: renewal.AttemptID, ClaimGeneration: renewal.Generation, LeaseExpiresAt: renewal.LeaseExpiresAt}
if err := transaction.Commit(); err != nil {
return RenewResponse{}, err
}
return response, nil
}
now, err := store.serverNow()
if err != nil {
return RenewResponse{}, err
}
record, found, err := store.loadClaimByAttempt(writeCtx, transaction, command.AttemptID)
if err != nil {
return RenewResponse{}, err
}
if !found || record.TaskID != command.TaskID || record.DeviceID != deviceID || record.SessionID != command.SessionID ||
record.Generation != command.ClaimGeneration || !matchingHash(record.TokenHash, providedHash) {
return RenewResponse{}, ErrNotCurrent
}
stateCurrent := record.ClosedAt == "" && record.LeaseExpiresAt.After(now) && record.AuthorizationExpiresAt.After(now) &&
record.CurrentAuthorizationExpiresAt.After(now) && record.AuthorizationStatus == "CLAIMED" &&
record.authorizationConsistent() && record.recoverableBusinessState()
if !stateCurrent || record.LeaseExpiresText != command.ExpectedLeaseExpiresAt {
return RenewResponse{}, ErrNotCurrent
}
leaseExpires := now.Add(store.leaseTTL)
if record.AuthorizationExpiresAt.Before(leaseExpires) {
leaseExpires = record.AuthorizationExpiresAt
}
leaseText := formatTime(leaseExpires)
updated, err := transaction.ExecContext(writeCtx, `UPDATE purchase_attempt_claims SET lease_expires_at = ?
WHERE attempt_id = ? AND lease_expires_at = ? AND closed_at IS NULL`, leaseText, command.AttemptID, command.ExpectedLeaseExpiresAt)
if err != nil {
return RenewResponse{}, err
}
if ok, err := exactlyOne(updated); err != nil || !ok {
if err == nil {
err = ErrNotCurrent
}
return RenewResponse{}, err
}
if _, err := transaction.ExecContext(writeCtx, `INSERT INTO purchase_attempt_lease_renewals
(renew_request_id, task_id, attempt_id, device_id, session_id, claim_generation,
claim_token_sha256, expected_lease_expires_at, lease_expires_at, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
command.RenewRequestID, command.TaskID, command.AttemptID, deviceID, command.SessionID,
command.ClaimGeneration, record.TokenHash, command.ExpectedLeaseExpiresAt, leaseText, formatTime(now)); err != nil {
return RenewResponse{}, err
}
response := RenewResponse{TaskID: command.TaskID, AttemptID: command.AttemptID, ClaimGeneration: command.ClaimGeneration, LeaseExpiresAt: leaseText}
if err := transaction.Commit(); err != nil {
return RenewResponse{}, err
}
return response, nil
}
type claimRequestRecord struct {
DeviceID, SessionID, Outcome, AttemptID, ResponseLeaseExpiresAt string
}
func findClaimRequest(ctx context.Context, transaction *sql.Tx, requestID string) (claimRequestRecord, bool, error) {
var record claimRequestRecord
var attemptID, responseLease sql.NullString
err := transaction.QueryRowContext(ctx, `SELECT device_id, session_id, outcome, attempt_id, response_lease_expires_at
FROM task_claim_requests WHERE claim_request_id = ?`, requestID).
Scan(&record.DeviceID, &record.SessionID, &record.Outcome, &attemptID, &responseLease)
if errors.Is(err, sql.ErrNoRows) {
return claimRequestRecord{}, false, nil
}
if err != nil {
return claimRequestRecord{}, false, err
}
record.AttemptID, record.ResponseLeaseExpiresAt = attemptID.String, responseLease.String
return record, true, nil
}
func insertClaimRequest(ctx context.Context, transaction *sql.Tx, requestID, deviceID, sessionID, outcome, attemptID, responseLease, errorCode string, now time.Time) error {
var attempt, lease, code any
if attemptID != "" {
attempt = attemptID
}
if responseLease != "" {
lease = responseLease
}
if errorCode != "" {
code = errorCode
}
_, err := transaction.ExecContext(ctx, `INSERT INTO task_claim_requests
(claim_request_id, device_id, session_id, outcome, attempt_id, response_lease_expires_at, error_code, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)`, requestID, deviceID, sessionID, outcome, attempt, lease, code, formatTime(now))
return err
}
type renewalRecord struct {
TaskID, AttemptID, DeviceID, SessionID string
Generation int
TokenHash []byte
ExpectedLeaseExpiresAt, LeaseExpiresAt string
}
func findRenewal(ctx context.Context, transaction *sql.Tx, requestID string) (renewalRecord, bool, error) {
var record renewalRecord
err := transaction.QueryRowContext(ctx, `SELECT task_id, attempt_id, device_id, session_id,
claim_generation, claim_token_sha256, expected_lease_expires_at, lease_expires_at
FROM purchase_attempt_lease_renewals WHERE renew_request_id = ?`, requestID).
Scan(&record.TaskID, &record.AttemptID, &record.DeviceID, &record.SessionID, &record.Generation,
&record.TokenHash, &record.ExpectedLeaseExpiresAt, &record.LeaseExpiresAt)
if errors.Is(err, sql.ErrNoRows) {
return renewalRecord{}, false, nil
}
return record, err == nil, err
}
type claimRecord struct {
AttemptID, TaskID, AuthorizationID, DeviceID, SessionID string
Generation, TaskVersion, CurrentTaskVersion int
TaskTitle string
Nonce, TokenHash []byte
LeaseExpiresText, ClaimedAt, ClosedAt string
LeaseExpiresAt time.Time
AuthorizationTaskVersion int
GoodsID, SKUColor, SKUSize, TotalPriceCap string
Quantity int
AuthorizationExpiresText, AuthorizationStatus string
AuthorizationExpiresAt time.Time
CurrentAuthorizationTaskVersion int
CurrentGoodsID, CurrentSKUColor, CurrentSKUSize string
CurrentQuantity int
CurrentTotalPriceCap, CurrentAuthorizationExpiresText string
CurrentAuthorizationExpiresAt time.Time
AttemptStatus, TaskStatus string
CurrentTaskTitle, CurrentTaskGoodsID string
CurrentTaskSKUColor, CurrentTaskSKUSize string
CurrentTaskQuantity int
CurrentTaskMaxTotalPrice string
CurrentAttemptGeneration int
}
const claimSelect = `SELECT claims.attempt_id, claims.task_id, claims.authorization_id,
claims.claimed_by_device_id, claims.session_id, claims.claim_generation, claims.task_version,
claims.task_title, claims.authorization_task_version, claims.goods_id, claims.sku_color,
claims.sku_size, claims.quantity, claims.total_price_cap, claims.authorization_expires_at,
claims.claim_nonce, claims.claim_token_sha256, claims.lease_expires_at,
claims.claimed_at, claims.closed_at, authorizations.task_version, authorizations.goods_id,
authorizations.sku_color, authorizations.sku_size, authorizations.quantity,
authorizations.total_price_cap, authorizations.expires_at, authorizations.status,
attempts.claim_generation, attempts.status, tasks.status, tasks.version, tasks.title, tasks.goods_id,
tasks.sku_color, tasks.sku_size, tasks.quantity, tasks.max_total_price
FROM purchase_attempt_claims AS claims
JOIN order_authorizations AS authorizations
ON authorizations.task_id = claims.task_id AND authorizations.id = claims.authorization_id
JOIN purchase_attempts AS attempts ON attempts.id = claims.attempt_id
JOIN tasks ON tasks.id = claims.task_id `
func (store *Store) loadOpenClaimByDevice(ctx context.Context, transaction *sql.Tx, deviceID string) (claimRecord, bool, error) {
return store.scanClaim(transaction.QueryRowContext(ctx, claimSelect+`WHERE claims.claimed_by_device_id = ? AND claims.closed_at IS NULL`, deviceID))
}
func (store *Store) loadClaimByAttempt(ctx context.Context, transaction *sql.Tx, attemptID string) (claimRecord, bool, error) {
return store.scanClaim(transaction.QueryRowContext(ctx, claimSelect+`WHERE claims.attempt_id = ?`, attemptID))
}
type rowScanner interface{ Scan(...any) error }
func (store *Store) scanClaim(row rowScanner) (claimRecord, bool, error) {
var record claimRecord
var closed sql.NullString
err := row.Scan(&record.AttemptID, &record.TaskID, &record.AuthorizationID, &record.DeviceID,
&record.SessionID, &record.Generation, &record.TaskVersion, &record.TaskTitle,
&record.AuthorizationTaskVersion, &record.GoodsID, &record.SKUColor, &record.SKUSize,
&record.Quantity, &record.TotalPriceCap, &record.AuthorizationExpiresText,
&record.Nonce, &record.TokenHash, &record.LeaseExpiresText, &record.ClaimedAt, &closed,
&record.CurrentAuthorizationTaskVersion, &record.CurrentGoodsID, &record.CurrentSKUColor,
&record.CurrentSKUSize, &record.CurrentQuantity, &record.CurrentTotalPriceCap,
&record.CurrentAuthorizationExpiresText,
&record.AuthorizationStatus, &record.CurrentAttemptGeneration, &record.AttemptStatus,
&record.TaskStatus, &record.CurrentTaskVersion,
&record.CurrentTaskTitle, &record.CurrentTaskGoodsID, &record.CurrentTaskSKUColor,
&record.CurrentTaskSKUSize, &record.CurrentTaskQuantity, &record.CurrentTaskMaxTotalPrice)
if errors.Is(err, sql.ErrNoRows) {
return claimRecord{}, false, nil
}
if err != nil {
return claimRecord{}, false, err
}
record.ClosedAt = closed.String
if !validUUID(record.AttemptID) || !validUUID(record.TaskID) || !validUUID(record.AuthorizationID) ||
!deviceauth.ValidDeviceID(record.DeviceID) || !validUUID(record.SessionID) || record.Generation <= 0 ||
record.CurrentAttemptGeneration != record.Generation ||
record.TaskVersion <= 0 || record.AuthorizationTaskVersion <= 0 || strings.TrimSpace(record.TaskTitle) == "" ||
!digitsOnly(record.GoodsID) || record.SKUColor == "" || record.SKUSize == "" || record.Quantity <= 0 ||
!canonicalMoney(record.TotalPriceCap) || len(record.Nonce) != sha256.Size || len(record.TokenHash) != sha256.Size {
return claimRecord{}, false, errors.New("stored task claim metadata is invalid")
}
record.LeaseExpiresAt, err = parseCanonicalTime(record.LeaseExpiresText)
if err != nil {
return claimRecord{}, false, errors.New("stored task claim lease is invalid")
}
record.AuthorizationExpiresAt, err = parseCanonicalTime(record.AuthorizationExpiresText)
if err != nil {
return claimRecord{}, false, errors.New("stored authorization expiry is invalid")
}
record.CurrentAuthorizationExpiresAt, err = parseCanonicalTime(record.CurrentAuthorizationExpiresText)
if err != nil {
return claimRecord{}, false, errors.New("current authorization expiry is invalid")
}
derived := deriveToken(store.secret, record.DeviceID, record.TaskID, record.AuthorizationID, record.AttemptID, record.Generation, record.Nonce)
if !matchingHash(tokenHash(derived), record.TokenHash) {
return claimRecord{}, false, errors.New("task claim secret does not match stored claim")
}
return record, true, nil
}
func (record claimRecord) authorizationConsistent() bool {
return record.AuthorizationTaskVersion == record.CurrentAuthorizationTaskVersion &&
record.GoodsID == record.CurrentGoodsID && record.SKUColor == record.CurrentSKUColor &&
record.SKUSize == record.CurrentSKUSize && record.Quantity == record.CurrentQuantity &&
record.TotalPriceCap == record.CurrentTotalPriceCap &&
record.AuthorizationExpiresText == record.CurrentAuthorizationExpiresText &&
record.TaskTitle == record.CurrentTaskTitle && record.GoodsID == record.CurrentTaskGoodsID &&
record.SKUColor == record.CurrentTaskSKUColor && record.SKUSize == record.CurrentTaskSKUSize &&
record.Quantity == record.CurrentTaskQuantity && record.TotalPriceCap == record.CurrentTaskMaxTotalPrice
}
func (record claimRecord) recoverableBusinessState() bool {
if record.TaskStatus == "CLAIMED" && record.AttemptStatus == "CLAIMED" {
return record.CurrentTaskVersion == record.TaskVersion
}
// A later server task may advance this same attempt to ORDERING. A valid lease and identical
// ownership recover that attempt; claim-next still cannot select another task.
return record.TaskStatus == "ORDERING" && record.AttemptStatus == "ORDERING" &&
record.TaskVersion < math.MaxInt && record.CurrentTaskVersion == record.TaskVersion+1
}
func (store *Store) responseFor(record claimRecord, responseLease string) (ClaimResponse, error) {
if !validCanonicalTime(responseLease) {
return ClaimResponse{}, errors.New("stored claim response lease is invalid")
}
token := deriveToken(store.secret, record.DeviceID, record.TaskID, record.AuthorizationID, record.AttemptID, record.Generation, record.Nonce)
return ClaimResponse{
Task: ClaimedTask{ID: record.TaskID, Version: record.TaskVersion, Title: record.TaskTitle,
ProductURL: productURL(record.GoodsID), GoodsID: record.GoodsID, SKUColor: record.SKUColor,
SKUSize: record.SKUSize, Quantity: record.Quantity, MaxTotalPrice: record.TotalPriceCap},
Authorization: ClaimedAuthorization{ID: record.AuthorizationID, TaskVersion: record.AuthorizationTaskVersion, ExpiresAt: record.AuthorizationExpiresText},
Attempt: ClaimedAttempt{ID: record.AttemptID, ClaimToken: hex.EncodeToString(token), ClaimGeneration: record.Generation, LeaseExpiresAt: responseLease},
}, nil
}
type candidate struct {
AuthorizationID, TaskID, Title, GoodsID, SKUColor, SKUSize, TotalPriceCap string
TaskVersion, Quantity int
AuthorizationExpiresText string
AuthorizationExpiresAt time.Time
}
func findCandidate(ctx context.Context, transaction *sql.Tx, now time.Time) (candidate, bool, error) {
rows, err := transaction.QueryContext(ctx, `SELECT authorizations.id, tasks.id, tasks.version,
tasks.title, tasks.goods_id, tasks.sku_color, tasks.sku_size, tasks.quantity,
tasks.max_total_price, authorizations.expires_at
FROM order_authorizations AS authorizations
JOIN tasks ON tasks.id = authorizations.task_id
WHERE authorizations.status = 'ACTIVE' AND tasks.status = 'PENDING'
AND authorizations.task_version = tasks.version
AND authorizations.goods_id = tasks.goods_id
AND authorizations.sku_color = tasks.sku_color
AND authorizations.sku_size = tasks.sku_size
AND authorizations.quantity = tasks.quantity
AND authorizations.total_price_cap = tasks.max_total_price
ORDER BY authorizations.created_at, authorizations.rowid, authorizations.id`)
if err != nil {
return candidate{}, false, err
}
defer rows.Close()
for rows.Next() {
var item candidate
if err := rows.Scan(&item.AuthorizationID, &item.TaskID, &item.TaskVersion, &item.Title,
&item.GoodsID, &item.SKUColor, &item.SKUSize, &item.Quantity, &item.TotalPriceCap,
&item.AuthorizationExpiresText); err != nil {
return candidate{}, false, err
}
item.AuthorizationExpiresAt, err = parseCanonicalTime(item.AuthorizationExpiresText)
if err != nil {
return candidate{}, false, errors.New("stored authorization expiry is invalid")
}
if !validCandidate(item) {
return candidate{}, false, errors.New("stored claim candidate is invalid")
}
if item.AuthorizationExpiresAt.After(now) {
if err := rows.Close(); err != nil {
return candidate{}, false, err
}
return item, true, nil
}
}
if err := rows.Err(); err != nil {
return candidate{}, false, err
}
return candidate{}, false, nil
}
func validCandidate(item candidate) bool {
return validUUID(item.AuthorizationID) && validUUID(item.TaskID) && item.TaskVersion > 0 && item.TaskVersion < math.MaxInt &&
strings.TrimSpace(item.Title) != "" && digitsOnly(item.GoodsID) && item.SKUColor != "" && item.SKUSize != "" &&
item.Quantity > 0 && canonicalMoney(item.TotalPriceCap)
}
func validAttemptStatus(value sql.NullString) bool {
return value.Valid && oneOf(value.String, "CLAIMED", "ORDERING", "FAILED", "FENCED", "ABANDONED")
}
func validAuthorizationStatus(value sql.NullString) bool {
return value.Valid && oneOf(value.String, "ACTIVE", "CLAIMED", "FENCED", "CONSUMED", "EXPIRED", "ABANDONED")
}
func validTaskStatus(value sql.NullString) bool {
return value.Valid && oneOf(value.String, "DRAFT", "PENDING", "CLAIMED", "ORDERING", "NEEDS_MANUAL",
"WAITING_PAYMENT", "RECONCILIATION_REQUIRED", "SUCCEEDED", "FAILED", "CANCELED")
}
func oneOf(value string, allowed ...string) bool {
for _, item := range allowed {
if value == item {
return true
}
}
return false
}
func nextGeneration(ctx context.Context, transaction *sql.Tx, taskID string) (int, error) {
var maximum int64
if err := transaction.QueryRowContext(ctx, `SELECT COALESCE(MAX(claim_generation), 0) FROM purchase_attempts WHERE task_id = ?`, taskID).Scan(&maximum); err != nil {
return 0, err
}
if maximum < 0 || maximum >= int64(math.MaxInt) {
return 0, errors.New("task claim generation is exhausted")
}
return int(maximum) + 1, nil
}
func (store *Store) serverNow() (time.Time, error) {
now := store.now().UTC()
if now.IsZero() {
return time.Time{}, errors.New("task claim clock is invalid")
}
return now, nil
}
func (store *Store) randomBytes(size int) ([]byte, error) {
value := make([]byte, size)
store.randomMu.Lock()
_, err := io.ReadFull(store.random, value)
store.randomMu.Unlock()
if err != nil {
return nil, fmt.Errorf("generate task claim randomness: %w", err)
}
return value, nil
}
func (store *Store) newUUID() (string, error) {
value, err := store.randomBytes(16)
if err != nil {
return "", err
}
value[6] = (value[6] & 0x0f) | 0x40
value[8] = (value[8] & 0x3f) | 0x80
encoded := hex.EncodeToString(value)
return encoded[:8] + "-" + encoded[8:12] + "-" + encoded[12:16] + "-" + encoded[16:20] + "-" + encoded[20:], nil
}
func exactlyOne(result sql.Result) (bool, error) {
rows, err := result.RowsAffected()
return rows == 1, err
}
func formatTime(value time.Time) string { return value.UTC().Format(time.RFC3339Nano) }
func parseCanonicalTime(value string) (time.Time, error) {
if !strings.HasSuffix(value, "Z") || strings.TrimSpace(value) != value {
return time.Time{}, ErrInvalid
}
parsed, err := time.Parse(time.RFC3339Nano, value)
if err != nil || parsed.Location() != time.UTC || formatTime(parsed) != value {
return time.Time{}, ErrInvalid
}
return parsed, nil
}
func validCanonicalTime(value string) bool {
_, err := parseCanonicalTime(value)
return err == nil
}
func validUUID(value string) bool {
if len(value) != 36 {
return false
}
for index, character := range value {
if index == 8 || index == 13 || index == 18 || index == 23 {
if character != '-' {
return false
}
continue
}
if !(character >= '0' && character <= '9' || character >= 'a' && character <= 'f') {
return false
}
}
return value[14] == '4' && (value[19] == '8' || value[19] == '9' || value[19] == 'a' || value[19] == 'b')
}
func digitsOnly(value string) bool {
if value == "" {
return false
}
for _, character := range value {
if character < '0' || character > '9' {
return false
}
}
return true
}
func canonicalMoney(value string) bool {
parts := strings.Split(value, ".")
if len(parts) != 2 || len(parts[0]) == 0 || len(parts[1]) != 2 || (len(parts[0]) > 1 && parts[0][0] == '0') {
return false
}
for _, part := range parts {
if !digitsOnly(part) {
return false
}
}
cents := new(big.Int)
_, ok := cents.SetString(parts[0]+parts[1], 10)
return ok && cents.Sign() > 0
}
func productURL(goodsID string) string {
return "https://mobile.yangkeduo.com/goods.html?goods_id=" + goodsID
}
+895
View File
@@ -0,0 +1,895 @@
package taskclaim
import (
"bytes"
"context"
"crypto/sha256"
"database/sql"
"encoding/hex"
"errors"
"path/filepath"
"reflect"
"runtime"
"strings"
"sync"
"testing"
"time"
"cmbuyer/admin/internal/migrations"
"cmbuyer/admin/internal/storage/sqlite"
)
const (
testDeviceA = "10000000-0000-4000-8000-000000000001"
testDeviceB = "10000000-0000-4000-8000-000000000002"
testSessionA = "20000000-0000-4000-8000-000000000001"
testSessionB = "20000000-0000-4000-8000-000000000002"
testTaskA = "30000000-0000-4000-8000-000000000001"
testTaskB = "30000000-0000-4000-8000-000000000002"
testAuthA = "40000000-0000-4000-8000-000000000001"
testAuthB = "40000000-0000-4000-8000-000000000002"
testClaimRequestA = "50000000-0000-4000-8000-000000000001"
testClaimRequestB = "50000000-0000-4000-8000-000000000002"
testClaimRequestC = "50000000-0000-4000-8000-000000000003"
testRenewRequestA = "60000000-0000-4000-8000-000000000001"
testRenewRequestB = "60000000-0000-4000-8000-000000000002"
)
var testNow = time.Date(2026, 8, 4, 1, 2, 3, 123000000, time.UTC)
func TestClaimReplayEmptyManualAndSecretRecovery(t *testing.T) {
database := openClaimTestDatabase(t)
insertDevice(t, database, testDeviceA, []byte("device-a"))
insertDevice(t, database, testDeviceB, []byte("device-b"))
insertCandidate(t, database, testTaskA, testAuthA, testNow.Add(-time.Minute), testNow.Add(10*time.Minute), true)
secret := bytes.Repeat([]byte{0x11}, 32)
store := mustStore(t, database, secret, 30*time.Second)
store.now = func() time.Time { return testNow }
command := ClaimCommand{SessionID: testSessionA, ClaimRequestID: testClaimRequestA}
claimed, found, err := store.ClaimNext(context.Background(), testDeviceA, command)
if err != nil || !found {
t.Fatalf("ClaimNext = found %v, err %v", found, err)
}
if claimed.Task.ID != testTaskA || claimed.Task.Version != 3 || claimed.Authorization.ID != testAuthA ||
claimed.Authorization.TaskVersion != 2 || claimed.Attempt.ClaimGeneration != 1 ||
len(claimed.Attempt.ClaimToken) != 64 || strings.ToLower(claimed.Attempt.ClaimToken) != claimed.Attempt.ClaimToken {
t.Fatalf("unexpected claim response: %#v", claimed)
}
assertClaimState(t, database, 1, "CLAIMED", "CLAIMED")
assertNoPlaintextTokenColumnOrValue(t, database, claimed.Attempt.ClaimToken)
replayed, found, err := store.ClaimNext(context.Background(), testDeviceA, command)
if err != nil || !found || !reflect.DeepEqual(replayed, claimed) {
t.Fatalf("same request replay = %#v, found %v, err %v", replayed, found, err)
}
restarted := mustStore(t, database, secret, 30*time.Second)
restarted.now = func() time.Time { return testNow.Add(5 * time.Second) }
replayed, found, err = restarted.ClaimNext(context.Background(), testDeviceA, command)
if err != nil || !found || !reflect.DeepEqual(replayed, claimed) {
t.Fatalf("restart replay = %#v, found %v, err %v", replayed, found, err)
}
if _, err := NewStore(database, bytes.Repeat([]byte{0x22}, 32), 30*time.Second); err == nil {
t.Fatal("NewStore accepted a secret that cannot rebuild existing claims")
}
sameSession, found, err := store.ClaimNext(context.Background(), testDeviceA, ClaimCommand{
SessionID: testSessionA, ClaimRequestID: "50000000-0000-4000-8000-000000000005",
})
if err != nil || !found || sameSession.Attempt.ID != claimed.Attempt.ID || sameSession.Attempt.ClaimToken != claimed.Attempt.ClaimToken {
t.Fatalf("same-session recovery = %#v, found %v, err %v", sameSession, found, err)
}
if _, err := database.Exec(`UPDATE order_authorizations SET goods_id='937122477376', sku_color='白色',
sku_size='L', quantity=3, total_price_cap='40.00', expires_at=? WHERE id=?`,
formatTime(testNow.Add(20*time.Minute)), testAuthA); err != nil {
t.Fatalf("mutate authorization source: %v", err)
}
if _, err := database.Exec(`UPDATE tasks SET title='漂移标题', goods_id='937122477376', sku_color='白色',
sku_size='L', quantity=3, max_total_price='40.00' WHERE id=?`, testTaskA); err != nil {
t.Fatalf("mutate task source: %v", err)
}
afterDrift := mustStore(t, database, secret, 30*time.Second)
afterDrift.now = func() time.Time { return testNow.Add(6 * time.Second) }
stable, found, err := afterDrift.ClaimNext(context.Background(), testDeviceA, command)
if err != nil || !found || !reflect.DeepEqual(stable, claimed) {
t.Fatalf("source-drift replay = %#v, found %v, err %v; want original %#v", stable, found, err, claimed)
}
if _, _, err := afterDrift.ClaimNext(context.Background(), testDeviceA, ClaimCommand{
SessionID: testSessionA, ClaimRequestID: "50000000-0000-4000-8000-000000000006",
}); !errors.Is(err, ErrRequiresManual) {
t.Fatalf("new recovery after source drift error = %v", err)
}
manualCommand := ClaimCommand{SessionID: testSessionB, ClaimRequestID: testClaimRequestB}
if _, _, err := store.ClaimNext(context.Background(), testDeviceA, manualCommand); !errors.Is(err, ErrRequiresManual) {
t.Fatalf("different session error = %v, want ErrRequiresManual", err)
}
if _, _, err := store.ClaimNext(context.Background(), testDeviceA, manualCommand); !errors.Is(err, ErrRequiresManual) {
t.Fatalf("manual replay error = %v, want ErrRequiresManual", err)
}
assertClaimState(t, database, 1, "CLAIMED", "CLAIMED")
emptyCommand := ClaimCommand{SessionID: testSessionB, ClaimRequestID: testClaimRequestC}
if _, found, err := store.ClaimNext(context.Background(), testDeviceB, emptyCommand); err != nil || found {
t.Fatalf("empty claim = found %v, err %v", found, err)
}
insertCandidate(t, database, testTaskB, testAuthB, testNow, testNow.Add(10*time.Minute), true)
if _, found, err := store.ClaimNext(context.Background(), testDeviceB, emptyCommand); err != nil || found {
t.Fatalf("persisted EMPTY replay = found %v, err %v", found, err)
}
claimedB, found, err := store.ClaimNext(context.Background(), testDeviceB, ClaimCommand{
SessionID: testSessionB, ClaimRequestID: "50000000-0000-4000-8000-000000000004",
})
if err != nil || !found || claimedB.Task.ID != testTaskB {
t.Fatalf("new request after EMPTY = %#v, found %v, err %v", claimedB, found, err)
}
var distinctNonces int
if err := database.QueryRow("SELECT COUNT(DISTINCT claim_nonce) FROM purchase_attempt_claims").Scan(&distinctNonces); err != nil || distinctNonces != 2 {
t.Fatalf("distinct claim nonces = %d, err %v", distinctNonces, err)
}
}
func TestSameSessionOrderingRecoveryNeverClaimsAnotherTask(t *testing.T) {
database := openClaimTestDatabase(t)
insertDevice(t, database, testDeviceA, []byte("device-a"))
insertCandidate(t, database, testTaskA, testAuthA, testNow, testNow.Add(10*time.Minute), true)
insertCandidate(t, database, testTaskB, testAuthB, testNow.Add(time.Second), testNow.Add(10*time.Minute), true)
store := mustStore(t, database, bytes.Repeat([]byte{0x21}, 32), time.Minute)
store.now = func() time.Time { return testNow }
claimed, found, err := store.ClaimNext(context.Background(), testDeviceA, ClaimCommand{SessionID: testSessionA, ClaimRequestID: testClaimRequestA})
if err != nil || !found || claimed.Task.ID != testTaskA {
t.Fatalf("initial claim = %#v, found %v, err %v", claimed, found, err)
}
if _, err := database.Exec("UPDATE tasks SET status='ORDERING', version=version+1 WHERE id=?", testTaskA); err != nil {
t.Fatal(err)
}
if _, err := database.Exec("UPDATE purchase_attempts SET status='ORDERING' WHERE id=?", claimed.Attempt.ID); err != nil {
t.Fatal(err)
}
recovered, found, err := store.ClaimNext(context.Background(), testDeviceA, ClaimCommand{SessionID: testSessionA, ClaimRequestID: testClaimRequestB})
if err != nil || !found || recovered.Attempt.ID != claimed.Attempt.ID || recovered.Task.ID != testTaskA {
t.Fatalf("ORDERING recovery = %#v, found %v, err %v", recovered, found, err)
}
var attempts int
var taskBStatus string
if err := database.QueryRow("SELECT COUNT(*) FROM purchase_attempts").Scan(&attempts); err != nil {
t.Fatal(err)
}
if err := database.QueryRow("SELECT status FROM tasks WHERE id=?", testTaskB).Scan(&taskBStatus); err != nil {
t.Fatal(err)
}
if attempts != 1 || taskBStatus != "PENDING" {
t.Fatalf("ORDERING recovery attempts/taskB = %d/%s", attempts, taskBStatus)
}
if _, err := database.Exec("UPDATE tasks SET version=version+1 WHERE id=?", testTaskA); err != nil {
t.Fatal(err)
}
if _, _, err := store.ClaimNext(context.Background(), testDeviceA, ClaimCommand{SessionID: testSessionA, ClaimRequestID: testClaimRequestC}); !errors.Is(err, ErrRequiresManual) {
t.Fatalf("ORDERING recovery with drifted version error = %v", err)
}
if err := database.QueryRow("SELECT COUNT(*) FROM purchase_attempts").Scan(&attempts); err != nil || attempts != 1 {
t.Fatalf("attempts after drifted ORDERING recovery = %d, err %v", attempts, err)
}
}
func TestClaimRollsBackEveryBusinessMutationOnLateFailure(t *testing.T) {
database := openClaimTestDatabase(t)
insertDevice(t, database, testDeviceA, []byte("device-a"))
insertCandidate(t, database, testTaskA, testAuthA, testNow, testNow.Add(time.Minute), true)
if _, err := database.Exec(`CREATE TRIGGER fail_claim_insert BEFORE INSERT ON purchase_attempt_claims
BEGIN SELECT RAISE(ABORT, 'injected claim failure'); END`); err != nil {
t.Fatal(err)
}
store := mustStore(t, database, bytes.Repeat([]byte{0x31}, 32), 30*time.Second)
store.now = func() time.Time { return testNow }
if _, _, err := store.ClaimNext(context.Background(), testDeviceA, ClaimCommand{SessionID: testSessionA, ClaimRequestID: testClaimRequestA}); err == nil {
t.Fatal("ClaimNext succeeded despite injected late failure")
}
assertClaimState(t, database, 0, "PENDING", "ACTIVE")
for _, table := range []string{"purchase_attempts", "task_claim_requests"} {
var count int
if err := database.QueryRow("SELECT COUNT(*) FROM " + table).Scan(&count); err != nil || count != 0 {
t.Fatalf("%s rows after rollback = %d, err %v", table, count, err)
}
}
}
func TestClaimEligibilityStableOrderAndConcurrentUniqueness(t *testing.T) {
database := openClaimTestDatabase(t)
insertDevice(t, database, testDeviceA, []byte("device-a"))
insertDevice(t, database, testDeviceB, []byte("device-b"))
// The oldest row has a mismatched snapshot and is ineligible; the next oldest valid row wins.
insertCandidate(t, database, testTaskA, testAuthA, testNow.Add(-2*time.Minute), testNow.Add(10*time.Minute), false)
insertCandidate(t, database, testTaskB, testAuthB, testNow.Add(-time.Minute), testNow.Add(10*time.Minute), true)
store := mustStore(t, database, bytes.Repeat([]byte{0x33}, 32), time.Minute)
store.now = func() time.Time { return testNow }
type result struct {
response ClaimResponse
found bool
err error
}
commands := []struct{ device, session, request string }{
{testDeviceA, testSessionA, testClaimRequestA},
{testDeviceB, testSessionB, testClaimRequestB},
}
results := make(chan result, 2)
var wait sync.WaitGroup
for _, command := range commands {
command := command
wait.Add(1)
go func() {
defer wait.Done()
response, found, err := store.ClaimNext(context.Background(), command.device, ClaimCommand{SessionID: command.session, ClaimRequestID: command.request})
results <- result{response, found, err}
}()
}
wait.Wait()
close(results)
foundCount := 0
for result := range results {
if result.err != nil {
t.Fatalf("concurrent ClaimNext error: %v", result.err)
}
if result.found {
foundCount++
if result.response.Task.ID != testTaskB {
t.Fatalf("claimed task = %s, want stable eligible task B", result.response.Task.ID)
}
}
}
if foundCount != 1 {
t.Fatalf("successful claims = %d, want 1", foundCount)
}
var claimCount int
if err := database.QueryRow("SELECT COUNT(*) FROM purchase_attempt_claims").Scan(&claimCount); err != nil || claimCount != 1 {
t.Fatalf("claim count = %d, err %v", claimCount, err)
}
var taskStatus, authorizationStatus string
if err := database.QueryRow(`SELECT tasks.status, order_authorizations.status FROM tasks
JOIN order_authorizations ON order_authorizations.task_id=tasks.id WHERE tasks.id=?`, testTaskB).
Scan(&taskStatus, &authorizationStatus); err != nil || taskStatus != "CLAIMED" || authorizationStatus != "CLAIMED" {
t.Fatalf("claimed B states = %s/%s, err %v", taskStatus, authorizationStatus, err)
}
}
func TestClaimConcurrencyAcrossDistinctDatabasesAndStores(t *testing.T) {
path := filepath.ToSlash(filepath.Join(t.TempDir(), "shared-claim.db"))
source := "file:" + path + "?_busy_timeout=5000&_journal_mode=WAL"
databaseA, err := sqlite.Open(source)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = databaseA.Close() })
if err := migrations.Up(context.Background(), databaseA, claimMigrationDirectory(t)); err != nil {
t.Fatal(err)
}
databaseB, err := sqlite.Open(source)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = databaseB.Close() })
databaseA.SetMaxOpenConns(1)
databaseB.SetMaxOpenConns(1)
insertDevice(t, databaseA, testDeviceA, []byte("device-a"))
insertDevice(t, databaseA, testDeviceB, []byte("device-b"))
insertCandidate(t, databaseA, testTaskA, testAuthA, testNow, testNow.Add(10*time.Minute), true)
secret := bytes.Repeat([]byte{0x39}, 32)
storeA := mustStore(t, databaseA, secret, time.Minute)
storeB := mustStore(t, databaseB, secret, time.Minute)
storeA.now = func() time.Time { return testNow }
storeB.now = func() time.Time { return testNow }
firstLinearized := make(chan struct{})
releaseFirst := make(chan struct{})
secondAtFirstWrite := make(chan struct{})
var releaseOnce sync.Once
release := func() { releaseOnce.Do(func() { close(releaseFirst) }) }
t.Cleanup(release)
storeA.afterLinearization = func() {
close(firstLinearized)
<-releaseFirst
}
storeB.beforeLinearization = func() {
// Reaching this hook means B has begun its own transaction and its very next
// database operation is the first-write UPDATE currently held by A.
close(secondAtFirstWrite)
}
type result struct {
found bool
err error
}
firstResult := make(chan result, 1)
secondResult := make(chan result, 1)
go func() {
_, found, err := storeA.ClaimNext(context.Background(), testDeviceA, ClaimCommand{
SessionID: testSessionA, ClaimRequestID: testClaimRequestA,
})
firstResult <- result{found: found, err: err}
}()
select {
case <-firstLinearized:
case result := <-firstResult:
t.Fatalf("first ClaimNext returned before holding SQLite write position: found %v, err %v", result.found, result.err)
case <-time.After(time.Second):
t.Fatal("first ClaimNext did not reach SQLite write position")
}
go func() {
_, found, err := storeB.ClaimNext(context.Background(), testDeviceB, ClaimCommand{
SessionID: testSessionB, ClaimRequestID: testClaimRequestB,
})
secondResult <- result{found: found, err: err}
}()
select {
case <-secondAtFirstWrite:
// A still owns the SQLite write position here. B cannot have observed or
// changed claim state, so releasing A below creates deterministic contention.
case result := <-secondResult:
release()
<-firstResult
t.Fatalf("second ClaimNext returned before reaching the contended first write: found %v, err %v", result.found, result.err)
case <-time.After(time.Second):
release()
<-firstResult
t.Fatal("second ClaimNext did not reach the contended SQLite first write")
}
select {
case result := <-secondResult:
release()
<-firstResult
t.Fatalf("second ClaimNext completed while first transaction held SQLite write position: found %v, err %v", result.found, result.err)
default:
}
release()
first := <-firstResult
second := <-secondResult
if first.err != nil || !first.found {
t.Fatalf("first cross-database ClaimNext = found %v, err %v", first.found, first.err)
}
if second.err != nil || second.found {
t.Fatalf("second cross-database ClaimNext = found %v, err %v", second.found, second.err)
}
var attempts, claims, requestsCount, claimedRequests, emptyRequests int
queries := []struct {
query string
value *int
}{
{"SELECT COUNT(*) FROM purchase_attempts", &attempts},
{"SELECT COUNT(*) FROM purchase_attempt_claims", &claims},
{"SELECT COUNT(*) FROM task_claim_requests", &requestsCount},
{"SELECT COUNT(*) FROM task_claim_requests WHERE outcome='CLAIMED'", &claimedRequests},
{"SELECT COUNT(*) FROM task_claim_requests WHERE outcome='EMPTY'", &emptyRequests},
}
for _, query := range queries {
if err := databaseA.QueryRow(query.query).Scan(query.value); err != nil {
t.Fatal(err)
}
}
if attempts != 1 || claims != 1 || requestsCount != 2 || claimedRequests != 1 || emptyRequests != 1 {
t.Fatalf("cross-database attempts/claims/requests/claimed/empty = %d/%d/%d/%d/%d",
attempts, claims, requestsCount, claimedRequests, emptyRequests)
}
}
func TestRenewCASReplayCapAndNoResurrection(t *testing.T) {
database := openClaimTestDatabase(t)
insertDevice(t, database, testDeviceA, []byte("device-a"))
insertCandidate(t, database, testTaskA, testAuthA, testNow, testNow.Add(40*time.Second), true)
store := mustStore(t, database, bytes.Repeat([]byte{0x44}, 32), 30*time.Second)
current := testNow
store.now = func() time.Time { return current }
claim, found, err := store.ClaimNext(context.Background(), testDeviceA, ClaimCommand{SessionID: testSessionA, ClaimRequestID: testClaimRequestA})
if err != nil || !found {
t.Fatalf("ClaimNext = found %v, err %v", found, err)
}
current = testNow.Add(20 * time.Second)
command := RenewCommand{TaskID: testTaskA, RenewRequestID: testRenewRequestA, SessionID: testSessionA,
AttemptID: claim.Attempt.ID, ClaimGeneration: claim.Attempt.ClaimGeneration,
ClaimToken: claim.Attempt.ClaimToken, ExpectedLeaseExpiresAt: claim.Attempt.LeaseExpiresAt}
renewed, err := store.Renew(context.Background(), testDeviceA, command)
if err != nil {
t.Fatalf("Renew: %v", err)
}
wantCap := formatTime(testNow.Add(40 * time.Second))
if renewed.LeaseExpiresAt != wantCap {
t.Fatalf("renewed lease = %s, want authorization cap %s", renewed.LeaseExpiresAt, wantCap)
}
current = testNow.Add(25 * time.Second)
replay, err := store.Renew(context.Background(), testDeviceA, command)
if err != nil || !reflect.DeepEqual(replay, renewed) {
t.Fatalf("renew replay = %#v, err %v", replay, err)
}
changed := command
changed.ExpectedLeaseExpiresAt = renewed.LeaseExpiresAt
if _, err := store.Renew(context.Background(), testDeviceA, changed); !errors.Is(err, ErrIdempotencyConflict) {
t.Fatalf("same key different payload error = %v", err)
}
stale := command
stale.RenewRequestID = "60000000-0000-4000-8000-000000000004"
if _, err := store.Renew(context.Background(), testDeviceA, stale); !errors.Is(err, ErrNotCurrent) {
t.Fatalf("out-of-order expected lease error = %v", err)
}
wrongToken := command
wrongToken.RenewRequestID = testRenewRequestB
wrongToken.ClaimToken = strings.Repeat("0", 64)
if _, err := store.Renew(context.Background(), testDeviceA, wrongToken); !errors.Is(err, ErrNotCurrent) {
t.Fatalf("wrong token error = %v", err)
}
current = testNow.Add(40 * time.Second) // equality is expired; no grace and no resurrection.
expired := command
expired.RenewRequestID = "60000000-0000-4000-8000-000000000003"
expired.ExpectedLeaseExpiresAt = renewed.LeaseExpiresAt
if _, err := store.Renew(context.Background(), testDeviceA, expired); !errors.Is(err, ErrNotCurrent) {
t.Fatalf("expired renewal error = %v", err)
}
var lease, taskStatus, attemptStatus, authorizationStatus string
if err := database.QueryRow(`SELECT claims.lease_expires_at, tasks.status, attempts.status, authorizations.status
FROM purchase_attempt_claims claims JOIN tasks ON tasks.id=claims.task_id
JOIN purchase_attempts attempts ON attempts.id=claims.attempt_id
JOIN order_authorizations authorizations ON authorizations.id=claims.authorization_id`).
Scan(&lease, &taskStatus, &attemptStatus, &authorizationStatus); err != nil {
t.Fatal(err)
}
if lease != wantCap || taskStatus != "CLAIMED" || attemptStatus != "CLAIMED" || authorizationStatus != "CLAIMED" {
t.Fatalf("renew changed business state: lease=%s task=%s attempt=%s auth=%s", lease, taskStatus, attemptStatus, authorizationStatus)
}
if _, err := database.Exec(`UPDATE device_credentials SET status='REVOKED', revoked_at=? WHERE device_id=?`, formatTime(current), testDeviceA); err != nil {
t.Fatal(err)
}
revoked := expired
revoked.RenewRequestID = "60000000-0000-4000-8000-000000000005"
if _, err := store.Renew(context.Background(), testDeviceA, revoked); !errors.Is(err, ErrDeviceInactive) {
t.Fatalf("renew after revocation error = %v", err)
}
}
func TestRenewRequiresPairedBusinessStateAndExactTaskVersion(t *testing.T) {
tests := []struct {
name string
mutate func(*testing.T, *sql.DB, ClaimResponse)
wantError bool
}{
{"claimed exact version", func(*testing.T, *sql.DB, ClaimResponse) {}, false},
{"ordering exact next version", func(t *testing.T, database *sql.DB, claim ClaimResponse) {
if _, err := database.Exec("UPDATE tasks SET status='ORDERING',version=version+1 WHERE id=?", testTaskA); err != nil {
t.Fatal(err)
}
if _, err := database.Exec("UPDATE purchase_attempts SET status='ORDERING' WHERE id=?", claim.Attempt.ID); err != nil {
t.Fatal(err)
}
}, false},
{"claimed version drift", func(t *testing.T, database *sql.DB, _ ClaimResponse) {
if _, err := database.Exec("UPDATE tasks SET version=version+1 WHERE id=?", testTaskA); err != nil {
t.Fatal(err)
}
}, true},
{"ordering version drift", func(t *testing.T, database *sql.DB, claim ClaimResponse) {
if _, err := database.Exec("UPDATE tasks SET status='ORDERING',version=version+2 WHERE id=?", testTaskA); err != nil {
t.Fatal(err)
}
if _, err := database.Exec("UPDATE purchase_attempts SET status='ORDERING' WHERE id=?", claim.Attempt.ID); err != nil {
t.Fatal(err)
}
}, true},
{"task ordering attempt claimed", func(t *testing.T, database *sql.DB, _ ClaimResponse) {
if _, err := database.Exec("UPDATE tasks SET status='ORDERING',version=version+1 WHERE id=?", testTaskA); err != nil {
t.Fatal(err)
}
}, true},
{"task claimed attempt ordering", func(t *testing.T, database *sql.DB, claim ClaimResponse) {
if _, err := database.Exec("UPDATE purchase_attempts SET status='ORDERING' WHERE id=?", claim.Attempt.ID); err != nil {
t.Fatal(err)
}
}, true},
}
for index, test := range tests {
t.Run(test.name, func(t *testing.T) {
database, store, claim := claimedRenewFixture(t, byte(0x50+index))
test.mutate(t, database, claim)
_, err := store.Renew(context.Background(), testDeviceA, renewCommandFor(claim, testRenewRequestA))
if test.wantError {
if !errors.Is(err, ErrNotCurrent) {
t.Fatalf("Renew error = %v, want ErrNotCurrent", err)
}
var lease string
var renewals int
if scanErr := database.QueryRow("SELECT lease_expires_at FROM purchase_attempt_claims WHERE attempt_id=?", claim.Attempt.ID).Scan(&lease); scanErr != nil {
t.Fatal(scanErr)
}
if scanErr := database.QueryRow("SELECT COUNT(*) FROM purchase_attempt_lease_renewals").Scan(&renewals); scanErr != nil {
t.Fatal(scanErr)
}
if lease != claim.Attempt.LeaseExpiresAt || renewals != 0 {
t.Fatalf("rejected renew changed lease/rows = %s/%d", lease, renewals)
}
return
}
if err != nil {
t.Fatalf("Renew valid state: %v", err)
}
})
}
}
func TestConcurrentRenewCASUsesSQLiteNotOneStoreGate(t *testing.T) {
database := openClaimTestDatabase(t)
insertDevice(t, database, testDeviceA, []byte("device-a"))
insertCandidate(t, database, testTaskA, testAuthA, testNow, testNow.Add(10*time.Minute), true)
secret := bytes.Repeat([]byte{0x48}, 32)
storeA := mustStore(t, database, secret, time.Minute)
storeA.now = func() time.Time { return testNow }
claim, found, err := storeA.ClaimNext(context.Background(), testDeviceA, ClaimCommand{SessionID: testSessionA, ClaimRequestID: testClaimRequestA})
if err != nil || !found {
t.Fatalf("ClaimNext = found %v, err %v", found, err)
}
storeB := mustStore(t, database, secret, time.Minute)
renewNow := testNow.Add(10 * time.Second)
storeA.now = func() time.Time { return renewNow }
storeB.now = func() time.Time { return renewNow }
base := RenewCommand{TaskID: testTaskA, SessionID: testSessionA, AttemptID: claim.Attempt.ID,
ClaimGeneration: claim.Attempt.ClaimGeneration, ClaimToken: claim.Attempt.ClaimToken,
ExpectedLeaseExpiresAt: claim.Attempt.LeaseExpiresAt}
commands := []RenewCommand{base, base}
commands[0].RenewRequestID = testRenewRequestA
commands[1].RenewRequestID = testRenewRequestB
type result struct{ err error }
results := make(chan result, 2)
var wait sync.WaitGroup
for index, claimStore := range []*Store{storeA, storeB} {
index, claimStore := index, claimStore
wait.Add(1)
go func() {
defer wait.Done()
_, err := claimStore.Renew(context.Background(), testDeviceA, commands[index])
results <- result{err: err}
}()
}
wait.Wait()
close(results)
successes, stale := 0, 0
for result := range results {
switch {
case result.err == nil:
successes++
case errors.Is(result.err, ErrNotCurrent):
stale++
default:
t.Fatalf("concurrent Renew error = %v", result.err)
}
}
var renewalCount int
if err := database.QueryRow("SELECT COUNT(*) FROM purchase_attempt_lease_renewals").Scan(&renewalCount); err != nil {
t.Fatal(err)
}
if successes != 1 || stale != 1 || renewalCount != 1 {
t.Fatalf("concurrent renew success/stale/rows = %d/%d/%d", successes, stale, renewalCount)
}
}
func TestRenewRevocationLinearizationBothOrders(t *testing.T) {
t.Run("revocation first", func(t *testing.T) {
database, store, claim := claimedRenewFixture(t, 0x49)
if _, err := database.Exec(`UPDATE device_credentials SET status='REVOKED', revoked_at=? WHERE device_id=?`, formatTime(testNow.Add(time.Second)), testDeviceA); err != nil {
t.Fatal(err)
}
command := renewCommandFor(claim, testRenewRequestA)
if _, err := store.Renew(context.Background(), testDeviceA, command); !errors.Is(err, ErrDeviceInactive) {
t.Fatalf("Renew after revocation error = %v", err)
}
})
t.Run("renew write position first", func(t *testing.T) {
database, store, claim := claimedRenewFixture(t, 0x4a)
linearized := make(chan struct{})
release := make(chan struct{})
store.afterLinearization = func() { close(linearized); <-release }
renewResult := make(chan error, 1)
go func() {
_, err := store.Renew(context.Background(), testDeviceA, renewCommandFor(claim, testRenewRequestA))
renewResult <- err
}()
<-linearized
revocationStarted := make(chan struct{})
revocationResult := make(chan error, 1)
go func() {
close(revocationStarted)
_, err := database.Exec(`UPDATE device_credentials SET status='REVOKED', revoked_at=?
WHERE device_id=? AND status='ACTIVE'`, formatTime(testNow.Add(2*time.Second)), testDeviceA)
revocationResult <- err
}()
<-revocationStarted
close(release)
if err := <-renewResult; err != nil {
t.Fatalf("renew holding first write position: %v", err)
}
if err := <-revocationResult; err != nil {
t.Fatalf("revocation after renew: %v", err)
}
var renewals int
if err := database.QueryRow("SELECT COUNT(*) FROM purchase_attempt_lease_renewals").Scan(&renewals); err != nil || renewals != 1 {
t.Fatalf("renewal rows = %d, err %v", renewals, err)
}
})
}
func TestStartupValidatesClosedClaimStorageTypesAndStatuses(t *testing.T) {
for _, test := range []struct {
name string
mutate func(*testing.T, *sql.DB)
}{
{"text nonce", func(t *testing.T, database *sql.DB) {
if _, err := database.Exec(`UPDATE purchase_attempt_claims
SET claim_nonce=CAST('12345678901234567890123456789012' AS TEXT)`); err != nil {
t.Fatal(err)
}
}},
{"invalid attempt status", func(t *testing.T, database *sql.DB) {
if _, err := database.Exec("UPDATE purchase_attempts SET status='CORRUPT'"); err != nil {
t.Fatal(err)
}
}},
} {
t.Run(test.name, func(t *testing.T) {
database := openClaimTestDatabase(t)
database.SetMaxOpenConns(1)
insertDevice(t, database, testDeviceA, []byte("device-a"))
insertCandidate(t, database, testTaskA, testAuthA, testNow, testNow.Add(10*time.Minute), true)
secret := bytes.Repeat([]byte{0x4b}, 32)
store := mustStore(t, database, secret, time.Minute)
store.now = func() time.Time { return testNow }
claim, found, err := store.ClaimNext(context.Background(), testDeviceA, ClaimCommand{SessionID: testSessionA, ClaimRequestID: testClaimRequestA})
if err != nil || !found {
t.Fatalf("ClaimNext = found %v, err %v", found, err)
}
if _, err := database.Exec("UPDATE purchase_attempt_claims SET closed_at=? WHERE attempt_id=?", formatTime(testNow.Add(2*time.Minute)), claim.Attempt.ID); err != nil {
t.Fatal(err)
}
if _, err := NewStore(database, secret, time.Minute); err != nil {
t.Fatalf("valid closed claim rejected: %v", err)
}
if _, err := database.Exec("PRAGMA ignore_check_constraints=ON"); err != nil {
t.Fatal(err)
}
test.mutate(t, database)
if _, err := NewStore(database, secret, time.Minute); err == nil {
t.Fatal("NewStore accepted corrupted closed claim storage")
}
})
}
}
func TestStartupRejectsClaimAttemptGenerationCorruption(t *testing.T) {
database := openClaimTestDatabase(t)
database.SetMaxOpenConns(1)
insertDevice(t, database, testDeviceA, []byte("device-a"))
insertCandidate(t, database, testTaskA, testAuthA, testNow, testNow.Add(10*time.Minute), true)
secret := bytes.Repeat([]byte{0x4c}, 32)
store := mustStore(t, database, secret, time.Minute)
store.now = func() time.Time { return testNow }
claim, found, err := store.ClaimNext(context.Background(), testDeviceA, ClaimCommand{SessionID: testSessionA, ClaimRequestID: testClaimRequestA})
if err != nil || !found {
t.Fatalf("ClaimNext = found %v, err %v", found, err)
}
var nonce []byte
if err := database.QueryRow("SELECT claim_nonce FROM purchase_attempt_claims WHERE attempt_id=?", claim.Attempt.ID).Scan(&nonce); err != nil {
t.Fatal(err)
}
corruptGeneration := claim.Attempt.ClaimGeneration + 1
corruptToken := deriveToken(secret, testDeviceA, testTaskA, testAuthA, claim.Attempt.ID, corruptGeneration, nonce)
if _, err := database.Exec("PRAGMA foreign_keys=OFF"); err != nil {
t.Fatal(err)
}
if _, err := database.Exec(`UPDATE purchase_attempt_claims SET claim_generation=?,claim_token_sha256=? WHERE attempt_id=?`,
corruptGeneration, tokenHash(corruptToken), claim.Attempt.ID); err != nil {
t.Fatalf("inject generation corruption: %v", err)
}
if _, err := NewStore(database, secret, time.Minute); err == nil {
t.Fatal("NewStore accepted claim generation different from its attempt")
}
}
func TestRevocationLinearizesBeforeOrAfterClaim(t *testing.T) {
t.Run("revocation first", func(t *testing.T) {
database := openClaimTestDatabase(t)
insertDevice(t, database, testDeviceA, []byte("device-a"))
insertCandidate(t, database, testTaskA, testAuthA, testNow, testNow.Add(time.Minute), true)
if _, err := database.Exec(`UPDATE device_credentials SET status='REVOKED', revoked_at=? WHERE device_id=?`, formatTime(testNow), testDeviceA); err != nil {
t.Fatal(err)
}
store := mustStore(t, database, bytes.Repeat([]byte{0x55}, 32), 30*time.Second)
if _, _, err := store.ClaimNext(context.Background(), testDeviceA, ClaimCommand{SessionID: testSessionA, ClaimRequestID: testClaimRequestA}); !errors.Is(err, ErrDeviceInactive) {
t.Fatalf("ClaimNext error = %v, want inactive", err)
}
assertClaimState(t, database, 0, "PENDING", "ACTIVE")
})
t.Run("claim write position first", func(t *testing.T) {
database := openClaimTestDatabase(t)
insertDevice(t, database, testDeviceA, []byte("device-a"))
insertCandidate(t, database, testTaskA, testAuthA, testNow, testNow.Add(time.Minute), true)
store := mustStore(t, database, bytes.Repeat([]byte{0x66}, 32), 30*time.Second)
store.now = func() time.Time { return testNow }
linearized := make(chan struct{})
release := make(chan struct{})
store.afterLinearization = func() { close(linearized); <-release }
claimResult := make(chan error, 1)
go func() {
_, found, err := store.ClaimNext(context.Background(), testDeviceA, ClaimCommand{SessionID: testSessionA, ClaimRequestID: testClaimRequestA})
if err == nil && !found {
err = errors.New("claim unexpectedly empty")
}
claimResult <- err
}()
<-linearized
revocationStarted := make(chan struct{})
revocationResult := make(chan error, 1)
go func() {
close(revocationStarted)
_, err := database.Exec(`UPDATE device_credentials SET status='REVOKED', revoked_at=? WHERE device_id=? AND status='ACTIVE'`, formatTime(testNow.Add(time.Second)), testDeviceA)
revocationResult <- err
}()
<-revocationStarted
close(release)
if err := <-claimResult; err != nil {
t.Fatalf("claim holding first write position: %v", err)
}
if err := <-revocationResult; err != nil {
t.Fatalf("revocation after claim: %v", err)
}
assertClaimState(t, database, 1, "CLAIMED", "CLAIMED")
})
}
func TestTokenDomainSeparationAndDeviceSecretIsolation(t *testing.T) {
secret := bytes.Repeat([]byte{0x77}, 32)
nonce := bytes.Repeat([]byte{0x88}, 32)
base := deriveToken(secret, testDeviceA, testTaskA, testAuthA, "70000000-0000-4000-8000-000000000001", 1, nonce)
variants := [][]byte{
deriveToken(secret, testDeviceB, testTaskA, testAuthA, "70000000-0000-4000-8000-000000000001", 1, nonce),
deriveToken(secret, testDeviceA, testTaskB, testAuthA, "70000000-0000-4000-8000-000000000001", 1, nonce),
deriveToken(secret, testDeviceA, testTaskA, testAuthB, "70000000-0000-4000-8000-000000000001", 1, nonce),
deriveToken(secret, testDeviceA, testTaskA, testAuthA, "70000000-0000-4000-8000-000000000002", 1, nonce),
deriveToken(secret, testDeviceA, testTaskA, testAuthA, "70000000-0000-4000-8000-000000000001", 2, nonce),
}
for index, variant := range variants {
if matchingHash(base, variant) {
t.Fatalf("token variant %d was not domain separated", index)
}
}
database := openClaimTestDatabase(t)
insertDevice(t, database, testDeviceA, secret)
if _, err := NewStore(database, secret, time.Minute); err == nil {
t.Fatal("NewStore accepted a key equal to a device token")
}
}
func openClaimTestDatabase(t *testing.T) *sql.DB {
t.Helper()
path := filepath.ToSlash(filepath.Join(t.TempDir(), "claim.db"))
database, err := sqlite.Open("file:" + path + "?_busy_timeout=5000&_journal_mode=WAL")
if err != nil {
t.Fatalf("open database: %v", err)
}
t.Cleanup(func() { _ = database.Close() })
if err := migrations.Up(context.Background(), database, claimMigrationDirectory(t)); err != nil {
t.Fatalf("migrate database: %v", err)
}
return database
}
func claimMigrationDirectory(t *testing.T) string {
t.Helper()
_, file, _, ok := runtime.Caller(0)
if !ok {
t.Fatal("locate test file")
}
return filepath.Join(filepath.Dir(file), "..", "..", "migrations")
}
func mustStore(t *testing.T, database *sql.DB, secret []byte, ttl time.Duration) *Store {
t.Helper()
store, err := NewStore(database, secret, ttl)
if err != nil {
t.Fatalf("NewStore: %v", err)
}
return store
}
func claimedRenewFixture(t *testing.T, secretByte byte) (*sql.DB, *Store, ClaimResponse) {
t.Helper()
database := openClaimTestDatabase(t)
insertDevice(t, database, testDeviceA, []byte("device-a"))
insertCandidate(t, database, testTaskA, testAuthA, testNow, testNow.Add(10*time.Minute), true)
store := mustStore(t, database, bytes.Repeat([]byte{secretByte}, 32), time.Minute)
store.now = func() time.Time { return testNow }
claim, found, err := store.ClaimNext(context.Background(), testDeviceA, ClaimCommand{SessionID: testSessionA, ClaimRequestID: testClaimRequestA})
if err != nil || !found {
t.Fatalf("ClaimNext = found %v, err %v", found, err)
}
store.now = func() time.Time { return testNow.Add(10 * time.Second) }
return database, store, claim
}
func renewCommandFor(claim ClaimResponse, requestID string) RenewCommand {
return RenewCommand{TaskID: testTaskA, RenewRequestID: requestID, SessionID: testSessionA,
AttemptID: claim.Attempt.ID, ClaimGeneration: claim.Attempt.ClaimGeneration,
ClaimToken: claim.Attempt.ClaimToken, ExpectedLeaseExpiresAt: claim.Attempt.LeaseExpiresAt}
}
func insertDevice(t *testing.T, database *sql.DB, deviceID string, token []byte) {
t.Helper()
digest := sha256.Sum256(token)
if _, err := database.Exec(`INSERT INTO device_credentials
(device_id,display_name,token_sha256,status,created_at,revoked_at)
VALUES (?, ?, ?, 'ACTIVE', ?, NULL)`, deviceID, "test device", digest[:], formatTime(testNow.Add(-time.Hour))); err != nil {
t.Fatalf("insert device: %v", err)
}
}
func insertCandidate(t *testing.T, database *sql.DB, taskID, authorizationID string, createdAt, expiresAt time.Time, snapshotMatches bool) {
t.Helper()
if _, err := database.Exec(`INSERT INTO tasks
(id,source,title,goods_id,sku_color,sku_size,quantity,max_total_price,status,version,created_at,updated_at)
VALUES (?, 'MANUAL', '测试商品', '937122477375', '黑色', 'M', 2, '30.00', 'PENDING', 2, ?, ?)`,
taskID, formatTime(createdAt), formatTime(createdAt)); err != nil {
t.Fatalf("insert task: %v", err)
}
color := "黑色"
if !snapshotMatches {
color = "白色"
}
if _, err := database.Exec(`INSERT INTO order_authorizations
(id,task_id,task_version,start_key,goods_id,sku_color,sku_size,quantity,total_price_cap,status,created_by,created_at,expires_at)
VALUES (?, ?, 2, ?, '937122477375', ?, 'M', 2, '30.00', 'ACTIVE', 'admin', ?, ?)`,
authorizationID, taskID, authorizationID, color, formatTime(createdAt), formatTime(expiresAt)); err != nil {
t.Fatalf("insert authorization: %v", err)
}
}
func assertClaimState(t *testing.T, database *sql.DB, wantClaims int, wantTaskStatus, wantAuthorizationStatus string) {
t.Helper()
var count int
if err := database.QueryRow("SELECT COUNT(*) FROM purchase_attempt_claims").Scan(&count); err != nil || count != wantClaims {
t.Fatalf("claim count = %d, err %v, want %d", count, err, wantClaims)
}
var taskStatus, authorizationStatus string
if err := database.QueryRow(`SELECT tasks.status, order_authorizations.status FROM tasks
JOIN order_authorizations ON order_authorizations.task_id=tasks.id
WHERE tasks.id=?`, testTaskA).Scan(&taskStatus, &authorizationStatus); err != nil {
t.Fatal(err)
}
if taskStatus != wantTaskStatus || authorizationStatus != wantAuthorizationStatus {
t.Fatalf("states = %s/%s, want %s/%s", taskStatus, authorizationStatus, wantTaskStatus, wantAuthorizationStatus)
}
}
func assertNoPlaintextTokenColumnOrValue(t *testing.T, database *sql.DB, token string) {
t.Helper()
rows, err := database.Query("PRAGMA table_info(purchase_attempt_claims)")
if err != nil {
t.Fatal(err)
}
defer rows.Close()
for rows.Next() {
var cid, notNull, primaryKey int
var name, kind string
var defaultValue any
if err := rows.Scan(&cid, &name, &kind, &notNull, &defaultValue, &primaryKey); err != nil {
t.Fatal(err)
}
if name == "claim_token" {
t.Fatal("schema contains a plaintext claim_token column")
}
}
decoded, _ := hex.DecodeString(token)
var nonce, storedHash []byte
if err := database.QueryRow("SELECT claim_nonce, claim_token_sha256 FROM purchase_attempt_claims").Scan(&nonce, &storedHash); err != nil {
t.Fatal(err)
}
if bytes.Equal(nonce, decoded) || bytes.Equal(storedHash, decoded) || len(nonce) != 32 || len(storedHash) != 32 {
t.Fatal("database contains plaintext token or malformed token metadata")
}
}
+52
View File
@@ -0,0 +1,52 @@
package taskclaim
import (
"crypto/hmac"
"crypto/sha256"
"crypto/subtle"
"encoding/binary"
"encoding/hex"
"hash"
)
const tokenDomain = "cmbuyer/task-claim-token/v1\x00"
func deriveToken(secret []byte, deviceID, taskID, authorizationID, attemptID string, generation int, nonce []byte) []byte {
mac := hmac.New(sha256.New, secret)
_, _ = mac.Write([]byte(tokenDomain))
writeTokenField(mac, deviceID)
writeTokenField(mac, taskID)
writeTokenField(mac, authorizationID)
writeTokenField(mac, attemptID)
var number [8]byte
binary.BigEndian.PutUint64(number[:], uint64(generation))
_, _ = mac.Write(number[:])
writeTokenBytes(mac, nonce)
return mac.Sum(nil)
}
func writeTokenField(writer hash.Hash, value string) { writeTokenBytes(writer, []byte(value)) }
func writeTokenBytes(writer hash.Hash, value []byte) {
var size [4]byte
binary.BigEndian.PutUint32(size[:], uint32(len(value)))
_, _ = writer.Write(size[:])
_, _ = writer.Write(value)
}
func tokenHash(token []byte) []byte {
sum := sha256.Sum256(token)
return sum[:]
}
func matchingHash(left, right []byte) bool {
return len(left) == sha256.Size && len(right) == sha256.Size && subtle.ConstantTimeCompare(left, right) == 1
}
func decodeToken(value string) ([]byte, bool) {
if len(value) != sha256.Size*2 {
return nil, false
}
decoded, err := hex.DecodeString(value)
return decoded, err == nil && hex.EncodeToString(decoded) == value
}
+74
View File
@@ -0,0 +1,74 @@
// Package taskclaim owns the atomic task-claim and lease-renewal boundary.
// A claim token proves only ownership of one attempt; it is never permission to submit an order.
package taskclaim
import (
"context"
"errors"
)
var (
ErrInvalid = errors.New("invalid task claim request")
ErrIdempotencyConflict = errors.New("task claim idempotency conflict")
ErrRequiresManual = errors.New("task claim requires manual recovery")
ErrNotCurrent = errors.New("task claim is not current")
ErrDeviceInactive = errors.New("task claim device is inactive")
)
type ClaimCommand struct {
SessionID string `json:"session_id"`
ClaimRequestID string `json:"claim_request_id"`
}
type RenewCommand struct {
TaskID string `json:"-"`
RenewRequestID string `json:"renew_request_id"`
SessionID string `json:"session_id"`
AttemptID string `json:"attempt_id"`
ClaimGeneration int `json:"claim_generation"`
ClaimToken string `json:"claim_token"`
ExpectedLeaseExpiresAt string `json:"expected_lease_expires_at"`
}
type ClaimedTask struct {
ID string `json:"id"`
Version int `json:"version"`
Title string `json:"title"`
ProductURL string `json:"product_url"`
GoodsID string `json:"goods_id"`
SKUColor string `json:"sku_color"`
SKUSize string `json:"sku_size"`
Quantity int `json:"quantity"`
MaxTotalPrice string `json:"max_total_price"`
}
type ClaimedAuthorization struct {
ID string `json:"id"`
TaskVersion int `json:"task_version"`
ExpiresAt string `json:"expires_at"`
}
type ClaimedAttempt struct {
ID string `json:"id"`
ClaimToken string `json:"claim_token"`
ClaimGeneration int `json:"claim_generation"`
LeaseExpiresAt string `json:"lease_expires_at"`
}
type ClaimResponse struct {
Task ClaimedTask `json:"task"`
Authorization ClaimedAuthorization `json:"authorization"`
Attempt ClaimedAttempt `json:"attempt"`
}
type RenewResponse struct {
TaskID string `json:"task_id"`
AttemptID string `json:"attempt_id"`
ClaimGeneration int `json:"claim_generation"`
LeaseExpiresAt string `json:"lease_expires_at"`
}
type Service interface {
ClaimNext(context.Context, string, ClaimCommand) (ClaimResponse, bool, error)
Renew(context.Context, string, RenewCommand) (RenewResponse, error)
}