easyai-ai-gateway/apps/api/internal/store/billing_v2_integration_test.go
chengcheng 8beb8501fa
Some checks failed
ci / verify (pull_request) Failing after 8s
fix(billing): 封住发布前计费竞态
阻止上游提交状态不明的任务被租约接管后重复执行,并为人工复核保留可操作的结算记录。

将生产提交绑定到当前估价签名,统一复用预处理快照,并补强规则形状、定点溢出与历史规则兼容校验。

已通过 PostgreSQL 16 集成测试、Go 全量测试与静态检查、前端测试与构建、OpenAPI、依赖审计、镜像、迁移、流水线和 SemVer 门禁。
2026-07-21 10:23:58 +08:00

409 lines
17 KiB
Go

package store
import (
"context"
"errors"
"os"
"strings"
"testing"
"time"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
"github.com/google/uuid"
"github.com/jackc/pgx/v5"
)
func TestTaskIdempotencyAndExecutionLease(t *testing.T) {
db := billingV2IntegrationStore(t)
ctx := context.Background()
tenantID, gatewayUserID := seedWalletReservationUser(t, ctx, db)
user := &auth.User{ID: "billing-v2-" + uuid.NewString(), GatewayUserID: gatewayUserID, GatewayTenantID: tenantID}
input := CreateTaskInput{
Kind: "images.generations", Model: "billing-v2-model", RunMode: "simulation",
Request: map[string]any{"model": "billing-v2-model", "prompt": "lease"},
IdempotencyKeyHash: "key-hash-" + uuid.NewString(), IdempotencyRequestHash: "request-a",
}
created, err := db.CreateTaskIdempotent(ctx, input, user)
if err != nil || created.Replayed {
t.Fatalf("create task result=%+v err=%v", created, err)
}
t.Cleanup(func() {
_, _ = db.pool.Exec(context.Background(), `DELETE FROM gateway_tasks WHERE id=$1::uuid`, created.Task.ID)
})
replayed, err := db.CreateTaskIdempotent(ctx, input, user)
if err != nil || !replayed.Replayed || replayed.Task.ID != created.Task.ID {
t.Fatalf("replay result=%+v err=%v", replayed, err)
}
input.IdempotencyRequestHash = "request-b"
if _, err := db.CreateTaskIdempotent(ctx, input, user); !errors.Is(err, ErrIdempotencyKeyReused) {
t.Fatalf("different request error=%v", err)
}
firstToken := uuid.NewString()
claimed, err := db.ClaimTaskExecution(ctx, created.Task.ID, firstToken, 5*time.Minute)
if err != nil || claimed.ExecutionToken != firstToken {
t.Fatalf("first claim task=%+v err=%v", claimed, err)
}
if err := db.RenewTaskExecutionLease(ctx, created.Task.ID, firstToken, 5*time.Minute); err != nil {
t.Fatalf("renew first lease: %v", err)
}
if _, err := db.pool.Exec(ctx, `UPDATE gateway_tasks SET execution_lease_expires_at=now()-interval '1 second' WHERE id=$1::uuid`, created.Task.ID); err != nil {
t.Fatal(err)
}
secondToken := uuid.NewString()
if _, err := db.ClaimTaskExecution(ctx, created.Task.ID, secondToken, 5*time.Minute); err != nil {
t.Fatalf("take over expired lease: %v", err)
}
if _, err := db.FinishTaskFailure(ctx, FinishTaskFailureInput{TaskID: created.Task.ID, ExecutionToken: firstToken, Code: "old_worker", Message: "old"}); !errors.Is(err, ErrTaskExecutionLeaseLost) {
t.Fatalf("old worker terminal error=%v", err)
}
finished, err := db.FinishTaskFailure(ctx, FinishTaskFailureInput{TaskID: created.Task.ID, ExecutionToken: secondToken, Code: "new_worker", Message: "new"})
if err != nil || finished.ErrorCode != "new_worker" {
t.Fatalf("new worker terminal task=%+v err=%v", finished, err)
}
if err := db.RenewTaskExecutionLease(ctx, created.Task.ID, secondToken, 5*time.Minute); !errors.Is(err, ErrTaskExecutionFinished) {
t.Fatalf("terminal task renewal error=%v", err)
}
}
func TestExpiredExecutionLeaseDoesNotReplayAmbiguousUpstreamSubmission(t *testing.T) {
db := billingV2IntegrationStore(t)
ctx := context.Background()
tenantID, gatewayUserID := seedWalletReservationUser(t, ctx, db)
user := &auth.User{ID: "billing-review-" + uuid.NewString(), GatewayUserID: gatewayUserID, GatewayTenantID: tenantID}
created, err := db.CreateTask(ctx, CreateTaskInput{
Kind: "images.generations", Model: "billing-v2-model", RunMode: "production",
Request: map[string]any{"model": "billing-v2-model"},
}, user)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
_, _ = db.pool.Exec(context.Background(), `DELETE FROM gateway_tasks WHERE id=$1::uuid`, created.ID)
})
firstToken := uuid.NewString()
if _, err := db.ClaimTaskExecution(ctx, created.ID, firstToken, 5*time.Minute); err != nil {
t.Fatal(err)
}
attemptID, err := db.CreateTaskAttempt(ctx, CreateTaskAttemptInput{
TaskID: created.ID, AttemptNo: 1, Status: "running", RequestSnapshot: map[string]any{"model": created.Model},
})
if err != nil {
t.Fatal(err)
}
if err := db.SetAttemptUpstreamSubmissionStatus(ctx, attemptID, "submitting"); err != nil {
t.Fatal(err)
}
if _, err := db.pool.Exec(ctx, `UPDATE gateway_tasks SET execution_lease_expires_at=now()-interval '1 second' WHERE id=$1::uuid`, created.ID); err != nil {
t.Fatal(err)
}
if _, err := db.ClaimTaskExecution(ctx, created.ID, uuid.NewString(), 5*time.Minute); !errors.Is(err, ErrTaskExecutionManualReview) {
t.Fatalf("ambiguous submission takeover error=%v", err)
}
review, err := db.GetTask(ctx, created.ID)
if err != nil {
t.Fatal(err)
}
if review.Status != "failed" || review.BillingStatus != "manual_review" || review.ErrorCode != "upstream_submission_unknown" {
t.Fatalf("manual review task=%+v", review)
}
var outboxStatus string
var outboxAction string
var reviewReason string
if err := db.pool.QueryRow(ctx, `
SELECT status, action, COALESCE(manual_review_reason, '')
FROM settlement_outbox
WHERE task_id=$1::uuid AND event_type='task.billing.review'`, created.ID).Scan(&outboxStatus, &outboxAction, &reviewReason); err != nil {
t.Fatal(err)
}
if outboxStatus != "manual_review" || outboxAction != "release" || reviewReason != "upstream_submission_unknown" {
t.Fatalf("review outbox status=%s action=%s reason=%s", outboxStatus, outboxAction, reviewReason)
}
}
func TestExpiredExecutionLeaseCanResumeAfterKnownRejectedResponse(t *testing.T) {
db := billingV2IntegrationStore(t)
ctx := context.Background()
_, gatewayUserID := seedWalletReservationUser(t, ctx, db)
user := &auth.User{ID: "billing-known-response-" + uuid.NewString(), GatewayUserID: gatewayUserID}
created, err := db.CreateTask(ctx, CreateTaskInput{
Kind: "images.generations", Model: "billing-v2-model", RunMode: "production",
Request: map[string]any{"model": "billing-v2-model"},
}, user)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
_, _ = db.pool.Exec(context.Background(), `DELETE FROM gateway_tasks WHERE id=$1::uuid`, created.ID)
})
if _, err := db.ClaimTaskExecution(ctx, created.ID, uuid.NewString(), 5*time.Minute); err != nil {
t.Fatal(err)
}
attemptID, err := db.CreateTaskAttempt(ctx, CreateTaskAttemptInput{TaskID: created.ID, AttemptNo: 1, Status: "running"})
if err != nil {
t.Fatal(err)
}
if err := db.SetAttemptUpstreamSubmissionStatus(ctx, attemptID, "response_received"); err != nil {
t.Fatal(err)
}
if err := db.FinishTaskAttempt(ctx, FinishTaskAttemptInput{AttemptID: attemptID, Status: "failed", ErrorCode: "upstream_rejected"}); err != nil {
t.Fatal(err)
}
if _, err := db.pool.Exec(ctx, `UPDATE gateway_tasks SET execution_lease_expires_at=now()-interval '1 second' WHERE id=$1::uuid`, created.ID); err != nil {
t.Fatal(err)
}
if _, err := db.ClaimTaskExecution(ctx, created.ID, uuid.NewString(), 5*time.Minute); err != nil {
t.Fatalf("known rejected response should remain retryable: %v", err)
}
}
func TestFinishTaskManualReviewCreatesVisibleBillingRecord(t *testing.T) {
db := billingV2IntegrationStore(t)
ctx := context.Background()
_, gatewayUserID := seedWalletReservationUser(t, ctx, db)
user := &auth.User{ID: "billing-direct-review-" + uuid.NewString(), GatewayUserID: gatewayUserID}
created, err := db.CreateTask(ctx, CreateTaskInput{
Kind: "images.generations", Model: "billing-v2-model", RunMode: "production",
Request: map[string]any{"model": "billing-v2-model"},
}, user)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
_, _ = db.pool.Exec(context.Background(), `DELETE FROM gateway_tasks WHERE id=$1::uuid`, created.ID)
})
token := uuid.NewString()
if _, err := db.ClaimTaskExecution(ctx, created.ID, token, 5*time.Minute); err != nil {
t.Fatal(err)
}
attemptID, err := db.CreateTaskAttempt(ctx, CreateTaskAttemptInput{TaskID: created.ID, AttemptNo: 1, Status: "running"})
if err != nil {
t.Fatal(err)
}
if err := db.SetAttemptUpstreamSubmissionStatus(ctx, attemptID, "submitting"); err != nil {
t.Fatal(err)
}
if _, err := db.FinishTaskManualReview(ctx, FinishTaskManualReviewInput{
TaskID: created.ID, ExecutionToken: token, AttemptID: attemptID, TaskStatus: "failed",
Code: "upstream_submission_unknown", Message: "upstream submission result is unknown",
}); err != nil {
t.Fatal(err)
}
var visible bool
if err := db.pool.QueryRow(ctx, `
SELECT EXISTS (
SELECT 1 FROM settlement_outbox
WHERE task_id=$1::uuid AND status='manual_review'
AND action='release' AND manual_review_reason='upstream_submission_unknown'
)`, created.ID).Scan(&visible); err != nil {
t.Fatal(err)
}
if !visible {
t.Fatal("manual review billing record is not visible in settlement outbox")
}
}
func TestBillingSettlementStaleTakeoverDebitsExactlyOnce(t *testing.T) {
db := billingV2IntegrationStore(t)
ctx := context.Background()
tenantID, gatewayUserID := seedWalletReservationUser(t, ctx, db)
user := &auth.User{ID: "billing-settle-" + uuid.NewString(), GatewayUserID: gatewayUserID, GatewayTenantID: tenantID}
if _, err := db.SetUserWalletBalance(ctx, WalletBalanceAdjustmentInput{GatewayUserID: gatewayUserID, Currency: "resource", Balance: 10, Reason: "billing v2 test"}); err != nil {
t.Fatal(err)
}
created, err := db.CreateTask(ctx, CreateTaskInput{Kind: "images.generations", Model: "billing-v2-model", RunMode: "production", Request: map[string]any{"model": "billing-v2-model"}}, user)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
cleanupCtx := context.Background()
_, _ = db.pool.Exec(cleanupCtx, `DELETE FROM gateway_wallet_transactions WHERE gateway_user_id=$1::uuid`, gatewayUserID)
_, _ = db.pool.Exec(cleanupCtx, `DELETE FROM gateway_tasks WHERE id=$1::uuid`, created.ID)
_, _ = db.pool.Exec(cleanupCtx, `DELETE FROM gateway_wallet_accounts WHERE gateway_user_id=$1::uuid`, gatewayUserID)
})
token := uuid.NewString()
claimed, err := db.ClaimTaskExecution(ctx, created.ID, token, 5*time.Minute)
if err != nil {
t.Fatal(err)
}
amount := "1.000000001"
reservations, err := db.ReserveTaskBilling(ctx, claimed, user, []any{map[string]any{"currency": "resource", "amount": amount}}, map[string]any{
"pricingVersion": "effective-pricing-v2", "reservationAmount": amount, "currency": "resource", "requestFingerprint": "billing-v2-test",
})
if err != nil || len(reservations) != 1 {
t.Fatalf("reserve=%+v err=%v", reservations, err)
}
if _, err := db.FinishTaskSuccess(ctx, FinishTaskSuccessInput{
TaskID: created.ID, ExecutionToken: token, Result: map[string]any{"ok": true},
FinalChargeAmountText: amount, BillingCurrency: "resource",
PricingSnapshot: map[string]any{"pricingVersion": "effective-pricing-v2"},
}); err != nil {
t.Fatal(err)
}
firstClaims, err := db.ClaimBillingSettlements(ctx, "worker-one", BillingSettlementBatchSize, BillingSettlementLockTimeout)
if err != nil {
t.Fatal(err)
}
first := settlementForTask(t, firstClaims, created.ID)
if _, err := db.pool.Exec(ctx, `UPDATE settlement_outbox SET locked_at=now()-interval '3 minutes' WHERE id=$1::uuid`, first.ID); err != nil {
t.Fatal(err)
}
secondClaims, err := db.ClaimBillingSettlements(ctx, "worker-two", BillingSettlementBatchSize, BillingSettlementLockTimeout)
if err != nil {
t.Fatal(err)
}
second := settlementForTask(t, secondClaims, created.ID)
if err := db.ProcessBillingSettlement(ctx, first); !errors.Is(err, pgx.ErrNoRows) {
t.Fatalf("stale settlement lock error=%v", err)
}
if err := db.ProcessBillingSettlement(ctx, second); err != nil {
t.Fatalf("process takeover: %v", err)
}
var walletExact bool
if err := db.pool.QueryRow(ctx, `
SELECT balance = 8.999999999::numeric
AND frozen_balance = 0::numeric
AND total_spent = 1.000000001::numeric
FROM gateway_wallet_accounts
WHERE gateway_user_id=$1::uuid AND currency='resource'`, gatewayUserID).Scan(&walletExact); err != nil {
t.Fatal(err)
}
if !walletExact {
t.Fatal("wallet amounts did not preserve nine-decimal settlement")
}
var billingTransactions int
if err := db.pool.QueryRow(ctx, `SELECT count(*) FROM gateway_wallet_transactions WHERE reference_id=$1 AND transaction_type='task_billing'`, created.ID).Scan(&billingTransactions); err != nil {
t.Fatal(err)
}
if billingTransactions != 1 {
t.Fatalf("task billing transactions=%d", billingTransactions)
}
if _, err := db.pool.Exec(ctx, `
DELETE FROM gateway_wallet_accounts
WHERE gateway_user_id=$1::uuid AND currency='resource'`, gatewayUserID); err == nil {
t.Fatal("wallet account with audit transactions must not be deletable")
}
if err := db.pool.QueryRow(ctx, `
SELECT count(*) FROM gateway_wallet_transactions
WHERE gateway_user_id=$1::uuid AND reference_id=$2`, gatewayUserID, created.ID).Scan(&billingTransactions); err != nil {
t.Fatal(err)
}
if billingTransactions == 0 {
t.Fatal("wallet audit transactions were lost after rejected account deletion")
}
settled, err := db.GetTask(ctx, created.ID)
if err != nil || settled.BillingStatus != "settled" {
t.Fatalf("settled task=%+v err=%v", settled, err)
}
}
func TestReleaseTaskBillingReservationsPreservesNineDecimalPlaces(t *testing.T) {
db := billingV2IntegrationStore(t)
ctx := context.Background()
tenantID, gatewayUserID := seedWalletReservationUser(t, ctx, db)
user := &auth.User{ID: "billing-release-" + uuid.NewString(), GatewayUserID: gatewayUserID, GatewayTenantID: tenantID}
if _, err := db.SetUserWalletBalance(ctx, WalletBalanceAdjustmentInput{
GatewayUserID: gatewayUserID, Currency: "resource", Balance: 1, Reason: "billing v2 release test",
}); err != nil {
t.Fatal(err)
}
created, err := db.CreateTask(ctx, CreateTaskInput{
Kind: "images.generations", Model: "billing-v2-model", RunMode: "production",
Request: map[string]any{"model": "billing-v2-model"},
}, user)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
cleanupCtx := context.Background()
_, _ = db.pool.Exec(cleanupCtx, `DELETE FROM gateway_wallet_transactions WHERE gateway_user_id=$1::uuid`, gatewayUserID)
_, _ = db.pool.Exec(cleanupCtx, `DELETE FROM gateway_tasks WHERE id=$1::uuid`, created.ID)
_, _ = db.pool.Exec(cleanupCtx, `DELETE FROM gateway_wallet_accounts WHERE gateway_user_id=$1::uuid`, gatewayUserID)
})
amount := "0.000000001"
reservations, err := db.ReserveTaskBilling(ctx, created, user, nil, map[string]any{
"pricingVersion": "effective-pricing-v2", "reservationAmount": amount, "currency": "resource",
})
if err != nil || len(reservations) != 1 {
t.Fatalf("reserve=%+v err=%v", reservations, err)
}
if _, err := db.SetUserWalletBalance(ctx, WalletBalanceAdjustmentInput{
GatewayUserID: gatewayUserID, Currency: "resource", BalanceText: "0", Reason: "must not cross frozen balance",
}); !errors.Is(err, ErrBalanceBelowFrozen) {
t.Fatalf("balance below frozen error=%v", err)
}
if err := db.ReleaseTaskBillingReservations(ctx, reservations, "integration_test"); err != nil {
t.Fatal(err)
}
if _, err := db.SetUserWalletBalance(ctx, WalletBalanceAdjustmentInput{
GatewayUserID: gatewayUserID, Currency: "resource", BalanceText: "0.123456789", Reason: "exact adjustment",
}); err != nil {
t.Fatal(err)
}
if _, err := db.RechargeUserWalletBalance(ctx, WalletRechargeInput{
GatewayUserID: gatewayUserID, Currency: "resource", AmountText: "0.000000001", Reason: "exact recharge",
}); err != nil {
t.Fatal(err)
}
var exact bool
if err := db.pool.QueryRow(ctx, `
SELECT account.balance = 0.123456790::numeric
AND account.frozen_balance = 0::numeric
AND task.reservation_amount = 0::numeric
AND task.billing_status = 'not_started'
AND EXISTS (
SELECT 1 FROM gateway_wallet_transactions transaction
WHERE transaction.reference_id = task.id::text
AND transaction.transaction_type = 'release'
AND transaction.amount = 0.000000001::numeric
)
FROM gateway_wallet_accounts account
JOIN gateway_tasks task ON task.gateway_user_id = account.gateway_user_id
WHERE task.id=$1::uuid AND account.currency='resource'`, created.ID).Scan(&exact); err != nil {
t.Fatal(err)
}
if !exact {
t.Fatal("nine-decimal reservation was not released exactly")
}
}
func billingV2IntegrationStore(t *testing.T) *Store {
t.Helper()
databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL"))
if databaseURL == "" {
t.Skip("set AI_GATEWAY_TEST_DATABASE_URL to run billing v2 PostgreSQL integration tests")
}
db, err := Connect(context.Background(), databaseURL)
if err != nil {
t.Fatal(err)
}
var databaseName string
if err := db.pool.QueryRow(context.Background(), `SELECT current_database()`).Scan(&databaseName); err != nil {
db.Close()
t.Fatal(err)
}
if !strings.Contains(strings.ToLower(databaseName), "test") {
db.Close()
t.Fatalf("refusing to use non-test database %q", databaseName)
}
t.Cleanup(db.Close)
return db
}
func settlementForTask(t *testing.T, items []BillingSettlement, taskID string) BillingSettlement {
t.Helper()
for _, item := range items {
if item.TaskID == taskID {
return item
}
}
t.Fatalf("settlement for task %s not claimed", taskID)
return BillingSettlement{}
}