Files

374 lines
8.5 KiB
Go

package sqlite
import (
"context"
"database/sql"
"errors"
"fmt"
"strings"
"time"
"cmroubao/backend-api/internal/domain"
"cmroubao/backend-api/internal/usecase"
sqlite3 "github.com/mattn/go-sqlite3"
)
const storageTimestampLayout = "2006-01-02T15:04:05.000000000Z"
type queryRower interface {
QueryRowContext(context.Context, string, ...any) *sql.Row
}
type queryer interface {
queryRower
QueryContext(context.Context, string, ...any) (*sql.Rows, error)
}
type rowScanner interface {
Scan(...any) error
}
func scanAsset(scanner rowScanner) (domain.Asset, error) {
var asset domain.Asset
var createdAt string
err := scanner.Scan(
&asset.ID,
&asset.CreatorSubject,
&asset.Purpose,
&asset.MediaType,
&asset.SizeBytes,
&asset.SHA256,
&asset.StorageKey,
&createdAt,
)
if err != nil {
return domain.Asset{}, err
}
asset.CreatedAt, err = parseTimestamp(createdAt)
if err != nil {
return domain.Asset{}, err
}
return asset, nil
}
func scanTask(scanner rowScanner) (domain.PurchaseTask, error) {
var task domain.PurchaseTask
var createdByUserID sql.NullString
var sourceRef sql.NullString
var maxBudget sql.NullInt64
var claimedByUserID sql.NullString
var claimedByDeviceID sql.NullString
var claimTokenHash sql.NullString
var claimIssuedAt sql.NullString
var claimExpiresAt sql.NullString
var cancelReason sql.NullString
var cancelRequestedAt sql.NullString
var cancelRequestedByUserID sql.NullString
var canceledAt sql.NullString
var createdAt string
var updatedAt string
err := scanner.Scan(
&task.ID,
&task.CreatorSubject,
&createdByUserID,
&sourceRef,
&task.Title,
&task.Description,
&task.SKU,
&task.ImageAssetID,
&task.Quantity,
&maxBudget,
&task.Currency,
&task.Status,
&task.Version,
&claimedByUserID,
&claimedByDeviceID,
&task.ClaimGeneration,
&claimTokenHash,
&claimIssuedAt,
&claimExpiresAt,
&cancelReason,
&cancelRequestedAt,
&cancelRequestedByUserID,
&canceledAt,
&createdAt,
&updatedAt,
)
if err != nil {
return domain.PurchaseTask{}, err
}
if sourceRef.Valid {
task.SourceRef = &sourceRef.String
}
if createdByUserID.Valid {
task.CreatedByUserID = &createdByUserID.String
}
if maxBudget.Valid {
task.MaxBudgetCents = &maxBudget.Int64
}
if claimedByUserID.Valid {
task.ClaimedByUserID = &claimedByUserID.String
}
if claimedByDeviceID.Valid {
task.ClaimedByDeviceID = &claimedByDeviceID.String
}
if claimTokenHash.Valid {
task.ClaimTokenHash = &claimTokenHash.String
}
task.ClaimIssuedAt, err = parseNullableTimestamp(claimIssuedAt)
if err != nil {
return domain.PurchaseTask{}, err
}
task.ClaimExpiresAt, err = parseNullableTimestamp(claimExpiresAt)
if err != nil {
return domain.PurchaseTask{}, err
}
if cancelReason.Valid {
task.CancelReason = &cancelReason.String
}
task.CancelRequestedAt, err = parseNullableTimestamp(cancelRequestedAt)
if err != nil {
return domain.PurchaseTask{}, err
}
if cancelRequestedByUserID.Valid {
task.CancelRequestedByUserID = &cancelRequestedByUserID.String
}
task.CanceledAt, err = parseNullableTimestamp(canceledAt)
if err != nil {
return domain.PurchaseTask{}, err
}
task.CreatedAt, err = parseTimestamp(createdAt)
if err != nil {
return domain.PurchaseTask{}, err
}
task.UpdatedAt, err = parseTimestamp(updatedAt)
if err != nil {
return domain.PurchaseTask{}, err
}
return task, nil
}
func scanExecution(scanner rowScanner) (domain.TaskExecution, error) {
var execution domain.TaskExecution
var lastHeartbeatAt sql.NullString
var finishedAt sql.NullString
var startedAt string
err := scanner.Scan(
&execution.ID,
&execution.TaskID,
&execution.AttemptNo,
&execution.ClaimGeneration,
&execution.UserID,
&execution.DeviceID,
&execution.CurrentStep,
&execution.OrderSubmitted,
&startedAt,
&lastHeartbeatAt,
&finishedAt,
)
if err != nil {
return domain.TaskExecution{}, err
}
execution.StartedAt, err = parseTimestamp(startedAt)
if err != nil {
return domain.TaskExecution{}, err
}
execution.LastHeartbeatAt, err = parseNullableTimestamp(lastHeartbeatAt)
if err != nil {
return domain.TaskExecution{}, err
}
execution.FinishedAt, err = parseNullableTimestamp(finishedAt)
if err != nil {
return domain.TaskExecution{}, err
}
return execution, nil
}
func getAssetByID(
ctx context.Context,
queryer queryRower,
creatorSubject string,
assetID string,
) (domain.Asset, error) {
asset, err := scanAsset(queryer.QueryRowContext(
ctx,
`SELECT
id, creator_subject, purpose, media_type, size_bytes,
sha256, storage_key, created_at
FROM assets
WHERE creator_subject = ? AND id = ?`,
creatorSubject,
assetID,
))
if errors.Is(err, sql.ErrNoRows) {
return domain.Asset{}, usecase.ErrRepositoryNotFound
}
if err != nil {
return domain.Asset{}, repositoryFailure(err)
}
return asset, nil
}
func getTaskByID(
ctx context.Context,
queryer queryRower,
creatorSubject string,
taskID string,
) (domain.PurchaseTask, error) {
task, err := scanTask(queryer.QueryRowContext(
ctx,
`SELECT
id, creator_subject, created_by_user_id, source_ref, title, description, sku,
image_asset_id, quantity, max_budget_cents, currency, status,
version, claimed_by_user_id, claimed_by_device_id, claim_generation,
claim_token_hash, claim_issued_at, claim_expires_at, cancel_reason,
cancel_requested_at, cancel_requested_by_user_id, canceled_at,
created_at, updated_at
FROM purchase_tasks
WHERE creator_subject = ? AND id = ?`,
creatorSubject,
taskID,
))
if errors.Is(err, sql.ErrNoRows) {
return domain.PurchaseTask{}, usecase.ErrRepositoryNotFound
}
if err != nil {
return domain.PurchaseTask{}, repositoryFailure(err)
}
return task, nil
}
func lookupIdempotency(
ctx context.Context,
tx *sql.Tx,
creatorSubject string,
operation string,
idempotencyKey string,
) (requestHash string, resourceID string, found bool, err error) {
err = tx.QueryRowContext(
ctx,
`SELECT request_sha256, resource_id
FROM idempotency_records
WHERE creator_subject = ?
AND operation = ?
AND idempotency_key = ?`,
creatorSubject,
operation,
idempotencyKey,
).Scan(&requestHash, &resourceID)
if errors.Is(err, sql.ErrNoRows) {
return "", "", false, nil
}
if err != nil {
return "", "", false, repositoryFailure(err)
}
return requestHash, resourceID, true, nil
}
func insertIdempotency(
ctx context.Context,
tx *sql.Tx,
creatorSubject string,
operation string,
idempotencyKey string,
requestHash string,
resourceType string,
resourceID string,
createdAt time.Time,
) error {
_, err := tx.ExecContext(
ctx,
`INSERT INTO idempotency_records (
creator_subject, operation, idempotency_key, request_sha256,
resource_type, resource_id, created_at
) VALUES (?, ?, ?, ?, ?, ?, ?)`,
creatorSubject,
operation,
idempotencyKey,
requestHash,
resourceType,
resourceID,
formatTimestamp(createdAt),
)
if err != nil {
return repositoryFailure(err)
}
return nil
}
func formatTimestamp(value time.Time) string {
return value.UTC().Format(storageTimestampLayout)
}
func parseTimestamp(value string) (time.Time, error) {
parsed, err := time.Parse(time.RFC3339Nano, value)
if err != nil {
return time.Time{}, fmt.Errorf(
"%w: invalid stored timestamp",
usecase.ErrRepositoryInvariant,
)
}
return parsed.UTC(), nil
}
func parseNullableTimestamp(value sql.NullString) (*time.Time, error) {
if !value.Valid {
return nil, nil
}
parsed, err := parseTimestamp(value.String)
if err != nil {
return nil, err
}
return &parsed, nil
}
func nullableString(value *string) any {
if value == nil {
return nil
}
return *value
}
func nullableInt64(value *int64) any {
if value == nil {
return nil
}
return *value
}
func nullableBool(value *bool) any {
if value == nil {
return nil
}
return *value
}
func repositoryFailure(err error) error {
if err == nil {
return nil
}
var sqliteError sqlite3.Error
if errors.As(err, &sqliteError) {
switch sqliteError.Code {
case sqlite3.ErrBusy, sqlite3.ErrLocked, sqlite3.ErrIoErr,
sqlite3.ErrCantOpen, sqlite3.ErrFull:
return fmt.Errorf("%w", usecase.ErrRepositoryUnavailable)
}
}
if errors.Is(err, context.Canceled) ||
errors.Is(err, context.DeadlineExceeded) {
return fmt.Errorf("%w", usecase.ErrRepositoryUnavailable)
}
return fmt.Errorf("%w", usecase.ErrRepositoryInvariant)
}
func isUniqueConstraint(err error, fragment string) bool {
var sqliteError sqlite3.Error
if !errors.As(err, &sqliteError) ||
sqliteError.ExtendedCode != sqlite3.ErrConstraintUnique {
return false
}
return strings.Contains(err.Error(), fragment)
}