feat(admin): add atomic task claim leases
This commit is contained in:
@@ -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
|
||||
}
|
||||
@@ -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, ¬Null, &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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
Reference in New Issue
Block a user