feat(billing): 完成异步结算与请求执行闭环

统一任务成功、Attempt 与结算 Outbox 的事务边界,增加多实例安全结算、人工复核、请求幂等与执行租约。钱包决策使用九位精确金额,并通过审计保护约束保留流水事实;同时补充管理接口、指标与 PostgreSQL 集成测试。
This commit is contained in:
2026-07-21 00:25:53 +08:00
parent 7cea21f765
commit 5b2b94b1bd
33 changed files with 2884 additions and 418 deletions
+6
View File
@@ -23,6 +23,12 @@ CONFIG_JWT_SECRET=this is a very secret secret
# - hybrid: both sources are accepted and separated by gateway_users.source. # - hybrid: both sources are accepted and separated by gateway_users.source.
IDENTITY_MODE=hybrid IDENTITY_MODE=hybrid
# Billing engine rollout mode:
# - observe: keep legacy billing decisions and compare effective-pricing-v2 in logs.
# - enforce: require v2 pricing, reserve the candidate maximum, then settle asynchronously.
# - hold: reject new production generation before any upstream request; existing settlements continue.
BILLING_ENGINE_MODE=observe
# Unified identity business settings are managed in System Settings > Unified # Unified identity business settings are managed in System Settings > Unified
# Identity. Deployment only supplies the SecretStore and infrastructure timing. # Identity. Deployment only supplies the SecretStore and infrastructure timing.
AI_GATEWAY_PUBLIC_BASE_URL=http://localhost:8088 AI_GATEWAY_PUBLIC_BASE_URL=http://localhost:8088
@@ -34,6 +34,11 @@ var gatewayOpenAIRequestExtensions = stringSet(
"requestId", "request_id", "signal", "userMessage", "user_message", "platformId", "requestId", "request_id", "signal", "userMessage", "user_message", "platformId",
"platform_id", "options", "enable_thinking", "thinking_budget_tokens", "enable_web_search", "platform_id", "options", "enable_thinking", "thinking_budget_tokens", "enable_web_search",
"modelType", "model_type", "capability", "capabilityType", "mode", "simulation", "testMode", "modelType", "model_type", "capability", "capabilityType", "mode", "simulation", "testMode",
"cacheAffinityKey", "cache_affinity_key", "simulationDurationMs", "simulationDurationSeconds",
"simulationMinDurationMs", "simulationMaxDurationMs", "simulationMinDurationSeconds",
"simulationMaxDurationSeconds", "simulationDurationMinMs", "simulationDurationMaxMs",
"simulationDurationMinSeconds", "simulationDurationMaxSeconds", "simulationFailure",
"simulationProfile", "simulationUsage",
) )
var gatewayResponsesRequestExtensions = stringSet("messages", "presence_penalty", "frequency_penalty") var gatewayResponsesRequestExtensions = stringSet("messages", "presence_penalty", "frequency_penalty")
@@ -60,6 +60,23 @@ func TestValidateOpenAIRequestParametersRejectsUnknownTopLevelField(t *testing.T
} }
} }
func TestValidateOpenAIRequestParametersAcceptsInternalSimulationAndAffinityFields(t *testing.T) {
body := map[string]any{
"model": "demo", "messages": []any{}, "simulation": true,
"simulationDurationMs": 5, "simulationUsage": map[string]any{"inputTokens": 10},
"cacheAffinityKey": "conversation-1",
}
if err := ValidateOpenAIRequestParameters("chat.completions", body); err != nil {
t.Fatalf("expected documented gateway extensions to remain accepted, got %v", err)
}
filtered := FilterOpenAIChatRequestBody(body)
for _, key := range []string{"simulation", "simulationDurationMs", "simulationUsage", "cacheAffinityKey"} {
if _, ok := filtered[key]; ok {
t.Fatalf("gateway-only field %q leaked upstream", key)
}
}
}
func TestResponsesFallbackMapsEquivalentCurrentParameters(t *testing.T) { func TestResponsesFallbackMapsEquivalentCurrentParameters(t *testing.T) {
body, err := ResponsesRequestToChat(map[string]any{ body, err := ResponsesRequestToChat(map[string]any{
"input": "hello", "store": false, "metadata": map[string]any{"trace": "1"}, "input": "hello", "store": false, "metadata": map[string]any{"trace": "1"},
+1
View File
@@ -36,6 +36,7 @@ type ResponseTurn struct {
} }
type Response struct { type Response struct {
AttemptID string
Result map[string]any Result map[string]any
RequestID string RequestID string
Usage Usage Usage Usage
+7
View File
@@ -47,6 +47,7 @@ type Config struct {
GlobalHTTPProxy string GlobalHTTPProxy string
GlobalHTTPProxySource string GlobalHTTPProxySource string
LogLevel slog.Level LogLevel slog.Level
BillingEngineMode string
} }
func Load() Config { func Load() Config {
@@ -90,10 +91,16 @@ func Load() Config {
GlobalHTTPProxy: globalProxy.HTTPProxy, GlobalHTTPProxy: globalProxy.HTTPProxy,
GlobalHTTPProxySource: globalProxy.Source, GlobalHTTPProxySource: globalProxy.Source,
LogLevel: logLevel(env("LOG_LEVEL", "info")), LogLevel: logLevel(env("LOG_LEVEL", "info")),
BillingEngineMode: strings.ToLower(env("BILLING_ENGINE_MODE", "observe")),
} }
} }
func (c Config) Validate() error { func (c Config) Validate() error {
switch strings.ToLower(strings.TrimSpace(c.BillingEngineMode)) {
case "", "observe", "enforce", "hold":
default:
return errors.New("BILLING_ENGINE_MODE must be observe, enforce, or hold")
}
switch strings.ToLower(strings.TrimSpace(c.IdentitySecretStore)) { switch strings.ToLower(strings.TrimSpace(c.IdentitySecretStore)) {
case "": case "":
case "file": case "file":
@@ -13,7 +13,7 @@ import (
type walletBalanceRequest struct { type walletBalanceRequest struct {
Currency string `json:"currency" example:"USD"` Currency string `json:"currency" example:"USD"`
Balance float64 `json:"balance" example:"100"` Balance json.Number `json:"balance" swaggertype:"number" example:"100"`
Reason string `json:"reason" example:"manual recharge"` Reason string `json:"reason" example:"manual recharge"`
IdempotencyKey string `json:"idempotencyKey" example:"wallet-set-20260514-001"` IdempotencyKey string `json:"idempotencyKey" example:"wallet-set-20260514-001"`
Metadata map[string]any `json:"metadata"` Metadata map[string]any `json:"metadata"`
@@ -21,7 +21,7 @@ type walletBalanceRequest struct {
type walletRechargeRequest struct { type walletRechargeRequest struct {
Currency string `json:"currency" example:"resource"` Currency string `json:"currency" example:"resource"`
Amount float64 `json:"amount" example:"100"` Amount json.Number `json:"amount" swaggertype:"number" example:"100"`
Reason string `json:"reason" example:"manual recharge"` Reason string `json:"reason" example:"manual recharge"`
IdempotencyKey string `json:"idempotencyKey" example:"wallet-recharge-20260514-001"` IdempotencyKey string `json:"idempotencyKey" example:"wallet-recharge-20260514-001"`
Metadata map[string]any `json:"metadata"` Metadata map[string]any `json:"metadata"`
@@ -50,10 +50,6 @@ func (s *Server) setUserWalletBalance(w http.ResponseWriter, r *http.Request) {
writeError(w, http.StatusBadRequest, "invalid json body") writeError(w, http.StatusBadRequest, "invalid json body")
return return
} }
if input.Balance < 0 {
writeError(w, http.StatusBadRequest, "wallet balance cannot be negative")
return
}
gatewayUserID := strings.TrimSpace(r.PathValue("userID")) gatewayUserID := strings.TrimSpace(r.PathValue("userID"))
reason := strings.TrimSpace(input.Reason) reason := strings.TrimSpace(input.Reason)
if reason == "" { if reason == "" {
@@ -67,7 +63,7 @@ func (s *Server) setUserWalletBalance(w http.ResponseWriter, r *http.Request) {
next, err := s.store.SetUserWalletBalanceTx(r.Context(), tx, store.WalletBalanceAdjustmentInput{ next, err := s.store.SetUserWalletBalanceTx(r.Context(), tx, store.WalletBalanceAdjustmentInput{
GatewayUserID: gatewayUserID, GatewayUserID: gatewayUserID,
Currency: input.Currency, Currency: input.Currency,
Balance: input.Balance, BalanceText: input.Balance.String(),
Reason: reason, Reason: reason,
IdempotencyKey: input.IdempotencyKey, IdempotencyKey: input.IdempotencyKey,
Metadata: input.Metadata, Metadata: input.Metadata,
@@ -89,6 +85,10 @@ func (s *Server) setUserWalletBalance(w http.ResponseWriter, r *http.Request) {
writeError(w, http.StatusNotFound, "user not found") writeError(w, http.StatusNotFound, "user not found")
case errors.Is(err, store.ErrWalletBalanceUnchanged): case errors.Is(err, store.ErrWalletBalanceUnchanged):
writeError(w, http.StatusBadRequest, "wallet balance is unchanged") writeError(w, http.StatusBadRequest, "wallet balance is unchanged")
case errors.Is(err, store.ErrInvalidWalletAmount):
writeError(w, http.StatusBadRequest, "wallet balance must be a non-negative decimal with at most nine fractional digits", "invalid_wallet_amount")
case errors.Is(err, store.ErrBalanceBelowFrozen):
writeError(w, http.StatusConflict, "wallet balance cannot be below frozen balance", "balance_below_frozen")
default: default:
s.logger.Error("set user wallet balance failed", "error", err) s.logger.Error("set user wallet balance failed", "error", err)
writeError(w, http.StatusInternalServerError, "set user wallet balance failed") writeError(w, http.StatusInternalServerError, "set user wallet balance failed")
@@ -126,10 +126,6 @@ func (s *Server) rechargeUserWalletBalance(w http.ResponseWriter, r *http.Reques
writeError(w, http.StatusBadRequest, "invalid json body") writeError(w, http.StatusBadRequest, "invalid json body")
return return
} }
if input.Amount <= 0 {
writeError(w, http.StatusBadRequest, "wallet recharge amount must be positive")
return
}
gatewayUserID := strings.TrimSpace(r.PathValue("userID")) gatewayUserID := strings.TrimSpace(r.PathValue("userID"))
reason := strings.TrimSpace(input.Reason) reason := strings.TrimSpace(input.Reason)
if reason == "" { if reason == "" {
@@ -143,7 +139,7 @@ func (s *Server) rechargeUserWalletBalance(w http.ResponseWriter, r *http.Reques
next, err := s.store.RechargeUserWalletBalanceTx(r.Context(), tx, store.WalletRechargeInput{ next, err := s.store.RechargeUserWalletBalanceTx(r.Context(), tx, store.WalletRechargeInput{
GatewayUserID: gatewayUserID, GatewayUserID: gatewayUserID,
Currency: input.Currency, Currency: input.Currency,
Amount: input.Amount, AmountText: input.Amount.String(),
Reason: reason, Reason: reason,
IdempotencyKey: input.IdempotencyKey, IdempotencyKey: input.IdempotencyKey,
Metadata: input.Metadata, Metadata: input.Metadata,
@@ -163,6 +159,8 @@ func (s *Server) rechargeUserWalletBalance(w http.ResponseWriter, r *http.Reques
switch { switch {
case store.IsNotFound(err): case store.IsNotFound(err):
writeError(w, http.StatusNotFound, "user not found") writeError(w, http.StatusNotFound, "user not found")
case errors.Is(err, store.ErrInvalidWalletAmount):
writeError(w, http.StatusBadRequest, "wallet recharge amount must be a positive decimal with at most nine fractional digits", "invalid_wallet_amount")
default: default:
s.logger.Error("recharge user wallet balance failed", "error", err) s.logger.Error("recharge user wallet balance failed", "error", err)
writeError(w, http.StatusInternalServerError, "recharge user wallet balance failed") writeError(w, http.StatusInternalServerError, "recharge user wallet balance failed")
@@ -0,0 +1,145 @@
package httpapi
import (
"crypto/sha256"
"encoding/hex"
"errors"
"net/http"
"strings"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
"github.com/google/uuid"
)
// listBillingSettlements godoc
// @Summary 查询计费结算队列
// @Description 管理端分页查询待结算、重试失败和人工复核记录。
// @Tags billing
// @Produce json
// @Security BearerAuth
// @Param status query string false "结算状态"
// @Param action query string false "动作:settle 或 release"
// @Param page query int false "页码" default(1)
// @Param pageSize query int false "每页数量" default(50)
// @Success 200 {object} BillingSettlementListResponse
// @Failure 400 {object} ErrorEnvelope
// @Failure 401 {object} ErrorEnvelope
// @Failure 403 {object} ErrorEnvelope
// @Failure 500 {object} ErrorEnvelope
// @Router /api/admin/runtime/billing-settlements [get]
func (s *Server) listBillingSettlements(w http.ResponseWriter, r *http.Request) {
query := r.URL.Query()
page, err := positiveQueryInt(query.Get("page"), 1)
if err != nil {
writeError(w, http.StatusBadRequest, "invalid page", "invalid_request")
return
}
pageSize, err := positiveQueryInt(query.Get("pageSize"), 50)
if err != nil {
writeError(w, http.StatusBadRequest, "invalid pageSize", "invalid_request")
return
}
result, err := s.store.ListBillingSettlements(r.Context(), store.BillingSettlementListFilter{
Status: query.Get("status"), Action: query.Get("action"), Page: page, PageSize: pageSize,
})
if err != nil {
s.logger.Error("list billing settlements failed", "error_category", "billing_settlement_list_failed")
writeError(w, http.StatusInternalServerError, "list billing settlements failed", "billing_settlement_list_failed")
return
}
writeJSON(w, http.StatusOK, result)
}
// retryBillingSettlement godoc
// @Summary 重试计费结算
// @Description Manager 使用单值 Idempotency-Key 将重试失败或人工复核记录重新放入队列,并记录审计日志。
// @Tags billing
// @Produce json
// @Security BearerAuth
// @Param settlementId path string true "结算记录 ID"
// @Param Idempotency-Key header string true "幂等键"
// @Success 200 {object} BillingSettlementRetryResponse
// @Failure 400 {object} ErrorEnvelope
// @Failure 401 {object} ErrorEnvelope
// @Failure 403 {object} ErrorEnvelope
// @Failure 404 {object} ErrorEnvelope
// @Failure 409 {object} ErrorEnvelope
// @Failure 500 {object} ErrorEnvelope
// @Router /api/admin/runtime/billing-settlements/{settlementId}/retry [post]
func (s *Server) retryBillingSettlement(w http.ResponseWriter, r *http.Request) {
key, ok := singleIdempotencyKey(r)
if !ok {
writeError(w, http.StatusBadRequest, "a single Idempotency-Key is required", "idempotency_key_required")
return
}
keyHash := sha256.Sum256([]byte(key))
settlementID := strings.TrimSpace(r.PathValue("settlementId"))
if _, err := uuid.Parse(settlementID); err != nil {
writeError(w, http.StatusBadRequest, "invalid billing settlement ID", "invalid_request")
return
}
actor, _ := auth.UserFromContext(r.Context())
var item store.BillingSettlement
var audit store.AuditLog
var replayed bool
err := s.store.InTx(r.Context(), func(tx store.Tx) error {
var err error
item, replayed, err = s.store.RetryBillingSettlementTx(r.Context(), tx, settlementID, hex.EncodeToString(keyHash[:]))
if err != nil || replayed {
return err
}
audit, err = s.store.RecordAuditLogTx(r.Context(), tx, billingSettlementRetryAuditInput(r, actor, item))
return err
})
if err != nil {
switch {
case store.IsNotFound(err):
writeError(w, http.StatusNotFound, "billing settlement not found", "billing_settlement_not_found")
case errors.Is(err, store.ErrBillingSettlementNotRetryable):
writeError(w, http.StatusConflict, "billing settlement is not retryable", "billing_settlement_not_retryable")
default:
s.logger.Error("retry billing settlement failed", "settlementID", settlementID, "error_category", "billing_settlement_retry_failed")
writeError(w, http.StatusInternalServerError, "retry billing settlement failed", "billing_settlement_retry_failed")
}
return
}
if replayed {
w.Header().Set("Idempotent-Replayed", "true")
}
response := map[string]any{"settlement": item}
if !replayed {
response["auditLog"] = audit
}
writeJSON(w, http.StatusOK, response)
}
func singleIdempotencyKey(r *http.Request) (string, bool) {
values := r.Header.Values("Idempotency-Key")
if len(values) != 1 {
return "", false
}
value := strings.TrimSpace(values[0])
if value == "" || len(value) > 255 || strings.Contains(value, ",") {
return "", false
}
return value, true
}
func billingSettlementRetryAuditInput(r *http.Request, actor *auth.User, item store.BillingSettlement) store.AuditLogInput {
input := store.AuditLogInput{
Category: "billing", Action: "billing.settlement.retry",
TargetType: "billing_settlement", TargetID: item.ID,
RequestIP: requestIP(r), UserAgent: r.UserAgent(),
AfterState: map[string]any{"status": item.Status, "attempts": item.Attempts},
Metadata: map[string]any{"taskId": item.TaskID, "action": item.Action, "currency": item.Currency},
}
if actor != nil {
input.ActorGatewayUserID = uuidText(firstNonEmptyText(actor.GatewayUserID, actor.ID))
input.ActorUserID = actor.ID
input.ActorUsername = actor.Username
input.ActorSource = actor.Source
input.ActorRoles = actor.Roles
}
return input
}
@@ -0,0 +1,21 @@
package httpapi
import (
"net/http/httptest"
"testing"
)
func TestSingleIdempotencyKey(t *testing.T) {
t.Parallel()
request := httptest.NewRequest("POST", "/", nil)
request.Header.Add("Idempotency-Key", "retry-1")
if got, ok := singleIdempotencyKey(request); !ok || got != "retry-1" {
t.Fatalf("got key=%q ok=%v", got, ok)
}
request.Header.Add("Idempotency-Key", "retry-2")
if _, ok := singleIdempotencyKey(request); ok {
t.Fatal("multiple Idempotency-Key values must be rejected")
}
}
@@ -444,7 +444,6 @@ VALUES ($1, 5, '{"purpose":"core-flow"}'::jsonb)`, inviteCode); err != nil {
t.Fatalf("unexpected compatible chat response: %+v", compatChat) t.Fatalf("unexpected compatible chat response: %+v", compatChat)
} }
cancelMarker := "cancel-stream-" + suffixText
cancelCtx, cancelRequest := context.WithCancel(context.Background()) cancelCtx, cancelRequest := context.WithCancel(context.Background())
cancelPayload := map[string]any{ cancelPayload := map[string]any{
"model": defaultTextModel, "model": defaultTextModel,
@@ -453,7 +452,6 @@ VALUES ($1, 5, '{"purpose":"core-flow"}'::jsonb)`, inviteCode); err != nil {
"stream": true, "stream": true,
"simulation": true, "simulation": true,
"simulationDurationMs": 250, "simulationDurationMs": 250,
"cancelTestId": cancelMarker,
} }
cancelRaw, err := json.Marshal(cancelPayload) cancelRaw, err := json.Marshal(cancelPayload)
if err != nil { if err != nil {
@@ -466,15 +464,27 @@ VALUES ($1, 5, '{"purpose":"core-flow"}'::jsonb)`, inviteCode); err != nil {
cancelReq.Header.Set("Authorization", "Bearer "+apiKeyResponse.Secret) cancelReq.Header.Set("Authorization", "Bearer "+apiKeyResponse.Secret)
cancelReq.Header.Set("Content-Type", "application/json") cancelReq.Header.Set("Content-Type", "application/json")
cancelErrCh := make(chan error, 1) cancelErrCh := make(chan error, 1)
cancelTaskIDCh := make(chan string, 1)
go func() { go func() {
resp, err := http.DefaultClient.Do(cancelReq) resp, err := http.DefaultClient.Do(cancelReq)
if resp != nil { if resp != nil {
cancelTaskIDCh <- strings.TrimSpace(resp.Header.Get("X-Gateway-Task-Id"))
_, _ = io.ReadAll(resp.Body) _, _ = io.ReadAll(resp.Body)
_ = resp.Body.Close() _ = resp.Body.Close()
} else {
cancelTaskIDCh <- ""
} }
cancelErrCh <- err cancelErrCh <- err
}() }()
cancelTaskID := waitForTaskIDByRequestMarker(t, ctx, testPool, cancelMarker, 2*time.Second) var cancelTaskID string
select {
case cancelTaskID = <-cancelTaskIDCh:
case <-time.After(2 * time.Second):
t.Fatal("cancelled stream did not return response headers")
}
if cancelTaskID == "" {
t.Fatal("cancelled stream response did not expose X-Gateway-Task-Id")
}
cancelRequest() cancelRequest()
select { select {
case <-cancelErrCh: case <-cancelErrCh:
@@ -845,7 +855,7 @@ WHERE gateway_user_id = $1::uuid
} }
doAPIV1ChatCompletionAndLoadTask(t, ctx, testPool, server.URL, apiKeyResponse.Secret, map[string]any{ doAPIV1ChatCompletionAndLoadTask(t, ctx, testPool, server.URL, apiKeyResponse.Secret, map[string]any{
"model": pricingModel, "model": pricingModel,
"runMode": "simulation", "runMode": "production",
"simulation": true, "simulation": true,
"simulationDurationMs": 5, "simulationDurationMs": 5,
"messages": []map[string]any{{"role": "user", "content": "priced ping"}}, "messages": []map[string]any{{"role": "user", "content": "priced ping"}},
@@ -853,6 +863,20 @@ WHERE gateway_user_id = $1::uuid
if pricingTask.Task.Status != "succeeded" || !floatNear(pricingTask.Task.FinalChargeAmount, 0.028) { if pricingTask.Task.Status != "succeeded" || !floatNear(pricingTask.Task.FinalChargeAmount, 0.028) {
t.Fatalf("custom pricing rule set should drive text billing, got task=%+v", pricingTask.Task) t.Fatalf("custom pricing rule set should drive text billing, got task=%+v", pricingTask.Task)
} }
settlementDeadline := time.Now().Add(3 * time.Second)
for {
var billingStatus string
if err := testPool.QueryRow(ctx, `SELECT billing_status FROM gateway_tasks WHERE id = $1::uuid`, pricingTask.Task.ID).Scan(&billingStatus); err != nil {
t.Fatalf("read pricing task billing status: %v", err)
}
if billingStatus == "settled" {
break
}
if time.Now().After(settlementDeadline) {
t.Fatalf("pricing task billing did not settle before deadline, status=%s", billingStatus)
}
time.Sleep(25 * time.Millisecond)
}
var walletBalanceAfter float64 var walletBalanceAfter float64
var walletSpentAfter float64 var walletSpentAfter float64
if err := testPool.QueryRow(ctx, ` if err := testPool.QueryRow(ctx, `
@@ -1484,7 +1508,7 @@ WHERE m.platform_id = $1::uuid
t.Fatalf("workspace task list should include persisted task records, got %+v", workspaceTaskList.Items) t.Fatalf("workspace task list should include persisted task records, got %+v", workspaceTaskList.Items)
} }
req, err := http.NewRequest(http.MethodGet, server.URL+"/api/v1/tasks/"+taskResponse.Task.ID+"/events", nil) req, err := http.NewRequest(http.MethodGet, server.URL+"/api/v1/tasks/"+pricingTask.Task.ID+"/events", nil)
if err != nil { if err != nil {
t.Fatalf("build events request: %v", err) t.Fatalf("build events request: %v", err)
} }
@@ -1864,7 +1888,7 @@ func doJSON(t *testing.T, baseURL string, method string, path string, token stri
doJSONWithHeaders(t, baseURL, method, path, token, payload, nil, expectedStatus, out) doJSONWithHeaders(t, baseURL, method, path, token, payload, nil, expectedStatus, out)
} }
func doJSONWithHeaders(t *testing.T, baseURL string, method string, path string, token string, payload any, headers map[string]string, expectedStatus int, out any) { func doJSONWithHeaders(t *testing.T, baseURL string, method string, path string, token string, payload any, headers map[string]string, expectedStatus int, out any) http.Header {
t.Helper() t.Helper()
var body io.Reader var body io.Reader
if payload != nil { if payload != nil {
@@ -1901,16 +1925,22 @@ func doJSONWithHeaders(t *testing.T, baseURL string, method string, path string,
t.Fatalf("decode %s %s response: %v body=%s", method, path, err, string(raw)) t.Fatalf("decode %s %s response: %v body=%s", method, path, err, string(raw))
} }
} }
return resp.Header.Clone()
} }
func doAPIV1ChatCompletionAndLoadTask(t *testing.T, ctx context.Context, pool *pgxpool.Pool, baseURL string, token string, payload map[string]any, marker string, expectedStatus int, responseOut any, taskDetailOut any) string { func doAPIV1ChatCompletionAndLoadTask(t *testing.T, ctx context.Context, pool *pgxpool.Pool, baseURL string, token string, payload map[string]any, marker string, expectedStatus int, responseOut any, taskDetailOut any) string {
t.Helper() t.Helper()
payload["integrationTestMarker"] = marker _ = ctx
_ = pool
_ = marker
if responseOut == nil { if responseOut == nil {
responseOut = &map[string]any{} responseOut = &map[string]any{}
} }
doJSON(t, baseURL, http.MethodPost, "/api/v1/chat/completions", token, payload, expectedStatus, responseOut) responseHeaders := doJSONWithHeaders(t, baseURL, http.MethodPost, "/api/v1/chat/completions", token, payload, nil, expectedStatus, responseOut)
taskID := waitForTaskIDByRequestField(t, ctx, pool, "integrationTestMarker", marker, 2*time.Second) taskID := strings.TrimSpace(responseHeaders.Get("X-Gateway-Task-Id"))
if taskID == "" {
t.Fatal("chat completion response did not expose X-Gateway-Task-Id")
}
if taskDetailOut != nil { if taskDetailOut != nil {
doJSON(t, baseURL, http.MethodGet, "/api/v1/tasks/"+taskID, token, nil, http.StatusOK, taskDetailOut) doJSON(t, baseURL, http.MethodGet, "/api/v1/tasks/"+taskID, token, nil, http.StatusOK, taskDetailOut)
} }
@@ -2046,11 +2076,6 @@ func waitForTaskStatus(t *testing.T, baseURL string, token string, taskID string
return detail return detail
} }
func waitForTaskIDByRequestMarker(t *testing.T, ctx context.Context, pool *pgxpool.Pool, marker string, timeout time.Duration) string {
t.Helper()
return waitForTaskIDByRequestField(t, ctx, pool, "cancelTestId", marker, timeout)
}
func waitForTaskIDByRequestField(t *testing.T, ctx context.Context, pool *pgxpool.Pool, key string, value string, timeout time.Duration) string { func waitForTaskIDByRequestField(t *testing.T, ctx context.Context, pool *pgxpool.Pool, key string, value string, timeout time.Duration) string {
t.Helper() t.Helper()
deadline := time.Now().Add(timeout) deadline := time.Now().Add(timeout)
+33 -4
View File
@@ -4,6 +4,7 @@ import (
"crypto/rand" "crypto/rand"
"encoding/hex" "encoding/hex"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"io" "io"
"mime" "mime"
@@ -94,6 +95,11 @@ func (s *Server) geminiGenerateContent(w http.ResponseWriter, r *http.Request) {
writeError(w, http.StatusForbidden, "api key scope does not allow this capability") writeError(w, http.StatusForbidden, "api key scope does not allow this capability")
return return
} }
idempotencyKey, hasIdempotencyKey, err := optionalTaskIdempotencyKey(r)
if err != nil {
writeError(w, http.StatusBadRequest, "Idempotency-Key must contain one non-empty value", "invalid_idempotency_key")
return
}
prepared, err := s.prepareTaskRequest(r.Context(), r, user, mapping.Body) prepared, err := s.prepareTaskRequest(r.Context(), r, user, mapping.Body)
if err != nil { if err != nil {
s.logger.Warn("prepare gemini task request failed", "kind", mapping.Kind, "error", err) s.logger.Warn("prepare gemini task request failed", "kind", mapping.Kind, "error", err)
@@ -104,7 +110,7 @@ func (s *Server) geminiGenerateContent(w http.ResponseWriter, r *http.Request) {
writeError(w, status, err.Error(), clients.ErrorCode(err)) writeError(w, status, err.Error(), clients.ErrorCode(err))
return return
} }
task, err := s.store.CreateTask(r.Context(), store.CreateTaskInput{ createInput := store.CreateTaskInput{
Kind: mapping.Kind, Kind: mapping.Kind,
Model: mapping.Model, Model: mapping.Model,
RunMode: runModeFromRequest(prepared.Body), RunMode: runModeFromRequest(prepared.Body),
@@ -113,10 +119,33 @@ func (s *Server) geminiGenerateContent(w http.ResponseWriter, r *http.Request) {
ConversationID: prepared.ConversationID, ConversationID: prepared.ConversationID,
NewMessageCount: prepared.NewMessageCount, NewMessageCount: prepared.NewMessageCount,
MessageRefs: prepared.MessageRefs, MessageRefs: prepared.MessageRefs,
}, user) }
if hasIdempotencyKey {
createInput.IdempotencyKeyHash = taskIdempotencyKeyHash(idempotencyKey)
createInput.IdempotencyRequestHash = taskIdempotencyRequestHash(mapping.Kind, false, false, prepared.Body)
}
created, err := s.store.CreateTaskIdempotent(r.Context(), createInput, user)
if err != nil { if err != nil {
s.logger.Error("create gemini task failed", "kind", mapping.Kind, "error", err) if errors.Is(err, store.ErrIdempotencyKeyReused) {
writeError(w, http.StatusInternalServerError, "create task failed") writeError(w, http.StatusConflict, "Idempotency-Key was reused for a different request", "idempotency_key_reused")
return
}
s.logger.Error("create gemini task failed", "kind", mapping.Kind, "error_category", "task_create_failed")
writeError(w, http.StatusInternalServerError, "create task failed", "task_create_failed")
return
}
task := created.Task
if created.Replayed {
if s.billingMetrics != nil {
s.billingMetrics.ObserveBillingEvent("idempotent_replay")
}
w.Header().Set("Idempotent-Replayed", "true")
w.Header().Set("X-Gateway-Task-Id", task.ID)
if task.Status == "succeeded" {
writeJSON(w, http.StatusOK, geminiGenerateContentResponse(task.Result, mapping.Model))
return
}
writeIdempotentTaskReplay(w, task, true)
return return
} }
runCtx, cancelRun := s.requestExecutionContext(r) runCtx, cancelRun := s.requestExecutionContext(r)
+40 -4
View File
@@ -917,7 +917,13 @@ func (s *Server) estimatePricing(w http.ResponseWriter, r *http.Request) {
} }
estimate, err := s.runner.Estimate(r.Context(), kind, model, body, user) estimate, err := s.runner.Estimate(r.Context(), kind, model, body, user)
if err != nil { if err != nil {
if s.billingMetrics != nil {
s.billingMetrics.ObserveBillingEvent("estimate_failed")
}
if runner.IsPricingUnavailable(err) { if runner.IsPricingUnavailable(err) {
if s.billingMetrics != nil {
s.billingMetrics.ObserveBillingEvent("pricing_unavailable")
}
writeErrorWithDetails(w, http.StatusServiceUnavailable, runErrorMessage(err), runErrorDetails(err), "pricing_unavailable") writeErrorWithDetails(w, http.StatusServiceUnavailable, runErrorMessage(err), runErrorDetails(err), "pricing_unavailable")
return return
} }
@@ -982,6 +988,7 @@ func (s *Server) listModelRateLimitStatuses(w http.ResponseWriter, r *http.Reque
// @Produce json // @Produce json
// @Security BearerAuth // @Security BearerAuth
// @Param X-Async header bool false "true 时异步创建任务并返回 202" // @Param X-Async header bool false "true 时异步创建任务并返回 202"
// @Param Idempotency-Key header string false "可选请求幂等键;同一用户范围内唯一"
// @Param input body TaskRequest true "AI 任务请求,字段随任务类型变化" // @Param input body TaskRequest true "AI 任务请求,字段随任务类型变化"
// @Success 200 {object} CompatibleResponse // @Success 200 {object} CompatibleResponse
// @Success 202 {object} TaskAcceptedResponse // @Success 202 {object} TaskAcceptedResponse
@@ -1050,6 +1057,11 @@ func (s *Server) createTask(kind string, compatible bool) http.Handler {
return return
} }
responsePlan := planTaskResponse(kind, compatible, body, r) responsePlan := planTaskResponse(kind, compatible, body, r)
idempotencyKey, hasIdempotencyKey, err := optionalTaskIdempotencyKey(r)
if err != nil {
writeError(w, http.StatusBadRequest, "Idempotency-Key must contain one non-empty value", "invalid_idempotency_key")
return
}
prepared, err := s.prepareTaskRequest(r.Context(), r, user, body) prepared, err := s.prepareTaskRequest(r.Context(), r, user, body)
if err != nil { if err != nil {
s.logger.Warn("prepare task request failed", "kind", kind, "error", err) s.logger.Warn("prepare task request failed", "kind", kind, "error", err)
@@ -1061,7 +1073,7 @@ func (s *Server) createTask(kind string, compatible bool) http.Handler {
return return
} }
task, err := s.store.CreateTask(r.Context(), store.CreateTaskInput{ createInput := store.CreateTaskInput{
Kind: kind, Kind: kind,
Model: model, Model: model,
RunMode: runModeFromRequest(prepared.Body), RunMode: runModeFromRequest(prepared.Body),
@@ -1070,10 +1082,32 @@ func (s *Server) createTask(kind string, compatible bool) http.Handler {
ConversationID: prepared.ConversationID, ConversationID: prepared.ConversationID,
NewMessageCount: prepared.NewMessageCount, NewMessageCount: prepared.NewMessageCount,
MessageRefs: prepared.MessageRefs, MessageRefs: prepared.MessageRefs,
}, user) }
if hasIdempotencyKey {
createInput.IdempotencyKeyHash = taskIdempotencyKeyHash(idempotencyKey)
createInput.IdempotencyRequestHash = taskIdempotencyRequestHash(kind, responsePlan.asyncMode, responsePlan.streamMode, prepared.Body)
}
created, err := s.store.CreateTaskIdempotent(r.Context(), createInput, user)
if err != nil { if err != nil {
s.logger.Error("create task failed", "kind", kind, "error", err) if errors.Is(err, store.ErrIdempotencyKeyReused) {
writeError(w, http.StatusInternalServerError, "create task failed") writeError(w, http.StatusConflict, "Idempotency-Key was reused for a different request", "idempotency_key_reused")
return
}
s.logger.Error("create task failed", "kind", kind, "error_category", "task_create_failed")
writeError(w, http.StatusInternalServerError, "create task failed", "task_create_failed")
return
}
task := created.Task
if created.Replayed {
if s.billingMetrics != nil {
s.billingMetrics.ObserveBillingEvent("idempotent_replay")
}
w.Header().Set("Idempotent-Replayed", "true")
if responsePlan.streamMode {
writeError(w, http.StatusConflict, "streaming idempotent replay is not supported", "idempotency_stream_replay_unsupported")
return
}
writeIdempotentTaskReplay(w, task, responsePlan.compatibleMode)
return return
} }
if responsePlan.asyncMode { if responsePlan.asyncMode {
@@ -1388,6 +1422,8 @@ func scopeForTaskKind(kind string) string {
func statusFromRunError(err error) int { func statusFromRunError(err error) int {
switch { switch {
case clients.ErrorCode(err) == "billing_hold":
return http.StatusServiceUnavailable
case runner.IsPricingUnavailable(err): case runner.IsPricingUnavailable(err):
return http.StatusServiceUnavailable return http.StatusServiceUnavailable
case clients.ErrorCode(err) == "invalid_previous_response_id" || clients.ErrorCode(err) == "response_chain_too_deep" || clients.ErrorCode(err) == "unsupported_model_protocol" || clients.ErrorCode(err) == "unsupported_response_tool" || clients.ErrorCode(err) == "unsupported_response_parameter": case clients.ErrorCode(err) == "invalid_previous_response_id" || clients.ErrorCode(err) == "response_chain_too_deep" || clients.ErrorCode(err) == "unsupported_model_protocol" || clients.ErrorCode(err) == "unsupported_response_tool" || clients.ErrorCode(err) == "unsupported_response_parameter":
@@ -43,7 +43,7 @@ func TestOIDCJITFullGatewayUserFlowAndSecurityBoundary(t *testing.T) {
if err != nil { if err != nil {
t.Fatalf("connect store: %v", err) t.Fatalf("connect store: %v", err)
} }
defer db.Close() t.Cleanup(db.Close)
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil { if err != nil {
@@ -247,6 +247,9 @@ DELETE FROM gateway_users WHERE source = 'oidc' AND external_user_id = ANY($1::t
t.Fatalf("rejected OIDC tokens created %d local users", rejectedWrites) t.Fatalf("rejected OIDC tokens created %d local users", rejectedWrites)
} }
if _, err := db.DisableActiveIdentityRevision(ctx, activeRevision.Version, "oidc-jit-disabled", "oidc-jit-disabled"); err != nil {
t.Fatalf("disable initial OIDC JIT revision: %v", err)
}
disabledJITRevision := prepareOIDCJITRevision(t, ctx, db, baseConfig, issuer, "default", false) disabledJITRevision := prepareOIDCJITRevision(t, ctx, db, baseConfig, issuer, "default", false)
testRevisionIDs = append(testRevisionIDs, disabledJITRevision.ID) testRevisionIDs = append(testRevisionIDs, disabledJITRevision.ID)
disabledJITRevision, _, err = db.ActivateIdentityRevision(ctx, disabledJITRevision.ID, disabledJITRevision.Version, "oidc-jit-disabled", "oidc-jit-disabled") disabledJITRevision, _, err = db.ActivateIdentityRevision(ctx, disabledJITRevision.ID, disabledJITRevision.Version, "oidc-jit-disabled", "oidc-jit-disabled")
@@ -298,6 +301,9 @@ func prepareOIDCJITRevision(t *testing.T, ctx context.Context, db *store.Store,
if err := secrets.Put(ctx, sessionReference, bytes.Repeat([]byte{8}, 32)); err != nil { if err := secrets.Put(ctx, sessionReference, bytes.Repeat([]byte{8}, 32)); err != nil {
t.Fatalf("store OIDC JIT session key: %v", err) t.Fatalf("store OIDC JIT session key: %v", err)
} }
if err := db.QueueIdentitySecretCleanup(ctx, sessionReference, time.Now().Add(10*time.Minute)); err != nil {
t.Fatalf("stage OIDC JIT session key for adoption: %v", err)
}
draft, err = db.ApplyIdentityManifest(ctx, draft.ID, draft.Version, identity.ManifestApplication{ draft, err = db.ApplyIdentityManifest(ctx, draft.ID, draft.Version, identity.ManifestApplication{
Manifest: identity.ManifestV1{ Manifest: identity.ManifestV1{
SchemaVersion: 1, Issuer: issuer, TenantID: oidcJITTenantID, ApplicationID: uuid.NewString(), SchemaVersion: 1, Issuer: issuer, TenantID: oidcJITTenantID, ApplicationID: uuid.NewString(),
@@ -360,6 +366,10 @@ var currentOIDCTestNonce string
func createOIDCBFFSessionCookie(t *testing.T, baseURL string) *http.Cookie { func createOIDCBFFSessionCookie(t *testing.T, baseURL string) *http.Cookie {
t.Helper() t.Helper()
parsedBaseURL, err := url.Parse(baseURL)
if err != nil {
t.Fatal(err)
}
jar, err := cookiejar.New(nil) jar, err := cookiejar.New(nil)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
@@ -371,6 +381,11 @@ func createOIDCBFFSessionCookie(t *testing.T, baseURL string) *http.Cookie {
if len(via) > 10 { if len(via) > 10 {
return errors.New("too many redirects") return errors.New("too many redirects")
} }
// Follow the synthetic issuer's authorization redirect, then stop before
// the callback redirects the browser to the separately hosted web app.
if request.URL.Host != parsedBaseURL.Host && request.URL.Path != "/authorize" {
return http.ErrUseLastResponse
}
return nil return nil
}} }}
response, err := client.Get(baseURL + "/api/v1/auth/oidc/login?returnTo=%2Fapi%2Fv1%2Fme") response, err := client.Get(baseURL + "/api/v1/auth/oidc/login?returnTo=%2Fapi%2Fv1%2Fme")
@@ -378,11 +393,10 @@ func createOIDCBFFSessionCookie(t *testing.T, baseURL string) *http.Cookie {
t.Fatalf("complete OIDC BFF login: %v", err) t.Fatalf("complete OIDC BFF login: %v", err)
} }
defer response.Body.Close() defer response.Body.Close()
if response.StatusCode != http.StatusOK { if response.StatusCode != http.StatusOK && response.StatusCode != http.StatusSeeOther {
body, _ := io.ReadAll(response.Body) body, _ := io.ReadAll(response.Body)
t.Fatalf("OIDC BFF login status = %d, want 200: %s", response.StatusCode, body) t.Fatalf("OIDC BFF login status = %d, want 200: %s", response.StatusCode, body)
} }
parsedBaseURL, _ := url.Parse(baseURL)
for _, cookie := range jar.Cookies(parsedBaseURL) { for _, cookie := range jar.Cookies(parsedBaseURL) {
if cookie.Name == auth.OIDCSessionCookieName { if cookie.Name == auth.OIDCSessionCookieName {
if strings.Count(cookie.Value, ".") == 2 { if strings.Count(cookie.Value, ".") == 2 {
@@ -114,6 +114,18 @@ type AuditLogListResponse struct {
Items []store.AuditLog `json:"items"` Items []store.AuditLog `json:"items"`
} }
type BillingSettlementListResponse struct {
Items []store.BillingSettlement `json:"items"`
Total int `json:"total" example:"42"`
Page int `json:"page" example:"1"`
PageSize int `json:"pageSize" example:"50"`
}
type BillingSettlementRetryResponse struct {
Settlement store.BillingSettlement `json:"settlement"`
AuditLog store.AuditLog `json:"auditLog"`
}
type WalletTransactionListResponse struct { type WalletTransactionListResponse struct {
Items []store.GatewayWalletTransaction `json:"items"` Items []store.GatewayWalletTransaction `json:"items"`
Total int `json:"total" example:"42"` Total int `json:"total" example:"42"`
+7 -2
View File
@@ -43,6 +43,7 @@ type Server struct {
identityTestRevision identity.Revision identityTestRevision identity.Revision
identityTestCookieSecure bool identityTestCookieSecure bool
identityTestBrowserEnabled bool identityTestBrowserEnabled bool
billingMetrics *ssfreceiver.Metrics
} }
type oidcPublicClient interface { type oidcPublicClient interface {
@@ -66,18 +67,19 @@ func NewServer(cfg config.Config, db *store.Store, logger *slog.Logger) http.Han
} }
func NewServerWithContext(ctx context.Context, cfg config.Config, db *store.Store, logger *slog.Logger) http.Handler { func NewServerWithContext(ctx context.Context, cfg config.Config, db *store.Store, logger *slog.Logger) http.Handler {
securityEventMetrics := &ssfreceiver.Metrics{}
server := &Server{ server := &Server{
ctx: ctx, ctx: ctx,
cfg: cfg, cfg: cfg,
store: db, store: db,
oidcUserResolver: db, oidcUserResolver: db,
auth: auth.New(cfg.JWTSecret, cfg.ServerMainBaseURL, cfg.ServerMainInternalToken), auth: auth.New(cfg.JWTSecret, cfg.ServerMainBaseURL, cfg.ServerMainInternalToken),
runner: runner.New(cfg, db, logger), runner: runner.New(cfg, db, logger, securityEventMetrics),
logger: logger, logger: logger,
billingMetrics: securityEventMetrics,
} }
server.auth.ServerMainInternalKey = cfg.ServerMainInternalKey server.auth.ServerMainInternalKey = cfg.ServerMainInternalKey
server.auth.ServerMainInternalSecret = cfg.ServerMainInternalSecret server.auth.ServerMainInternalSecret = cfg.ServerMainInternalSecret
securityEventMetrics := &ssfreceiver.Metrics{}
secretStore, err := identitySecretStore(cfg) secretStore, err := identitySecretStore(cfg)
if err != nil { if err != nil {
panic("invalid identity SecretStore: " + err.Error()) panic("invalid identity SecretStore: " + err.Error())
@@ -118,6 +120,7 @@ func NewServerWithContext(ctx context.Context, cfg config.Config, db *store.Stor
} }
server.auth.LocalAPIKeyVerifier = db.VerifyLocalAPIKey server.auth.LocalAPIKeyVerifier = db.VerifyLocalAPIKey
server.runner.StartAsyncQueueWorker(ctx) server.runner.StartAsyncQueueWorker(ctx)
server.runner.StartBillingSettlementWorker(ctx)
server.startLocalTempAssetCleanup(ctx) server.startLocalTempAssetCleanup(ctx)
server.startOIDCSessionCleanup(ctx) server.startOIDCSessionCleanup(ctx)
@@ -204,6 +207,8 @@ func NewServerWithContext(ctx context.Context, cfg config.Config, db *store.Stor
mux.Handle("DELETE /api/admin/runtime/policy-sets/{policySetID}", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.deleteRuntimePolicySet))) mux.Handle("DELETE /api/admin/runtime/policy-sets/{policySetID}", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.deleteRuntimePolicySet)))
mux.Handle("GET /api/admin/runtime/runner-policy", server.requireAdmin(auth.PermissionPower, http.HandlerFunc(server.getRunnerPolicy))) mux.Handle("GET /api/admin/runtime/runner-policy", server.requireAdmin(auth.PermissionPower, http.HandlerFunc(server.getRunnerPolicy)))
mux.Handle("PATCH /api/admin/runtime/runner-policy", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.updateRunnerPolicy))) mux.Handle("PATCH /api/admin/runtime/runner-policy", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.updateRunnerPolicy)))
mux.Handle("GET /api/admin/runtime/billing-settlements", server.requireAdmin(auth.PermissionPower, http.HandlerFunc(server.listBillingSettlements)))
mux.Handle("POST /api/admin/runtime/billing-settlements/{settlementId}/retry", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.retryBillingSettlement)))
mux.Handle("POST /api/admin/runtime/model-rate-limits/{platformModelID}/restore", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.restorePlatformModelRuntimeStatus))) mux.Handle("POST /api/admin/runtime/model-rate-limits/{platformModelID}/restore", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.restorePlatformModelRuntimeStatus)))
mux.Handle("GET /api/admin/config/network-proxy", server.requireAdmin(auth.PermissionPower, http.HandlerFunc(server.getNetworkProxyConfig))) mux.Handle("GET /api/admin/config/network-proxy", server.requireAdmin(auth.PermissionPower, http.HandlerFunc(server.getNetworkProxyConfig)))
mux.Handle("GET /api/admin/system/file-storage/settings", server.requireAdmin(auth.PermissionPower, http.HandlerFunc(server.getFileStorageSettings))) mux.Handle("GET /api/admin/system/file-storage/settings", server.requireAdmin(auth.PermissionPower, http.HandlerFunc(server.getFileStorageSettings)))
@@ -0,0 +1,85 @@
package httpapi
import (
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"net/http"
"strings"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
)
var errInvalidTaskIdempotencyKey = errors.New("invalid Idempotency-Key header")
func optionalTaskIdempotencyKey(r *http.Request) (string, bool, error) {
values := r.Header.Values("Idempotency-Key")
if len(values) == 0 {
return "", false, nil
}
if len(values) != 1 {
return "", false, errInvalidTaskIdempotencyKey
}
value := strings.TrimSpace(values[0])
if value == "" || len(value) > 255 || strings.Contains(value, ",") {
return "", false, errInvalidTaskIdempotencyKey
}
return value, true, nil
}
func taskIdempotencyKeyHash(key string) string {
digest := sha256.Sum256([]byte(key))
return hex.EncodeToString(digest[:])
}
func taskIdempotencyRequestHash(kind string, async bool, stream bool, body map[string]any) string {
payload, _ := json.Marshal(map[string]any{
"kind": kind, "async": async, "stream": stream, "request": body,
})
digest := sha256.Sum256(payload)
return hex.EncodeToString(digest[:])
}
func writeIdempotentTaskReplay(w http.ResponseWriter, task store.GatewayTask, compatible bool) {
w.Header().Set("Idempotent-Replayed", "true")
w.Header().Set("X-Gateway-Task-Id", task.ID)
if !compatible || task.AsyncMode || (task.Status != "succeeded" && task.Status != "failed" && task.Status != "cancelled") {
writeTaskAccepted(w, task)
return
}
if task.Status == "succeeded" {
writeJSON(w, http.StatusOK, task.Result)
return
}
status := storedTaskErrorStatus(task.ErrorCode)
message := strings.TrimSpace(task.ErrorMessage)
if message == "" {
message = strings.TrimSpace(task.Error)
}
if message == "" {
message = "task failed"
}
code := strings.TrimSpace(task.ErrorCode)
if code == "" {
code = "task_failed"
}
writeError(w, status, message, code)
}
func storedTaskErrorStatus(code string) int {
switch strings.TrimSpace(code) {
case "pricing_unavailable", "response_chain_unavailable", "billing_hold":
return http.StatusServiceUnavailable
case "insufficient_balance":
return http.StatusPaymentRequired
case "bad_request", "invalid_parameter", "invalid_previous_response_id", "unsupported_operation":
return http.StatusBadRequest
case "no_model_candidate", "cloned_voice_not_found":
return http.StatusNotFound
case "rate_limit", "platform_cooling_down", "model_cooling_down":
return http.StatusTooManyRequests
default:
return http.StatusBadGateway
}
}
@@ -0,0 +1,31 @@
package httpapi
import (
"net/http/httptest"
"testing"
)
func TestOptionalTaskIdempotencyKeyRejectsMultipleValues(t *testing.T) {
t.Parallel()
request := httptest.NewRequest("POST", "/", nil)
if _, present, err := optionalTaskIdempotencyKey(request); err != nil || present {
t.Fatalf("missing key present=%v err=%v", present, err)
}
request.Header.Add("Idempotency-Key", "one")
request.Header.Add("Idempotency-Key", "two")
if _, _, err := optionalTaskIdempotencyKey(request); err == nil {
t.Fatal("multiple keys must be rejected")
}
}
func TestTaskIdempotencyRequestHashIsCanonical(t *testing.T) {
t.Parallel()
first := map[string]any{"model": "m", "n": float64(1), "nested": map[string]any{"b": true, "a": "x"}}
second := map[string]any{"nested": map[string]any{"a": "x", "b": true}, "n": float64(1), "model": "m"}
if taskIdempotencyRequestHash("images.generations", false, false, first) != taskIdempotencyRequestHash("images.generations", false, false, second) {
t.Fatal("equivalent JSON objects must have the same request hash")
}
if taskIdempotencyRequestHash("images.generations", true, false, first) == taskIdempotencyRequestHash("images.generations", false, false, first) {
t.Fatal("async response semantics must be part of the request hash")
}
}
@@ -0,0 +1,27 @@
package runner
import "testing"
func TestNormalizedBillingEngineMode(t *testing.T) {
t.Parallel()
if got := normalizedBillingEngineMode(""); got != "observe" {
t.Fatalf("empty mode = %q", got)
}
if got := normalizedBillingEngineMode("ENFORCE"); got != "enforce" {
t.Fatalf("enforce mode = %q", got)
}
if got := normalizedBillingEngineMode("hold"); got != "hold" {
t.Fatalf("hold mode = %q", got)
}
}
func TestBillingItemsFixedTotalKeepsNineDecimalPlaces(t *testing.T) {
t.Parallel()
items := []any{
map[string]any{"amount": "0.000000001"},
map[string]any{"amount": float64(0.000000002)},
}
if got := billingItemsFixedTotal(items).String(); got != "0.000000003" {
t.Fatalf("total = %s", got)
}
}
@@ -0,0 +1,96 @@
package runner
import (
"context"
"errors"
"time"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
"github.com/google/uuid"
)
const billingSettlementPollInterval = time.Second
func billingSettlementRetryDelay(attempt int) time.Duration {
if attempt < 1 {
attempt = 1
}
delay := time.Second
for current := 1; current < attempt && delay < 15*time.Minute; current++ {
delay *= 2
}
if delay > 15*time.Minute {
return 15 * time.Minute
}
return delay
}
func billingSettlementErrorCode(err error) string {
if errors.Is(err, store.ErrInsufficientWalletBalance) {
return "insufficient_balance"
}
return "settlement_failed"
}
func billingSettlementErrorMessage(code string) string {
if code == "insufficient_balance" {
return "wallet balance is insufficient for settlement"
}
return "billing settlement processing failed"
}
func (s *Service) StartBillingSettlementWorker(ctx context.Context) {
workerID := "billing-" + uuid.NewString()
go s.runBillingSettlementWorker(ctx, workerID)
}
func (s *Service) runBillingSettlementWorker(ctx context.Context, workerID string) {
ticker := time.NewTicker(billingSettlementPollInterval)
defer ticker.Stop()
for {
s.processBillingSettlementBatch(ctx, workerID)
select {
case <-ctx.Done():
return
case <-ticker.C:
}
}
}
func (s *Service) processBillingSettlementBatch(ctx context.Context, workerID string) {
items, err := s.store.ClaimBillingSettlements(ctx, workerID, store.BillingSettlementBatchSize, store.BillingSettlementLockTimeout)
if err != nil {
if ctx.Err() == nil {
s.logger.Error("claim billing settlements failed", "error_category", "billing_settlement_claim_failed")
}
return
}
for _, item := range items {
if ctx.Err() != nil {
return
}
if err := s.store.ProcessBillingSettlement(ctx, item); err != nil {
code := billingSettlementErrorCode(err)
markErr := s.store.MarkBillingSettlementFailed(
context.WithoutCancel(ctx),
item,
code,
billingSettlementErrorMessage(code),
billingSettlementRetryDelay(item.Attempts),
)
if markErr != nil {
s.logger.Error("mark billing settlement failed", "settlementID", item.ID, "taskID", item.TaskID, "error_category", "billing_settlement_state_failed")
continue
}
s.observeBillingEvent("settlement_retry")
if item.Attempts >= store.BillingSettlementMaxAttempts {
s.observeBillingEvent("manual_review")
}
s.logger.Warn("billing settlement scheduled for retry", "settlementID", item.ID, "taskID", item.TaskID, "action", item.Action, "error_category", code, "attempt", item.Attempts)
continue
}
s.observeBillingEvent("settlement_completed")
s.logger.Debug("billing settlement completed", "settlementID", item.ID, "taskID", item.TaskID, "action", item.Action)
}
}
@@ -0,0 +1,40 @@
package runner
import (
"errors"
"testing"
"time"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
)
func TestBillingSettlementRetryDelay(t *testing.T) {
t.Parallel()
tests := []struct {
attempt int
want time.Duration
}{
{attempt: 1, want: time.Second},
{attempt: 2, want: 2 * time.Second},
{attempt: 10, want: 512 * time.Second},
{attempt: 11, want: 15 * time.Minute},
{attempt: 20, want: 15 * time.Minute},
}
for _, test := range tests {
if got := billingSettlementRetryDelay(test.attempt); got != test.want {
t.Fatalf("attempt %d: got %s, want %s", test.attempt, got, test.want)
}
}
}
func TestBillingSettlementErrorCode(t *testing.T) {
t.Parallel()
if got := billingSettlementErrorCode(store.ErrInsufficientWalletBalance); got != "insufficient_balance" {
t.Fatalf("got %q", got)
}
if got := billingSettlementErrorCode(errors.New("boom")); got != "settlement_failed" {
t.Fatalf("got %q", got)
}
}
@@ -0,0 +1,28 @@
package runner
import (
"context"
"time"
)
const (
taskExecutionLeaseTTL = 5 * time.Minute
taskExecutionRenewInterval = 30 * time.Second
)
func (s *Service) renewTaskExecutionLease(ctx context.Context, cancel context.CancelFunc, taskID string, executionToken string) {
ticker := time.NewTicker(taskExecutionRenewInterval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
if err := s.store.RenewTaskExecutionLease(ctx, taskID, executionToken, taskExecutionLeaseTTL); err != nil {
s.logger.Warn("task execution lease lost", "taskID", taskID, "error_category", "task_execution_lease_lost")
cancel()
return
}
}
}
}
+1 -6
View File
@@ -37,19 +37,14 @@ func (s *Service) Estimate(ctx context.Context, kind string, model string, body
return EstimateResult{}, err return EstimateResult{}, err
} }
estimates := make([]candidateEstimate, 0, len(candidates)) estimates := make([]candidateEstimate, 0, len(candidates))
var pricingErr error
for _, candidate := range candidates { for _, candidate := range candidates {
candidateBody := preprocessRequest(kind, cloneMap(body), candidate) candidateBody := preprocessRequest(kind, cloneMap(body), candidate)
estimate, candidateErr := s.estimateCandidateV2(ctx, user, kind, candidateBody, candidate) estimate, candidateErr := s.estimateCandidateV2(ctx, user, kind, candidateBody, candidate)
if candidateErr != nil { if candidateErr != nil {
pricingErr = candidateErr return EstimateResult{}, candidateErr
continue
} }
estimates = append(estimates, estimate) estimates = append(estimates, estimate)
} }
if len(estimates) == 0 && pricingErr != nil {
return EstimateResult{}, pricingErr
}
return buildEstimateResult(estimates, pricingRequestFingerprint(kind, model, body)) return buildEstimateResult(estimates, pricingRequestFingerprint(kind, model, body))
} }
+19 -3
View File
@@ -10,6 +10,7 @@ import (
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth" "github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store" "github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
"github.com/google/uuid"
"github.com/riverqueue/river" "github.com/riverqueue/river"
"github.com/riverqueue/river/riverdriver/riverpgxv5" "github.com/riverqueue/river/riverdriver/riverpgxv5"
"github.com/riverqueue/river/rivermigrate" "github.com/riverqueue/river/rivermigrate"
@@ -38,16 +39,22 @@ func (w *asyncTaskWorker) Work(ctx context.Context, job *river.Job[asyncTaskArgs
if task.Status == "succeeded" || task.Status == "failed" || task.Status == "cancelled" { if task.Status == "succeeded" || task.Status == "failed" || task.Status == "cancelled" {
return nil return nil
} }
result, runErr := w.service.Execute(ctx, task, authUserFromTask(task)) executionToken := uuid.NewString()
result, runErr := w.service.executeWithToken(ctx, task, authUserFromTask(task), nil, executionToken)
if runErr == nil { if runErr == nil {
w.service.logger.Debug("river async task completed", "taskID", task.ID, "status", result.Task.Status, "riverJobID", job.ID) w.service.logger.Debug("river async task completed", "taskID", task.ID, "status", result.Task.Status, "riverJobID", job.ID)
return nil return nil
} }
if errors.Is(runErr, store.ErrTaskExecutionLeaseUnavailable) {
w.service.logger.Debug("river async task execution lease already held", "taskID", task.ID, "riverJobID", job.ID)
return nil
}
var queuedErr *TaskQueuedError var queuedErr *TaskQueuedError
if errors.As(runErr, &queuedErr) { if errors.As(runErr, &queuedErr) {
return river.JobSnooze(queuedErr.Delay) return river.JobSnooze(queuedErr.Delay)
} }
if ctx.Err() != nil { if ctx.Err() != nil {
task.ExecutionToken = executionToken
queued, queueErr := w.service.requeueInterruptedAsyncTask(context.WithoutCancel(ctx), task) queued, queueErr := w.service.requeueInterruptedAsyncTask(context.WithoutCancel(ctx), task)
if queueErr != nil { if queueErr != nil {
return queueErr return queueErr
@@ -145,8 +152,7 @@ func (s *Service) recoverAsyncRiverJobs(ctx context.Context) error {
return err return err
} }
for _, item := range items { for _, item := range items {
task := store.GatewayTask{ID: item.ID} result, err := s.riverClient.Insert(ctx, asyncTaskArgs{TaskID: item.ID}, asyncTaskRecoveryInsertOpts(item))
result, err := s.riverClient.Insert(ctx, asyncTaskArgs{TaskID: item.ID}, asyncTaskInsertOpts(task))
if err != nil { if err != nil {
return err return err
} }
@@ -186,6 +192,16 @@ func asyncTaskInsertOpts(task store.GatewayTask) *river.InsertOpts {
} }
} }
func asyncTaskRecoveryInsertOpts(item store.AsyncTaskQueueItem) *river.InsertOpts {
opts := asyncTaskInsertOpts(store.GatewayTask{ID: item.ID})
opts.ScheduledAt = item.NextRunAt
// A replacement process must not be blocked by a River row that the dead
// process left in running state. PostgreSQL execution leases still ensure
// that only one recovery job can call the upstream provider.
opts.UniqueOpts = river.UniqueOpts{}
return opts
}
func authUserFromTask(task store.GatewayTask) *auth.User { func authUserFromTask(task store.GatewayTask) *auth.User {
roles := []string{"user"} roles := []string{"user"}
if strings.TrimSpace(task.UserID) == "" { if strings.TrimSpace(task.UserID) == "" {
+260 -95
View File
@@ -14,6 +14,7 @@ import (
"github.com/easyai/easyai-ai-gateway/apps/api/internal/config" "github.com/easyai/easyai-ai-gateway/apps/api/internal/config"
scriptengine "github.com/easyai/easyai-ai-gateway/apps/api/internal/script" scriptengine "github.com/easyai/easyai-ai-gateway/apps/api/internal/script"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store" "github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
"github.com/google/uuid"
"github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5"
"github.com/riverqueue/river" "github.com/riverqueue/river"
) )
@@ -26,6 +27,11 @@ type Service struct {
scriptExecutor *scriptengine.Executor scriptExecutor *scriptengine.Executor
httpClients *httpClientCache httpClients *httpClientCache
riverClient *river.Client[pgx.Tx] riverClient *river.Client[pgx.Tx]
billingMetrics billingMetricsObserver
}
type billingMetricsObserver interface {
ObserveBillingEvent(string)
} }
type Result struct { type Result struct {
@@ -39,6 +45,19 @@ type TaskQueuedError struct {
Delay time.Duration Delay time.Duration
} }
type upstreamSubmissionUnknownError struct {
AttemptID string
Cause error
}
func (e *upstreamSubmissionUnknownError) Error() string {
return "upstream submission result is unknown"
}
func (e *upstreamSubmissionUnknownError) Unwrap() error {
return e.Cause
}
func (e *TaskQueuedError) Error() string { func (e *TaskQueuedError) Error() string {
return ErrTaskQueued.Error() return ErrTaskQueued.Error()
} }
@@ -47,10 +66,10 @@ func (e *TaskQueuedError) Is(target error) bool {
return target == ErrTaskQueued return target == ErrTaskQueued
} }
func New(cfg config.Config, db *store.Store, logger *slog.Logger) *Service { func New(cfg config.Config, db *store.Store, logger *slog.Logger, observers ...billingMetricsObserver) *Service {
httpClients := newHTTPClientCache() httpClients := newHTTPClientCache()
scriptExecutor := &scriptengine.Executor{Logger: logger} scriptExecutor := &scriptengine.Executor{Logger: logger}
return &Service{ service := &Service{
cfg: cfg, cfg: cfg,
store: db, store: db,
logger: logger, logger: logger,
@@ -76,6 +95,16 @@ func New(cfg config.Config, db *store.Store, logger *slog.Logger) *Service {
}, },
httpClients: httpClients, httpClients: httpClients,
} }
if len(observers) > 0 {
service.billingMetrics = observers[0]
}
return service
}
func (s *Service) observeBillingEvent(event string) {
if s.billingMetrics != nil {
s.billingMetrics.ObserveBillingEvent(event)
}
} }
func (s *Service) Execute(ctx context.Context, task store.GatewayTask, user *auth.User) (Result, error) { func (s *Service) Execute(ctx context.Context, task store.GatewayTask, user *auth.User) (Result, error) {
@@ -87,6 +116,20 @@ func (s *Service) ExecuteStream(ctx context.Context, task store.GatewayTask, use
} }
func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *auth.User, onDelta clients.StreamDelta) (Result, error) { func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *auth.User, onDelta clients.StreamDelta) (Result, error) {
return s.executeWithToken(ctx, task, user, onDelta, uuid.NewString())
}
func (s *Service) executeWithToken(ctx context.Context, task store.GatewayTask, user *auth.User, onDelta clients.StreamDelta, executionToken string) (Result, error) {
wasRunning := task.Status == "running"
claimed, err := s.store.ClaimTaskExecution(ctx, task.ID, executionToken, taskExecutionLeaseTTL)
if err != nil {
return Result{}, err
}
task = claimed
executionCtx, stopExecution := context.WithCancel(ctx)
defer stopExecution()
go s.renewTaskExecutionLease(executionCtx, stopExecution, task.ID, task.ExecutionToken)
ctx = executionCtx
executeStartedAt := time.Now() executeStartedAt := time.Now()
restoredRequest, err := s.restoreTaskRequestReferences(ctx, task) restoredRequest, err := s.restoreTaskRequestReferences(ctx, task)
if err != nil { if err != nil {
@@ -95,10 +138,10 @@ func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *aut
body := normalizeRequest(task.Kind, restoredRequest) body := normalizeRequest(task.Kind, restoredRequest)
responseExecution := responseExecutionContext{} responseExecution := responseExecutionContext{}
modelType := modelTypeFromKind(task.Kind, body) modelType := modelTypeFromKind(task.Kind, body)
if err := s.store.MarkTaskRunning(ctx, task.ID, modelType, s.slimTaskRequestSnapshot(task, body)); err != nil { if err := s.store.MarkTaskRunning(ctx, task.ID, task.ExecutionToken, modelType, s.slimTaskRequestSnapshot(task, body)); err != nil {
return Result{}, err return Result{}, err
} }
if task.Status != "running" { if !wasRunning {
if err := s.emit(ctx, task.ID, "task.running", "running", "starting", 0.12, "task pulled from queue and started", map[string]any{"modelType": modelType}, task.RunMode == "simulation"); err != nil { if err := s.emit(ctx, task.ID, "task.running", "running", "starting", 0.12, "task pulled from queue and started", map[string]any{"modelType": modelType}, task.RunMode == "simulation"); err != nil {
return Result{}, err return Result{}, err
} }
@@ -115,7 +158,7 @@ func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *aut
Reason: "request_validation_failed", Reason: "request_validation_failed",
ModelType: modelType, ModelType: modelType,
}) })
failed, finishErr := s.failTask(ctx, task.ID, "bad_request", err.Error(), task.RunMode == "simulation", err) failed, finishErr := s.failTask(ctx, task.ID, task.ExecutionToken, "bad_request", err.Error(), task.RunMode == "simulation", err)
if finishErr != nil { if finishErr != nil {
return Result{}, finishErr return Result{}, finishErr
} }
@@ -135,14 +178,14 @@ func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *aut
Reason: "cloned_voice_binding_failed", Reason: "cloned_voice_binding_failed",
ModelType: modelType, ModelType: modelType,
}) })
failed, finishErr := s.failTask(ctx, task.ID, clients.ErrorCode(err), err.Error(), task.RunMode == "simulation", err) failed, finishErr := s.failTask(ctx, task.ID, task.ExecutionToken, clients.ErrorCode(err), err.Error(), task.RunMode == "simulation", err)
if finishErr != nil { if finishErr != nil {
return Result{}, finishErr return Result{}, finishErr
} }
return Result{Task: failed, Output: failed.Result}, err return Result{Task: failed, Output: failed.Result}, err
} }
if clonedVoice.Found { if clonedVoice.Found {
if err := s.store.MarkTaskRunning(ctx, task.ID, modelType, s.slimTaskRequestSnapshot(task, body)); err != nil { if err := s.store.MarkTaskRunning(ctx, task.ID, task.ExecutionToken, modelType, s.slimTaskRequestSnapshot(task, body)); err != nil {
return Result{}, err return Result{}, err
} }
} }
@@ -151,7 +194,7 @@ func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *aut
if err != nil { if err != nil {
code, message := responseExecutionFailure(err) code, message := responseExecutionFailure(err)
s.recordFailedAttempt(ctx, failedAttemptRecord{Task: task, Body: body, AttemptNo: task.AttemptCount + 1, Code: code, Cause: err, Simulated: task.RunMode == "simulation", Scope: "response_chain", Reason: code, ModelType: modelType}) s.recordFailedAttempt(ctx, failedAttemptRecord{Task: task, Body: body, AttemptNo: task.AttemptCount + 1, Code: code, Cause: err, Simulated: task.RunMode == "simulation", Scope: "response_chain", Reason: code, ModelType: modelType})
failed, finishErr := s.failTask(ctx, task.ID, code, message, task.RunMode == "simulation", err) failed, finishErr := s.failTask(ctx, task.ID, task.ExecutionToken, code, message, task.RunMode == "simulation", err)
if finishErr != nil { if finishErr != nil {
return Result{}, finishErr return Result{}, finishErr
} }
@@ -183,7 +226,7 @@ func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *aut
Reason: "candidate_selection_failed", Reason: "candidate_selection_failed",
ModelType: modelType, ModelType: modelType,
}) })
failed, finishErr := s.failTask(ctx, task.ID, store.ModelCandidateErrorCode(err), err.Error(), task.RunMode == "simulation", err) failed, finishErr := s.failTask(ctx, task.ID, task.ExecutionToken, store.ModelCandidateErrorCode(err), err.Error(), task.RunMode == "simulation", err)
if finishErr != nil { if finishErr != nil {
return Result{}, finishErr return Result{}, finishErr
} }
@@ -202,7 +245,7 @@ func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *aut
Reason: store.ModelCandidateErrorCode(err), Reason: store.ModelCandidateErrorCode(err),
ModelType: modelType, ModelType: modelType,
}) })
failed, finishErr := s.failTask(ctx, task.ID, store.ModelCandidateErrorCode(err), err.Error(), task.RunMode == "simulation", err) failed, finishErr := s.failTask(ctx, task.ID, task.ExecutionToken, store.ModelCandidateErrorCode(err), err.Error(), task.RunMode == "simulation", err)
if finishErr != nil { if finishErr != nil {
return Result{}, finishErr return Result{}, finishErr
} }
@@ -224,7 +267,7 @@ func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *aut
ExtraMetrics: []map[string]any{candidateFilterMetrics}, ExtraMetrics: []map[string]any{candidateFilterMetrics},
ModelType: modelType, ModelType: modelType,
}) })
failed, finishErr := s.failTask(ctx, task.ID, store.ModelCandidateErrorCode(err), err.Error(), task.RunMode == "simulation", err, candidateFilterMetrics) failed, finishErr := s.failTask(ctx, task.ID, task.ExecutionToken, store.ModelCandidateErrorCode(err), err.Error(), task.RunMode == "simulation", err, candidateFilterMetrics)
if finishErr != nil { if finishErr != nil {
return Result{}, finishErr return Result{}, finishErr
} }
@@ -240,7 +283,7 @@ func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *aut
Simulated: task.RunMode == "simulation", Scope: "candidate_output_token_filter", Reason: store.ModelCandidateErrorCode(err), Simulated: task.RunMode == "simulation", Scope: "candidate_output_token_filter", Reason: store.ModelCandidateErrorCode(err),
ExtraMetrics: []map[string]any{candidateFilterMetrics}, ModelType: modelType, ExtraMetrics: []map[string]any{candidateFilterMetrics}, ModelType: modelType,
}) })
failed, finishErr := s.failTask(ctx, task.ID, store.ModelCandidateErrorCode(err), err.Error(), task.RunMode == "simulation", err, candidateFilterMetrics) failed, finishErr := s.failTask(ctx, task.ID, task.ExecutionToken, store.ModelCandidateErrorCode(err), err.Error(), task.RunMode == "simulation", err, candidateFilterMetrics)
if finishErr != nil { if finishErr != nil {
return Result{}, finishErr return Result{}, finishErr
} }
@@ -251,7 +294,7 @@ func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *aut
if err != nil { if err != nil {
code, message := responseExecutionFailure(err) code, message := responseExecutionFailure(err)
s.recordFailedAttempt(ctx, failedAttemptRecord{Task: task, Body: body, AttemptNo: task.AttemptCount + 1, Code: code, Cause: err, Simulated: task.RunMode == "simulation", Scope: "response_protocol", Reason: code, ModelType: modelType}) s.recordFailedAttempt(ctx, failedAttemptRecord{Task: task, Body: body, AttemptNo: task.AttemptCount + 1, Code: code, Cause: err, Simulated: task.RunMode == "simulation", Scope: "response_protocol", Reason: code, ModelType: modelType})
failed, finishErr := s.failTask(ctx, task.ID, code, message, task.RunMode == "simulation", err) failed, finishErr := s.failTask(ctx, task.ID, task.ExecutionToken, code, message, task.RunMode == "simulation", err)
if finishErr != nil { if finishErr != nil {
return Result{}, finishErr return Result{}, finishErr
} }
@@ -260,8 +303,17 @@ func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *aut
} }
pricingByCandidate := map[string]resolvedPricing{} pricingByCandidate := map[string]resolvedPricing{}
reservationBillings := []any(nil) reservationBillings := []any(nil)
reservationPricingSnapshot := map[string]any(nil)
if task.RunMode == "production" { if task.RunMode == "production" {
pricedCandidates := make([]store.RuntimeModelCandidate, 0, len(candidates)) billingMode := normalizedBillingEngineMode(s.cfg.BillingEngineMode)
if billingMode == "hold" {
holdErr := &clients.ClientError{Code: "billing_hold", Message: "production billing is temporarily on hold", Retryable: true}
failed, finishErr := s.failTask(ctx, task.ID, task.ExecutionToken, "billing_hold", holdErr.Error(), false, holdErr)
if finishErr != nil {
return Result{}, finishErr
}
return Result{Task: failed, Output: failed.Result}, holdErr
}
estimates := make([]candidateEstimate, 0, len(candidates)) estimates := make([]candidateEstimate, 0, len(candidates))
var pricingErr error var pricingErr error
for _, candidate := range candidates { for _, candidate := range candidates {
@@ -269,32 +321,68 @@ func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *aut
estimate, estimateErr := s.estimateCandidateV2(ctx, user, task.Kind, pricingBody, candidate) estimate, estimateErr := s.estimateCandidateV2(ctx, user, task.Kind, pricingBody, candidate)
if estimateErr != nil { if estimateErr != nil {
pricingErr = estimateErr pricingErr = estimateErr
if billingMode == "enforce" {
break
}
continue continue
} }
pricedCandidates = append(pricedCandidates, candidate)
estimates = append(estimates, estimate) estimates = append(estimates, estimate)
pricingByCandidate[pricingCandidateKey(candidate)] = estimate.Pricing pricingByCandidate[pricingCandidateKey(candidate)] = estimate.Pricing
} }
if len(pricedCandidates) == 0 { if billingMode == "observe" {
if pricingErr == nil { legacyItems, legacyAmount := s.maximumLegacyCandidateEstimate(ctx, user, task.Kind, body, candidates)
pricingErr = &PricingUnavailableError{Reason: "no candidate has effective pricing"} reservationBillings = legacyItems
candidateSnapshots := make([]any, 0, len(estimates))
for _, estimate := range estimates {
candidateSnapshots = append(candidateSnapshots, estimate.Snapshot)
} }
failed, finishErr := s.failTask(ctx, task.ID, "pricing_unavailable", pricingErr.Error(), false, pricingErr) observedAmount := ""
if maximumEstimate, estimateErr := maximumCandidateEstimate(estimates); estimateErr == nil {
observedAmount = maximumEstimate.Amount.String()
}
reservationPricingSnapshot = map[string]any{
"pricingVersion": "legacy-observe", "observedPricingVersion": pricingVersionV2,
"requestFingerprint": pricingRequestFingerprint(task.Kind, task.Model, body),
"reservationAmount": legacyAmount.String(), "observedReservationAmount": observedAmount,
"candidateCount": len(candidates), "observedCandidates": candidateSnapshots,
}
s.logger.Info("billing observe comparison", "taskID", task.ID, "legacyAmount", legacyAmount.String(), "v2Amount", observedAmount, "pricedCandidates", len(estimates), "candidateCount", len(candidates))
} else if pricingErr != nil || len(estimates) != len(candidates) {
if pricingErr == nil {
pricingErr = &PricingUnavailableError{Reason: "not every candidate has effective pricing"}
}
s.observeBillingEvent("pricing_unavailable")
failed, finishErr := s.failTask(ctx, task.ID, task.ExecutionToken, "pricing_unavailable", pricingErr.Error(), false, pricingErr)
if finishErr != nil { if finishErr != nil {
return Result{}, finishErr return Result{}, finishErr
} }
return Result{Task: failed, Output: failed.Result}, pricingErr return Result{Task: failed, Output: failed.Result}, pricingErr
} } else {
maximumEstimate, estimateErr := maximumCandidateEstimate(estimates) maximumEstimate, estimateErr := maximumCandidateEstimate(estimates)
if estimateErr != nil { if estimateErr != nil {
failed, finishErr := s.failTask(ctx, task.ID, "pricing_unavailable", estimateErr.Error(), false, estimateErr) failed, finishErr := s.failTask(ctx, task.ID, task.ExecutionToken, "pricing_unavailable", estimateErr.Error(), false, estimateErr)
if finishErr != nil { if finishErr != nil {
return Result{}, finishErr return Result{}, finishErr
}
return Result{Task: failed, Output: failed.Result}, estimateErr
}
pricedCandidates := make([]store.RuntimeModelCandidate, 0, len(estimates))
candidateSnapshots := make([]any, 0, len(estimates))
for _, candidate := range candidates {
if _, ok := pricingByCandidate[pricingCandidateKey(candidate)]; ok {
pricedCandidates = append(pricedCandidates, candidate)
}
}
for _, estimate := range estimates {
candidateSnapshots = append(candidateSnapshots, estimate.Snapshot)
}
candidates = pricedCandidates
reservationBillings = maximumEstimate.Items
reservationPricingSnapshot = map[string]any{
"pricingVersion": pricingVersionV2, "requestFingerprint": pricingRequestFingerprint(task.Kind, task.Model, body),
"reservationAmount": maximumEstimate.Amount.String(), "candidateCount": len(estimates), "candidates": candidateSnapshots,
} }
return Result{Task: failed, Output: failed.Result}, estimateErr
} }
candidates = pricedCandidates
reservationBillings = maximumEstimate.Items
} }
firstCandidateBody := body firstCandidateBody := body
normalizedModelType := modelType normalizedModelType := modelType
@@ -328,17 +416,17 @@ func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *aut
Preprocessing: &firstPreprocessing, Preprocessing: &firstPreprocessing,
ModelType: normalizedModelType, ModelType: normalizedModelType,
}) })
failed, finishErr := s.failTask(ctx, task.ID, clients.ErrorCode(clientErr), clientErr.Error(), task.RunMode == "simulation", clientErr, parameterPreprocessingMetrics(firstPreprocessing)) failed, finishErr := s.failTask(ctx, task.ID, task.ExecutionToken, clients.ErrorCode(clientErr), clientErr.Error(), task.RunMode == "simulation", clientErr, parameterPreprocessingMetrics(firstPreprocessing))
if finishErr != nil { if finishErr != nil {
return Result{}, finishErr return Result{}, finishErr
} }
return Result{Task: failed, Output: failed.Result}, clientErr return Result{Task: failed, Output: failed.Result}, clientErr
} }
if err := s.store.MarkTaskRunning(ctx, task.ID, candidates[0].ModelType, s.slimTaskRequestSnapshot(task, firstCandidateBody)); err != nil { if err := s.store.MarkTaskRunning(ctx, task.ID, task.ExecutionToken, candidates[0].ModelType, s.slimTaskRequestSnapshot(task, firstCandidateBody)); err != nil {
return Result{}, err return Result{}, err
} }
var reserveErr error var reserveErr error
walletReservations, reserveErr = s.store.ReserveTaskBilling(ctx, task, user, reservationBillings) walletReservations, reserveErr = s.store.ReserveTaskBilling(ctx, task, user, reservationBillings, reservationPricingSnapshot)
if reserveErr != nil { if reserveErr != nil {
if errors.Is(reserveErr, store.ErrInsufficientWalletBalance) { if errors.Is(reserveErr, store.ErrInsufficientWalletBalance) {
attemptNo = s.recordFailedAttempt(ctx, failedAttemptRecord{ attemptNo = s.recordFailedAttempt(ctx, failedAttemptRecord{
@@ -355,7 +443,7 @@ func (s *Service) execute(ctx context.Context, task store.GatewayTask, user *aut
Preprocessing: &firstPreprocessing, Preprocessing: &firstPreprocessing,
ModelType: normalizedModelType, ModelType: normalizedModelType,
}) })
failed, finishErr := s.failTask(ctx, task.ID, "insufficient_balance", reserveErr.Error(), task.RunMode == "simulation", reserveErr, parameterPreprocessingMetrics(firstPreprocessing)) failed, finishErr := s.failTask(ctx, task.ID, task.ExecutionToken, "insufficient_balance", reserveErr.Error(), task.RunMode == "simulation", reserveErr, parameterPreprocessingMetrics(firstPreprocessing))
if finishErr != nil { if finishErr != nil {
return Result{}, finishErr return Result{}, finishErr
} }
@@ -406,19 +494,46 @@ candidatesLoop:
break candidatesLoop break candidatesLoop
} }
candidateBody := preprocessing.Body candidateBody := preprocessing.Body
response, err := s.runCandidate(ctx, task, user, candidateBody, preprocessing.Log, candidate, nextAttemptNo, onDelta, responseExecution, singleSourceProtected, runnerPolicy.CacheAffinityPolicy, cacheAffinityKeys.Record) candidatePricing := pricingByCandidate[pricingCandidateKey(candidate)]
response, err := s.runCandidate(ctx, task, user, candidateBody, preprocessing.Log, candidate, candidatePricing, nextAttemptNo, onDelta, responseExecution, singleSourceProtected, runnerPolicy.CacheAffinityPolicy, cacheAffinityKeys.Record)
if err == nil { if err == nil {
attemptNo = nextAttemptNo attemptNo = nextAttemptNo
var billings []any var billings []any
finalAmount := fixedAmount(0)
finalAmountText := ""
pricingSnapshot := map[string]any{"pricingVersion": pricingVersionV2, "source": "simulation"}
if task.RunMode == "production" { if task.RunMode == "production" {
pricing := pricingByCandidate[pricingCandidateKey(candidate)] billingMode := normalizedBillingEngineMode(s.cfg.BillingEngineMode)
var billingErr error var billingErr error
billings, _, _, billingErr = s.billingsWithResolvedPricingV2(ctx, user, task.Kind, candidateBody, candidate, response, false, pricing) if billingMode == "observe" {
if billingErr != nil { billings = s.billings(ctx, user, task.Kind, candidateBody, candidate, response, false)
// The upstream result may already exist. Preserve the reservation for manual review. finalAmount = billingItemsFixedTotal(billings)
walletReservationFinalized = true pricingSnapshot = map[string]any{
return Result{}, billingErr "pricingVersion": "legacy-observe", "observedPricingVersion": pricingVersionV2,
"observedPricing": candidatePricing.Snapshot,
}
} else {
billings, finalAmount, _, billingErr = s.billingsWithResolvedPricingV2(ctx, user, task.Kind, candidateBody, candidate, response, false, candidatePricing)
pricingSnapshot = candidatePricing.Snapshot
} }
if billingErr != nil {
review, reviewErr := s.store.FinishTaskManualReview(context.WithoutCancel(ctx), store.FinishTaskManualReviewInput{
TaskID: task.ID, ExecutionToken: task.ExecutionToken, AttemptID: response.AttemptID, TaskStatus: "succeeded",
Code: "billing_calculation_failed", Message: "billing calculation requires manual review",
Result: response.Result, RequestID: response.RequestID,
PricingSnapshot: candidatePricing.Snapshot,
RequestFingerprint: pricingRequestFingerprint(task.Kind, task.Model, candidateBody),
ResponseStartedAt: response.ResponseStartedAt, ResponseFinishedAt: response.ResponseFinishedAt,
ResponseDurationMS: response.ResponseDurationMS,
})
if reviewErr != nil {
return Result{}, reviewErr
}
walletReservationFinalized = true
s.logger.Warn("task succeeded but billing requires manual review", "taskID", task.ID, "error_category", "billing_calculation_failed")
return Result{Task: review, Output: response.Result}, nil
}
finalAmountText = finalAmount.String()
} else { } else {
billings = s.billings(ctx, user, task.Kind, candidateBody, candidate, response, true) billings = s.billings(ctx, user, task.Kind, candidateBody, candidate, response, true)
} }
@@ -427,38 +542,32 @@ candidatesLoop:
record.Metrics = mergeMetrics(record.Metrics, parameterPreprocessingMetrics(preprocessing.Log)) record.Metrics = mergeMetrics(record.Metrics, parameterPreprocessingMetrics(preprocessing.Log))
record.Metrics = s.withAttemptHistory(ctx, task.ID, record.Metrics) record.Metrics = s.withAttemptHistory(ctx, task.ID, record.Metrics)
finished, finishErr := s.store.FinishTaskSuccess(ctx, store.FinishTaskSuccessInput{ finished, finishErr := s.store.FinishTaskSuccess(ctx, store.FinishTaskSuccessInput{
TaskID: task.ID, TaskID: task.ID,
Result: response.Result, ExecutionToken: task.ExecutionToken,
Billings: billings, AttemptID: response.AttemptID,
RequestID: record.RequestID, Result: response.Result,
ResolvedModel: record.ResolvedModel, Billings: billings,
Usage: record.Usage, RequestID: record.RequestID,
Metrics: record.Metrics, ResolvedModel: record.ResolvedModel,
BillingSummary: record.BillingSummary, Usage: record.Usage,
FinalChargeAmount: record.FinalChargeAmount, Metrics: record.Metrics,
ResponseStartedAt: record.ResponseStartedAt, BillingSummary: record.BillingSummary,
ResponseFinishedAt: record.ResponseFinishedAt, FinalChargeAmount: record.FinalChargeAmount,
ResponseDurationMS: record.ResponseDurationMS, FinalChargeAmountText: finalAmountText,
BillingCurrency: stringFromAny(record.BillingSummary["currency"]),
PricingSnapshot: pricingSnapshot,
RequestFingerprint: pricingRequestFingerprint(task.Kind, task.Model, candidateBody),
ResponseStartedAt: record.ResponseStartedAt,
ResponseFinishedAt: record.ResponseFinishedAt,
ResponseDurationMS: record.ResponseDurationMS,
}) })
if finishErr != nil { if finishErr != nil {
return Result{}, finishErr return Result{}, finishErr
} }
if finished.FinalChargeAmount > 0 {
walletReservationFinalized = true
if settleErr := s.store.SettleTaskBilling(ctx, finished); settleErr != nil {
return Result{}, settleErr
}
} else if len(walletReservations) > 0 {
if releaseErr := s.store.ReleaseTaskBillingReservations(ctx, walletReservations, "task_billing_zero"); releaseErr != nil {
return Result{}, releaseErr
}
walletReservationFinalized = true
}
walletReservationFinalized = true walletReservationFinalized = true
if finished.FinalChargeAmount > 0 { if finished.BillingStatus == "pending" {
if err := s.emit(ctx, task.ID, "task.billing.settled", "succeeded", "billing", 0.98, "task billing settled", map[string]any{ if err := s.emit(ctx, task.ID, "task.billing.pending", "succeeded", "billing", 0.98, "task billing queued", map[string]any{
"amount": finished.FinalChargeAmount, "amount": finished.FinalChargeAmount, "currency": finished.BillingCurrency,
"currency": stringFromAny(record.BillingSummary["currency"]),
}, isSimulation(task, candidate)); err != nil { }, isSimulation(task, candidate)); err != nil {
return Result{}, err return Result{}, err
} }
@@ -475,6 +584,21 @@ candidatesLoop:
} }
return Result{Task: finished, Output: response.Result}, nil return Result{Task: finished, Output: response.Result}, nil
} }
var submissionUnknown *upstreamSubmissionUnknownError
if errors.As(err, &submissionUnknown) {
review, reviewErr := s.store.FinishTaskManualReview(context.WithoutCancel(ctx), store.FinishTaskManualReviewInput{
TaskID: task.ID, ExecutionToken: task.ExecutionToken, AttemptID: submissionUnknown.AttemptID, TaskStatus: "failed",
Code: "upstream_submission_unknown", Message: submissionUnknown.Error(),
PricingSnapshot: candidatePricing.Snapshot,
RequestFingerprint: pricingRequestFingerprint(task.Kind, task.Model, candidateBody),
})
if reviewErr != nil {
return Result{}, reviewErr
}
walletReservationFinalized = true
s.logger.Warn("upstream submission requires manual review", "taskID", task.ID, "attemptID", submissionUnknown.AttemptID, "error_category", "upstream_submission_unknown")
return Result{Task: review, Output: review.Result}, submissionUnknown
}
if isLocalRateLimitError(err) { if isLocalRateLimitError(err) {
lastErr = err lastErr = err
candidateErr = err candidateErr = err
@@ -622,10 +746,12 @@ candidatesLoop:
if lastPreprocessing != nil { if lastPreprocessing != nil {
extraMetrics = append(extraMetrics, parameterPreprocessingMetrics(*lastPreprocessing)) extraMetrics = append(extraMetrics, parameterPreprocessingMetrics(*lastPreprocessing))
} }
failed, err := s.failTask(ctx, task.ID, code, message, task.RunMode == "simulation", lastErr, extraMetrics...) failed, err := s.failTask(ctx, task.ID, task.ExecutionToken, code, message, task.RunMode == "simulation", lastErr, extraMetrics...)
if err != nil { if err != nil {
return Result{}, err return Result{}, err
} }
// FinishTaskFailure atomically transfers ownership of any reservation to the release Outbox.
walletReservationFinalized = true
return Result{Task: failed, Output: failed.Result}, lastErr return Result{Task: failed, Output: failed.Result}, lastErr
} }
@@ -633,7 +759,46 @@ func pricingCandidateKey(candidate store.RuntimeModelCandidate) string {
return firstNonEmptyString(candidate.PlatformModelID, candidate.PlatformID+":"+candidate.ModelName) return firstNonEmptyString(candidate.PlatformModelID, candidate.PlatformID+":"+candidate.ModelName)
} }
func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user *auth.User, body map[string]any, preprocessing parameterPreprocessingLog, candidate store.RuntimeModelCandidate, attemptNo int, onDelta clients.StreamDelta, responseExecution responseExecutionContext, singleSourceProtected bool, cacheAffinityPolicy map[string]any, cacheAffinityRecordKeys []string) (clients.Response, error) { func normalizedBillingEngineMode(value string) string {
switch strings.ToLower(strings.TrimSpace(value)) {
case "enforce", "hold":
return strings.ToLower(strings.TrimSpace(value))
default:
return "observe"
}
}
func (s *Service) maximumLegacyCandidateEstimate(ctx context.Context, user *auth.User, kind string, body map[string]any, candidates []store.RuntimeModelCandidate) ([]any, fixedAmount) {
var maximumItems []any
maximumAmount := fixedAmount(0)
for index, candidate := range candidates {
candidateBody := preprocessRequest(kind, cloneMap(body), candidate)
items := s.estimatedBillings(ctx, user, kind, candidateBody, candidate)
amount := billingItemsFixedTotal(items)
if index == 0 || amount > maximumAmount {
maximumItems = items
maximumAmount = amount
}
}
return maximumItems, maximumAmount
}
func billingItemsFixedTotal(items []any) fixedAmount {
total := fixedAmount(0)
for _, raw := range items {
line, _ := raw.(map[string]any)
if line == nil {
continue
}
amount, err := fixedAmountFromAny(line["amount"])
if err == nil && amount > 0 {
total = total.Add(amount)
}
}
return total
}
func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user *auth.User, body map[string]any, preprocessing parameterPreprocessingLog, candidate store.RuntimeModelCandidate, pricing resolvedPricing, attemptNo int, onDelta clients.StreamDelta, responseExecution responseExecutionContext, singleSourceProtected bool, cacheAffinityPolicy map[string]any, cacheAffinityRecordKeys []string) (clients.Response, error) {
simulated := isSimulation(task, candidate) simulated := isSimulation(task, candidate)
baseAttemptMetrics := mergeMetrics(attemptMetrics(candidate, attemptNo, simulated), parameterPreprocessingMetrics(preprocessing)) baseAttemptMetrics := mergeMetrics(attemptMetrics(candidate, attemptNo, simulated), parameterPreprocessingMetrics(preprocessing))
reservations := s.rateLimitReservations(ctx, user, candidate, body) reservations := s.rateLimitReservations(ctx, user, candidate, body)
@@ -652,16 +817,18 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
defer s.store.ReleaseConcurrencyLeases(context.WithoutCancel(ctx), limitResult.LeaseIDs) defer s.store.ReleaseConcurrencyLeases(context.WithoutCancel(ctx), limitResult.LeaseIDs)
attemptID, err := s.store.CreateTaskAttempt(ctx, store.CreateTaskAttemptInput{ attemptID, err := s.store.CreateTaskAttempt(ctx, store.CreateTaskAttemptInput{
TaskID: task.ID, TaskID: task.ID,
AttemptNo: attemptNo, AttemptNo: attemptNo,
PlatformID: candidate.PlatformID, PlatformID: candidate.PlatformID,
PlatformModelID: candidate.PlatformModelID, PlatformModelID: candidate.PlatformModelID,
ClientID: candidate.ClientID, ClientID: candidate.ClientID,
QueueKey: candidate.QueueKey, QueueKey: candidate.QueueKey,
Status: "running", Status: "running",
Simulated: simulated, Simulated: simulated,
RequestSnapshot: s.slimTaskRequestSnapshot(task, body), RequestSnapshot: s.slimTaskRequestSnapshot(task, body),
Metrics: baseAttemptMetrics, Metrics: baseAttemptMetrics,
PricingSnapshot: pricing.Snapshot,
RequestFingerprint: pricingRequestFingerprint(task.Kind, task.Model, body),
}) })
if err != nil { if err != nil {
return clients.Response{}, fmt.Errorf("create task attempt: %w", err) return clients.Response{}, fmt.Errorf("create task attempt: %w", err)
@@ -731,6 +898,9 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
publicResponseID = responseExecution.PublicResponseID publicResponseID = responseExecution.PublicResponseID
publicPreviousResponseID = responseExecution.PublicPreviousResponseID publicPreviousResponseID = responseExecution.PublicPreviousResponseID
} }
if err := s.store.SetAttemptUpstreamSubmissionStatus(ctx, attemptID, "submitting"); err != nil {
return clients.Response{}, fmt.Errorf("mark upstream submission: %w", err)
}
response, err := client.Run(ctx, clients.Request{ response, err := client.Run(ctx, clients.Request{
Kind: task.Kind, Kind: task.Kind,
ModelType: candidate.ModelType, ModelType: candidate.ModelType,
@@ -744,7 +914,7 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
if strings.TrimSpace(remoteTaskID) == "" { if strings.TrimSpace(remoteTaskID) == "" {
return nil return nil
} }
return s.store.SetTaskRemoteTask(context.WithoutCancel(ctx), task.ID, attemptID, remoteTaskID, payload) return s.store.SetTaskRemoteTask(context.WithoutCancel(ctx), task.ID, task.ExecutionToken, attemptID, remoteTaskID, payload)
}, },
Stream: boolFromMap(providerBody, "stream"), Stream: boolFromMap(providerBody, "stream"),
StreamDelta: onDelta, StreamDelta: onDelta,
@@ -768,6 +938,9 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
} }
} }
if err != nil { if err != nil {
if clients.ErrorResponseMetadata(err).StatusCode > 0 {
_ = s.store.SetAttemptUpstreamSubmissionStatus(context.WithoutCancel(ctx), attemptID, "response_received")
}
retryable := clients.IsRetryable(err) retryable := clients.IsRetryable(err)
requestID, metrics, responseStartedAt, responseFinishedAt, responseDurationMS := failureMetrics(err, simulated) requestID, metrics, responseStartedAt, responseFinishedAt, responseDurationMS := failureMetrics(err, simulated)
if responseStartedAt.IsZero() { if responseStartedAt.IsZero() {
@@ -796,6 +969,9 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
ErrorMessage: err.Error(), ErrorMessage: err.Error(),
}) })
_ = s.emit(ctx, task.ID, "task.attempt.failed", "running", "attempt_failed", 0.45, err.Error(), map[string]any{"attempt": attemptNo, "retryable": retryable, "requestId": requestID, "statusCode": clients.ErrorResponseMetadata(err).StatusCode, "metrics": metrics}, simulated) _ = s.emit(ctx, task.ID, "task.attempt.failed", "running", "attempt_failed", 0.45, err.Error(), map[string]any{"attempt": attemptNo, "retryable": retryable, "requestId": requestID, "statusCode": clients.ErrorResponseMetadata(err).StatusCode, "metrics": metrics}, simulated)
if !simulated && clients.ErrorResponseMetadata(err).StatusCode == 0 {
return clients.Response{}, &upstreamSubmissionUnknownError{AttemptID: attemptID, Cause: err}
}
s.applyCandidateFailurePolicies(ctx, task.ID, candidate, err, simulated, singleSourceProtected) s.applyCandidateFailurePolicies(ctx, task.ID, candidate, err, simulated, singleSourceProtected)
return clients.Response{}, err return clients.Response{}, err
} }
@@ -877,19 +1053,7 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
return clients.Response{}, fmt.Errorf("commit rate limit reservations: %w", err) return clients.Response{}, fmt.Errorf("commit rate limit reservations: %w", err)
} }
rateReservationsFinalized = true rateReservationsFinalized = true
if err := s.store.FinishTaskAttempt(ctx, store.FinishTaskAttemptInput{ response.AttemptID = attemptID
AttemptID: attemptID,
Status: "succeeded",
RequestID: response.RequestID,
Usage: usageToMap(response.Usage),
Metrics: mergeMetrics(taskMetrics(task, user, body, candidate, response, simulated), parameterPreprocessingMetrics(preprocessing)),
ResponseSnapshot: response.Result,
ResponseStartedAt: response.ResponseStartedAt,
ResponseFinishedAt: response.ResponseFinishedAt,
ResponseDurationMS: response.ResponseDurationMS,
}); err != nil {
return clients.Response{}, fmt.Errorf("finish task attempt: %w", err)
}
if err := s.store.RecordCacheAffinityObservation(context.WithoutCancel(ctx), store.CacheAffinityObservationInput{ if err := s.store.RecordCacheAffinityObservation(context.WithoutCancel(ctx), store.CacheAffinityObservationInput{
CacheAffinityKey: candidate.CacheAffinity.Key, CacheAffinityKey: candidate.CacheAffinity.Key,
CacheAffinityKeys: cacheAffinityRecordKeys, CacheAffinityKeys: cacheAffinityRecordKeys,
@@ -957,7 +1121,7 @@ func (s *Service) clientFor(candidate store.RuntimeModelCandidate, simulated boo
return s.clients["openai"] return s.clients["openai"]
} }
func (s *Service) failTask(ctx context.Context, taskID string, code string, message string, simulated bool, cause error, extraMetrics ...map[string]any) (store.GatewayTask, error) { func (s *Service) failTask(ctx context.Context, taskID string, executionToken string, code string, message string, simulated bool, cause error, extraMetrics ...map[string]any) (store.GatewayTask, error) {
requestID, metrics, responseStartedAt, responseFinishedAt, responseDurationMS := failureMetrics(cause, simulated) requestID, metrics, responseStartedAt, responseFinishedAt, responseDurationMS := failureMetrics(cause, simulated)
if len(extraMetrics) > 0 { if len(extraMetrics) > 0 {
values := append([]map[string]any{metrics}, extraMetrics...) values := append([]map[string]any{metrics}, extraMetrics...)
@@ -966,6 +1130,7 @@ func (s *Service) failTask(ctx context.Context, taskID string, code string, mess
metrics = s.withAttemptHistory(ctx, taskID, metrics) metrics = s.withAttemptHistory(ctx, taskID, metrics)
failed, err := s.store.FinishTaskFailure(ctx, store.FinishTaskFailureInput{ failed, err := s.store.FinishTaskFailure(ctx, store.FinishTaskFailureInput{
TaskID: taskID, TaskID: taskID,
ExecutionToken: executionToken,
Code: code, Code: code,
Message: message, Message: message,
Result: buildFailureResult(code, message, requestID, cause), Result: buildFailureResult(code, message, requestID, cause),
@@ -1094,7 +1259,7 @@ func (s *Service) requeueRateLimitedTask(ctx context.Context, task store.Gateway
if delay <= 0 { if delay <= 0 {
delay = 5 * time.Second delay = 5 * time.Second
} }
queued, err := s.store.RequeueTask(ctx, task.ID, delay, candidate.QueueKey) queued, err := s.store.RequeueTask(ctx, task.ID, task.ExecutionToken, delay, candidate.QueueKey)
if err != nil { if err != nil {
return store.GatewayTask{}, 0, err return store.GatewayTask{}, 0, err
} }
@@ -1110,7 +1275,7 @@ func (s *Service) requeueRateLimitedTask(ctx context.Context, task store.Gateway
} }
func (s *Service) requeueInterruptedAsyncTask(ctx context.Context, task store.GatewayTask) (store.GatewayTask, error) { func (s *Service) requeueInterruptedAsyncTask(ctx context.Context, task store.GatewayTask) (store.GatewayTask, error) {
queued, err := s.store.RequeueTask(ctx, task.ID, 0, "") queued, err := s.store.RequeueTask(ctx, task.ID, task.ExecutionToken, 0, "")
if err != nil { if err != nil {
return store.GatewayTask{}, err return store.GatewayTask{}, err
} }
@@ -182,6 +182,20 @@ func TestExecuteWithMockClientRejectsConcurrentTasksBeyondWalletBalance(t *testi
if got := mockClient.calls.Load(); got != 1 { if got := mockClient.calls.Load(); got != 1 {
t.Fatalf("mock client calls = %d, want 1", got) t.Fatalf("mock client calls = %d, want 1", got)
} }
settlements, err := db.ClaimBillingSettlements(ctx, "wallet-execute-test", store.BillingSettlementBatchSize, store.BillingSettlementLockTimeout)
if err != nil {
t.Fatalf("claim billing settlements: %v", err)
}
processed := 0
for _, settlement := range settlements {
if err := db.ProcessBillingSettlement(ctx, settlement); err != nil {
t.Fatalf("process billing settlement %s: %v", settlement.ID, err)
}
processed++
}
if processed != 1 {
t.Fatalf("processed billing settlements = %d, want 1", processed)
}
summary, err := db.GetWalletSummary(ctx, user, "resource") summary, err := db.GetWalletSummary(ctx, user, "resource")
if err != nil { if err != nil {
+72 -16
View File
@@ -18,25 +18,32 @@ type MetricsSnapshotProvider interface {
type DynamicMetricsSnapshotProvider interface { type DynamicMetricsSnapshotProvider interface {
MetricsSnapshotProvider MetricsSnapshotProvider
SecurityEventConnection(context.Context) (store.SecurityEventConnection, error) SecurityEventConnection(context.Context) (store.SecurityEventConnection, error)
BillingMetrics(context.Context) (store.BillingMetricsSnapshot, error)
} }
type Metrics struct { type Metrics struct {
accepted atomic.Uint64 accepted atomic.Uint64
rejected atomic.Uint64 rejected atomic.Uint64
duplicate atomic.Uint64 duplicate atomic.Uint64
sessionsDeleted atomic.Uint64 sessionsDeleted atomic.Uint64
watermarkRejected atomic.Uint64 watermarkRejected atomic.Uint64
verificationAccepted atomic.Uint64 verificationAccepted atomic.Uint64
heartbeatAccepted atomic.Uint64 heartbeatAccepted atomic.Uint64
heartbeatFailed atomic.Uint64 heartbeatFailed atomic.Uint64
introspectionActive atomic.Uint64 introspectionActive atomic.Uint64
introspectionInactive atomic.Uint64 introspectionInactive atomic.Uint64
introspectionFailed atomic.Uint64 introspectionFailed atomic.Uint64
jwksSSFFailed atomic.Uint64 jwksSSFFailed atomic.Uint64
jwksOIDCFailed atomic.Uint64 jwksOIDCFailed atomic.Uint64
processingCount atomic.Uint64 processingCount atomic.Uint64
processingNanos atomic.Uint64 processingNanos atomic.Uint64
processingBuckets [6]atomic.Uint64 processingBuckets [6]atomic.Uint64
billingSettlementCompleted atomic.Uint64
billingSettlementRetry atomic.Uint64
billingManualReview atomic.Uint64
billingEstimateFailed atomic.Uint64
billingIdempotentReplay atomic.Uint64
billingPricingUnavailable atomic.Uint64
} }
var processingDurationBounds = [...]time.Duration{ var processingDurationBounds = [...]time.Duration{
@@ -99,6 +106,23 @@ func (m *Metrics) ObserveJWKSRefreshFailure(source string) {
m.jwksOIDCFailed.Add(1) m.jwksOIDCFailed.Add(1)
} }
func (m *Metrics) ObserveBillingEvent(event string) {
switch event {
case "settlement_completed":
m.billingSettlementCompleted.Add(1)
case "settlement_retry":
m.billingSettlementRetry.Add(1)
case "manual_review":
m.billingManualReview.Add(1)
case "estimate_failed":
m.billingEstimateFailed.Add(1)
case "idempotent_replay":
m.billingIdempotentReplay.Add(1)
case "pricing_unavailable":
m.billingPricingUnavailable.Add(1)
}
}
func (m *Metrics) Handler(provider MetricsSnapshotProvider, issuer, audience string, enabled bool) http.Handler { func (m *Metrics) Handler(provider MetricsSnapshotProvider, issuer, audience string, enabled bool) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
mode := "disabled" mode := "disabled"
@@ -113,6 +137,17 @@ func (m *Metrics) Handler(provider MetricsSnapshotProvider, issuer, audience str
return return
} }
} }
billing := store.BillingMetricsSnapshot{}
if billingProvider, ok := provider.(interface {
BillingMetrics(context.Context) (store.BillingMetricsSnapshot, error)
}); ok {
var err error
billing, err = billingProvider.BillingMetrics(r.Context())
if err != nil {
http.Error(w, "metrics unavailable", http.StatusServiceUnavailable)
return
}
}
w.Header().Set("Content-Type", "text/plain; version=0.0.4; charset=utf-8") w.Header().Set("Content-Type", "text/plain; version=0.0.4; charset=utf-8")
outcomeCounters(w, "easyai_gateway_ssf_receipts_total", "Received SETs by bounded outcome.", []outcomeValue{ outcomeCounters(w, "easyai_gateway_ssf_receipts_total", "Received SETs by bounded outcome.", []outcomeValue{
{"accepted", m.accepted.Load()}, {"rejected", m.rejected.Load()}, {"duplicate", m.duplicate.Load()}, {"accepted", m.accepted.Load()}, {"rejected", m.rejected.Load()}, {"duplicate", m.duplicate.Load()},
@@ -156,6 +191,27 @@ func (m *Metrics) Handler(provider MetricsSnapshotProvider, issuer, audience str
} }
fmt.Fprintf(w, "easyai_gateway_ssf_mode{mode=\"%s\"} %d\n", modeName, value) fmt.Fprintf(w, "easyai_gateway_ssf_mode{mode=\"%s\"} %d\n", modeName, value)
} }
fmt.Fprintln(w, "# HELP easyai_gateway_billing_settlement_backlog Current unsettled Outbox records.")
fmt.Fprintln(w, "# TYPE easyai_gateway_billing_settlement_backlog gauge")
fmt.Fprintf(w, "easyai_gateway_billing_settlement_backlog %d\n", billing.SettlementBacklog)
fmt.Fprintln(w, "# HELP easyai_gateway_billing_settlement_delay_seconds Age of the oldest unsettled Outbox record.")
fmt.Fprintln(w, "# TYPE easyai_gateway_billing_settlement_delay_seconds gauge")
fmt.Fprintf(w, "easyai_gateway_billing_settlement_delay_seconds %.6f\n", billing.SettlementDelaySecs)
fmt.Fprintln(w, "# HELP easyai_gateway_billing_manual_review Current billing records requiring manual review.")
fmt.Fprintln(w, "# TYPE easyai_gateway_billing_manual_review gauge")
fmt.Fprintf(w, "easyai_gateway_billing_manual_review %d\n", billing.ManualReview)
fmt.Fprintln(w, "# HELP easyai_gateway_billing_orphan_frozen Wallets whose frozen amount differs from active reservations.")
fmt.Fprintln(w, "# TYPE easyai_gateway_billing_orphan_frozen gauge")
fmt.Fprintf(w, "easyai_gateway_billing_orphan_frozen %d\n", billing.OrphanFrozen)
fmt.Fprintln(w, "# HELP easyai_gateway_billing_pricing_unavailable_tasks Current tasks rejected because no effective price was available.")
fmt.Fprintln(w, "# TYPE easyai_gateway_billing_pricing_unavailable_tasks gauge")
fmt.Fprintf(w, "easyai_gateway_billing_pricing_unavailable_tasks %d\n", billing.PricingUnavailable)
plainCounter(w, "easyai_gateway_billing_settlements_completed_total", "Billing settlements completed by this process.", m.billingSettlementCompleted.Load())
plainCounter(w, "easyai_gateway_billing_settlement_retries_total", "Billing settlement retries scheduled by this process.", m.billingSettlementRetry.Load())
plainCounter(w, "easyai_gateway_billing_manual_review_transitions_total", "Billing records transitioned to manual review by this process.", m.billingManualReview.Load())
plainCounter(w, "easyai_gateway_billing_estimate_failures_total", "Pricing estimate requests that failed.", m.billingEstimateFailed.Load())
plainCounter(w, "easyai_gateway_billing_idempotent_replays_total", "Generation requests replayed idempotently.", m.billingIdempotentReplay.Load())
plainCounter(w, "easyai_gateway_billing_pricing_unavailable_total", "Pricing requests rejected because no effective price was available.", m.billingPricingUnavailable.Load())
}) })
} }
@@ -0,0 +1,637 @@
package store
import (
"context"
"encoding/json"
"errors"
"fmt"
"strconv"
"strings"
"time"
"github.com/jackc/pgx/v5"
)
const (
BillingSettlementBatchSize = 50
BillingSettlementLockTimeout = 120 * time.Second
BillingSettlementMaxAttempts = 20
)
var ErrBillingSettlementNotRetryable = errors.New("billing settlement is not retryable")
type BillingSettlement struct {
ID string `json:"id"`
TaskID string `json:"taskId"`
Action string `json:"action"`
Amount float64 `json:"amount"`
Currency string `json:"currency"`
Status string `json:"status"`
Attempts int `json:"attempts"`
NextAttemptAt time.Time `json:"nextAttemptAt"`
LockedBy string `json:"lockedBy,omitempty"`
LockedAt string `json:"lockedAt,omitempty"`
LastErrorCode string `json:"lastErrorCode,omitempty"`
LastErrorMessage string `json:"lastErrorMessage,omitempty"`
ManualReviewReason string `json:"manualReviewReason,omitempty"`
PricingSnapshot map[string]any `json:"pricingSnapshot,omitempty"`
Payload map[string]any `json:"payload,omitempty"`
CompletedAt string `json:"completedAt,omitempty"`
CreatedAt time.Time `json:"createdAt"`
UpdatedAt time.Time `json:"updatedAt"`
LockToken string `json:"-"`
amountExact string
}
type BillingSettlementListFilter struct {
Status string
Action string
Page int
PageSize int
}
type BillingSettlementListResult struct {
Items []BillingSettlement `json:"items"`
Total int `json:"total"`
Page int `json:"page"`
PageSize int `json:"pageSize"`
}
type BillingMetricsSnapshot struct {
SettlementBacklog int64
SettlementDelaySecs float64
ManualReview int64
OrphanFrozen int64
PricingUnavailable int64
}
func (s *Store) BillingMetrics(ctx context.Context) (BillingMetricsSnapshot, error) {
var snapshot BillingMetricsSnapshot
err := s.pool.QueryRow(ctx, `
WITH open_outbox AS (
SELECT created_at
FROM settlement_outbox
WHERE status IN ('pending', 'processing', 'retryable_failed')
), active_reservations AS (
SELECT reserve.account_id, COALESCE(SUM(reserve.amount), 0) AS amount
FROM gateway_wallet_transactions reserve
WHERE reserve.transaction_type = 'reserve'
AND NOT EXISTS (
SELECT 1 FROM gateway_wallet_transactions release
WHERE release.account_id = reserve.account_id
AND release.transaction_type = 'release'
AND release.idempotency_key = reserve.idempotency_key || ':release'
)
GROUP BY reserve.account_id
)
SELECT
(SELECT COUNT(*) FROM open_outbox),
COALESCE((SELECT EXTRACT(EPOCH FROM now() - MIN(created_at)) FROM open_outbox), 0),
(SELECT COUNT(*) FROM gateway_tasks WHERE billing_status = 'manual_review'),
(SELECT COUNT(*) FROM gateway_wallet_accounts account
LEFT JOIN active_reservations reservation ON reservation.account_id = account.id
WHERE account.frozen_balance <> COALESCE(reservation.amount, 0)),
(SELECT COUNT(*) FROM gateway_tasks WHERE error_code = 'pricing_unavailable')`).Scan(
&snapshot.SettlementBacklog,
&snapshot.SettlementDelaySecs,
&snapshot.ManualReview,
&snapshot.OrphanFrozen,
&snapshot.PricingUnavailable,
)
return snapshot, err
}
func (s *Store) ClaimBillingSettlements(ctx context.Context, workerID string, limit int, staleAfter time.Duration) ([]BillingSettlement, error) {
if limit <= 0 || limit > BillingSettlementBatchSize {
limit = BillingSettlementBatchSize
}
if staleAfter <= 0 {
staleAfter = BillingSettlementLockTimeout
}
rows, err := s.pool.Query(ctx, `
WITH picked AS (
SELECT id
FROM settlement_outbox
WHERE (
status IN ('pending', 'retryable_failed')
AND next_attempt_at <= now()
)
OR (
status = 'processing'
AND locked_at < now() - ($3::int * interval '1 second')
)
ORDER BY next_attempt_at ASC, created_at ASC
LIMIT $2
FOR UPDATE SKIP LOCKED
)
UPDATE settlement_outbox outbox
SET status = 'processing',
attempts = outbox.attempts + 1,
locked_by = $1,
lock_token = gen_random_uuid(),
locked_at = now(),
updated_at = now()
FROM picked
WHERE outbox.id = picked.id
RETURNING outbox.id::text, outbox.task_id::text, outbox.action, outbox.amount::text,
outbox.currency, outbox.status, outbox.attempts, outbox.next_attempt_at,
COALESCE(outbox.locked_by, ''), COALESCE(outbox.lock_token::text, ''),
COALESCE(outbox.locked_at::text, ''), COALESCE(outbox.last_error_code, ''),
COALESCE(outbox.last_error_message, ''), COALESCE(outbox.manual_review_reason, ''),
outbox.pricing_snapshot, outbox.payload, COALESCE(outbox.completed_at::text, ''),
outbox.created_at, outbox.updated_at`,
workerID, limit, int(staleAfter/time.Second))
if err != nil {
return nil, err
}
defer rows.Close()
items := make([]BillingSettlement, 0)
for rows.Next() {
item, err := scanBillingSettlement(rows)
if err != nil {
return nil, err
}
items = append(items, item)
}
return items, rows.Err()
}
func (s *Store) ProcessBillingSettlement(ctx context.Context, settlement BillingSettlement) error {
return pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error {
var action string
var amount string
var currency string
var taskID string
var taskStatus string
var gatewayUserID string
var gatewayTenantID string
var userID string
var tenantID string
var tenantKey string
err := tx.QueryRow(ctx, `
SELECT outbox.action, outbox.amount::text, outbox.currency, outbox.task_id::text,
task.status, COALESCE(task.gateway_user_id::text, ''),
COALESCE(task.gateway_tenant_id::text, ''), task.user_id,
COALESCE(task.tenant_id, ''), COALESCE(task.tenant_key, '')
FROM settlement_outbox outbox
JOIN gateway_tasks task ON task.id = outbox.task_id
WHERE outbox.id = $1::uuid
AND outbox.status = 'processing'
AND outbox.lock_token = $2::uuid
FOR UPDATE OF outbox, task`, settlement.ID, settlement.LockToken).Scan(
&action, &amount, &currency, &taskID, &taskStatus, &gatewayUserID,
&gatewayTenantID, &userID, &tenantID, &tenantKey,
)
if err != nil {
return err
}
if _, err := tx.Exec(ctx, `
UPDATE gateway_tasks
SET billing_status = 'processing', billing_updated_at = now(), updated_at = now()
WHERE id = $1::uuid`, taskID); err != nil {
return err
}
switch action {
case "settle":
if taskStatus != "succeeded" {
return fmt.Errorf("settle action requires succeeded task")
}
if err := settleBillingOutboxTx(ctx, tx, taskID, gatewayUserID, gatewayTenantID, userID, tenantID, tenantKey, currency, amount); err != nil {
return err
}
return completeBillingOutboxTx(ctx, tx, settlement.ID, settlement.LockToken, taskID, "settled")
case "release":
if err := releaseBillingOutboxTx(ctx, tx, taskID, gatewayUserID, gatewayTenantID, currency); err != nil {
return err
}
return completeBillingOutboxTx(ctx, tx, settlement.ID, settlement.LockToken, taskID, "released")
default:
return fmt.Errorf("unsupported billing settlement action %q", action)
}
})
}
func settleBillingOutboxTx(
ctx context.Context,
tx pgx.Tx,
taskID string,
gatewayUserID string,
gatewayTenantID string,
userID string,
tenantID string,
tenantKey string,
currency string,
amount string,
) error {
if gatewayUserID == "" {
return fmt.Errorf("task %s has no gateway wallet user", taskID)
}
if _, err := tx.Exec(ctx, `
INSERT INTO gateway_wallet_accounts (
gateway_tenant_id, gateway_user_id, tenant_id, tenant_key, user_id, currency
)
VALUES (NULLIF($1, '')::uuid, $2::uuid, NULLIF($3, ''), NULLIF($4, ''), NULLIF($5, ''), $6)
ON CONFLICT (gateway_user_id, currency) DO NOTHING`,
gatewayTenantID, gatewayUserID, tenantID, tenantKey, userID, currency); err != nil {
return err
}
if err := ensureWalletAccountAuditGuard(ctx, tx, gatewayUserID, currency); err != nil {
return err
}
var accountID string
var balanceBefore string
var frozenBefore string
if err := tx.QueryRow(ctx, `
SELECT id::text, balance::text, frozen_balance::text
FROM gateway_wallet_accounts
WHERE gateway_user_id = $1::uuid AND currency = $2
FOR UPDATE`, gatewayUserID, currency).Scan(&accountID, &balanceBefore, &frozenBefore); err != nil {
return err
}
var alreadySettled bool
if err := tx.QueryRow(ctx, `
SELECT EXISTS (
SELECT 1 FROM gateway_wallet_transactions
WHERE account_id = $1::uuid AND idempotency_key = $2
)`, accountID, billingIdempotencyKey(taskID)).Scan(&alreadySettled); err != nil {
return err
}
if alreadySettled {
return nil
}
var reservationKey string
var reservedAmount string
if err := tx.QueryRow(ctx, `
SELECT COALESCE(reserve.idempotency_key, ''), reserve.amount::text
FROM gateway_wallet_transactions reserve
WHERE reserve.account_id = $1::uuid
AND reserve.reference_type = 'gateway_task'
AND reserve.reference_id = $2
AND reserve.transaction_type = 'reserve'
AND NOT EXISTS (
SELECT 1 FROM gateway_wallet_transactions release
WHERE release.account_id = reserve.account_id
AND release.transaction_type = 'release'
AND release.idempotency_key = reserve.idempotency_key || ':release'
)
ORDER BY reserve.created_at DESC
LIMIT 1`, accountID, taskID).Scan(&reservationKey, &reservedAmount); errors.Is(err, pgx.ErrNoRows) {
reservationKey = ""
reservedAmount = "0"
} else if err != nil {
return err
}
var balanceAfter string
var frozenAfter string
err := tx.QueryRow(ctx, `
UPDATE gateway_wallet_accounts
SET balance = balance - $2::numeric,
total_spent = total_spent + $2::numeric,
frozen_balance = GREATEST(0, frozen_balance - $3::numeric),
updated_at = now()
WHERE id = $1::uuid
AND balance - frozen_balance + $3::numeric >= $2::numeric
RETURNING balance::text, frozen_balance::text`,
accountID, amount, reservedAmount).Scan(&balanceAfter, &frozenAfter)
if errors.Is(err, pgx.ErrNoRows) {
return fmt.Errorf("%w: task %s settlement amount exceeds spendable balance", ErrInsufficientWalletBalance, taskID)
}
if err != nil {
return err
}
var reservedPositive bool
if err := tx.QueryRow(ctx, `SELECT $1::numeric > 0`, reservedAmount).Scan(&reservedPositive); err != nil {
return err
}
if reservedPositive {
releaseMetadata, _ := json.Marshal(map[string]any{
"taskId": taskID, "reason": "task_billing_settled",
"reserved": reservedAmount, "frozenBefore": frozenBefore, "frozenAfter": frozenAfter,
})
if _, err := tx.Exec(ctx, `
INSERT INTO gateway_wallet_transactions (
account_id, gateway_tenant_id, gateway_user_id, direction, transaction_type,
amount, balance_before, balance_after, idempotency_key, reference_type, reference_id, metadata
)
VALUES (
$1::uuid, NULLIF($2, '')::uuid, $3::uuid, 'credit', 'release',
$4::numeric, $5::numeric, $5::numeric, $6, 'gateway_task', $7, $8::jsonb
)
ON CONFLICT (account_id, idempotency_key) WHERE idempotency_key IS NOT NULL DO NOTHING`,
accountID, gatewayTenantID, gatewayUserID, reservedAmount, balanceBefore,
billingReservationReleaseIdempotencyKey(reservationKey), taskID, string(releaseMetadata)); err != nil {
return err
}
}
metadata, _ := json.Marshal(map[string]any{
"taskId": taskID, "reservedAmount": reservedAmount,
"frozenBefore": frozenBefore, "frozenAfter": frozenAfter,
})
_, err = tx.Exec(ctx, `
INSERT INTO gateway_wallet_transactions (
account_id, gateway_tenant_id, gateway_user_id, direction, transaction_type,
amount, balance_before, balance_after, idempotency_key, reference_type, reference_id, metadata
)
VALUES (
$1::uuid, NULLIF($2, '')::uuid, $3::uuid, 'debit', 'task_billing',
$4::numeric, $5::numeric, $6::numeric, $7, 'gateway_task', $8, $9::jsonb
)
ON CONFLICT (account_id, idempotency_key) WHERE idempotency_key IS NOT NULL DO NOTHING`,
accountID, gatewayTenantID, gatewayUserID, amount, balanceBefore, balanceAfter,
billingIdempotencyKey(taskID), taskID, string(metadata))
return err
}
func releaseBillingOutboxTx(ctx context.Context, tx pgx.Tx, taskID string, gatewayUserID string, gatewayTenantID string, currency string) error {
if gatewayUserID == "" {
return nil
}
var accountID string
var balance string
var frozenBefore string
err := tx.QueryRow(ctx, `
SELECT id::text, balance::text, frozen_balance::text
FROM gateway_wallet_accounts
WHERE gateway_user_id = $1::uuid AND currency = $2
FOR UPDATE`, gatewayUserID, currency).Scan(&accountID, &balance, &frozenBefore)
if errors.Is(err, pgx.ErrNoRows) {
return nil
}
if err != nil {
return err
}
var reservationKey string
var reservedAmount string
err = tx.QueryRow(ctx, `
SELECT COALESCE(reserve.idempotency_key, ''), reserve.amount::text
FROM gateway_wallet_transactions reserve
WHERE reserve.account_id = $1::uuid
AND reserve.reference_type = 'gateway_task'
AND reserve.reference_id = $2
AND reserve.transaction_type = 'reserve'
AND NOT EXISTS (
SELECT 1 FROM gateway_wallet_transactions release
WHERE release.account_id = reserve.account_id
AND release.transaction_type = 'release'
AND release.idempotency_key = reserve.idempotency_key || ':release'
)
ORDER BY reserve.created_at DESC
LIMIT 1`, accountID, taskID).Scan(&reservationKey, &reservedAmount)
if errors.Is(err, pgx.ErrNoRows) {
return nil
}
if err != nil {
return err
}
var frozenAfter string
if err := tx.QueryRow(ctx, `
UPDATE gateway_wallet_accounts
SET frozen_balance = frozen_balance - $2::numeric,
updated_at = now()
WHERE id = $1::uuid
AND frozen_balance >= $2::numeric
RETURNING frozen_balance::text`, accountID, reservedAmount).Scan(&frozenAfter); err != nil {
return err
}
metadata, _ := json.Marshal(map[string]any{
"taskId": taskID, "reason": "task_terminal_release",
"reserved": reservedAmount, "frozenBefore": frozenBefore, "frozenAfter": frozenAfter,
})
_, err = tx.Exec(ctx, `
INSERT INTO gateway_wallet_transactions (
account_id, gateway_tenant_id, gateway_user_id, direction, transaction_type,
amount, balance_before, balance_after, idempotency_key, reference_type, reference_id, metadata
)
VALUES (
$1::uuid, NULLIF($2, '')::uuid, $3::uuid, 'credit', 'release',
$4::numeric, $5::numeric, $5::numeric, $6, 'gateway_task', $7, $8::jsonb
)
ON CONFLICT (account_id, idempotency_key) WHERE idempotency_key IS NOT NULL DO NOTHING`,
accountID, gatewayTenantID, gatewayUserID, reservedAmount, balance,
billingReservationReleaseIdempotencyKey(reservationKey), taskID, string(metadata))
return err
}
func completeBillingOutboxTx(ctx context.Context, tx pgx.Tx, settlementID string, lockToken string, taskID string, billingStatus string) error {
if _, err := tx.Exec(ctx, `
UPDATE gateway_tasks
SET billing_status = $2,
reservation_amount = 0,
billing_updated_at = now(),
billing_settled_at = CASE WHEN $2 = 'settled' THEN now() ELSE billing_settled_at END,
updated_at = now()
WHERE id = $1::uuid`, taskID, billingStatus); err != nil {
return err
}
if _, err := tx.Exec(ctx, `
INSERT INTO gateway_task_events (task_id, seq, event_type, status, phase, progress, message, payload, simulated)
SELECT task.id,
COALESCE((SELECT MAX(event.seq) + 1 FROM gateway_task_events event WHERE event.task_id = task.id), 1),
CASE WHEN $2 = 'settled' THEN 'task.billing.settled' ELSE 'task.billing.released' END,
task.status,
'billing',
1,
CASE WHEN $2 = 'settled' THEN 'task billing settled' ELSE 'task billing reservation released' END,
jsonb_build_object('settlementId', $3::text, 'billingStatus', $2::text),
task.run_mode = 'simulation'
FROM gateway_tasks task
WHERE task.id = $1::uuid`, taskID, billingStatus, settlementID); err != nil {
return err
}
tag, err := tx.Exec(ctx, `
UPDATE settlement_outbox
SET status = 'completed',
completed_at = now(),
locked_by = NULL,
lock_token = NULL,
locked_at = NULL,
last_error_code = NULL,
last_error_message = NULL,
updated_at = now()
WHERE id = $1::uuid
AND status = 'processing'
AND lock_token = $2::uuid`, settlementID, lockToken)
if err != nil {
return err
}
if tag.RowsAffected() != 1 {
return pgx.ErrNoRows
}
return nil
}
func (s *Store) MarkBillingSettlementFailed(ctx context.Context, settlement BillingSettlement, code string, message string, delay time.Duration) error {
if delay < time.Second {
delay = time.Second
}
if delay > 15*time.Minute {
delay = 15 * time.Minute
}
manualReview := settlement.Attempts >= BillingSettlementMaxAttempts
status := "retryable_failed"
taskStatus := "retryable_failed"
manualReason := ""
if manualReview {
status = "manual_review"
taskStatus = "manual_review"
manualReason = "maximum settlement attempts exceeded"
}
return pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error {
tag, err := tx.Exec(ctx, `
UPDATE settlement_outbox
SET status = $3,
next_attempt_at = now() + ($4::int * interval '1 second'),
locked_by = NULL,
lock_token = NULL,
locked_at = NULL,
last_error_code = NULLIF($5, ''),
last_error_message = NULLIF($6, ''),
manual_review_reason = NULLIF($7, ''),
updated_at = now()
WHERE id = $1::uuid
AND lock_token = $2::uuid
AND status = 'processing'`,
settlement.ID, settlement.LockToken, status, int(delay/time.Second), code, message, manualReason)
if err != nil {
return err
}
if tag.RowsAffected() != 1 {
return pgx.ErrNoRows
}
_, err = tx.Exec(ctx, `
UPDATE gateway_tasks
SET billing_status = $2,
billing_updated_at = now(),
updated_at = now()
WHERE id = $1::uuid`, settlement.TaskID, taskStatus)
return err
})
}
func (s *Store) ListBillingSettlements(ctx context.Context, filter BillingSettlementListFilter) (BillingSettlementListResult, error) {
page := filter.Page
if page <= 0 {
page = 1
}
pageSize := filter.PageSize
if pageSize <= 0 {
pageSize = 50
}
if pageSize > 100 {
pageSize = 100
}
status := strings.TrimSpace(filter.Status)
action := strings.TrimSpace(filter.Action)
var total int
if err := s.pool.QueryRow(ctx, `
SELECT count(*)
FROM settlement_outbox
WHERE (NULLIF($1, '') IS NULL OR status = $1)
AND (NULLIF($2, '') IS NULL OR action = $2)`, status, action).Scan(&total); err != nil {
return BillingSettlementListResult{}, err
}
rows, err := s.pool.Query(ctx, `
SELECT id::text, task_id::text, action, amount::text, currency, status, attempts,
next_attempt_at, COALESCE(locked_by, ''), COALESCE(lock_token::text, ''),
COALESCE(locked_at::text, ''), COALESCE(last_error_code, ''),
COALESCE(last_error_message, ''), COALESCE(manual_review_reason, ''),
pricing_snapshot, payload, COALESCE(completed_at::text, ''), created_at, updated_at
FROM settlement_outbox
WHERE (NULLIF($1, '') IS NULL OR status = $1)
AND (NULLIF($2, '') IS NULL OR action = $2)
ORDER BY created_at DESC
LIMIT $3 OFFSET $4`, status, action, pageSize, (page-1)*pageSize)
if err != nil {
return BillingSettlementListResult{}, err
}
defer rows.Close()
items := make([]BillingSettlement, 0)
for rows.Next() {
item, err := scanBillingSettlement(rows)
if err != nil {
return BillingSettlementListResult{}, err
}
items = append(items, item)
}
if err := rows.Err(); err != nil {
return BillingSettlementListResult{}, err
}
return BillingSettlementListResult{Items: items, Total: total, Page: page, PageSize: pageSize}, nil
}
func (s *Store) RetryBillingSettlementTx(ctx context.Context, tx Tx, id string, idempotencyKeyHash string) (BillingSettlement, bool, error) {
current, err := scanBillingSettlement(tx.QueryRow(ctx, `
SELECT id::text, task_id::text, action, amount::text, currency, status, attempts,
next_attempt_at, COALESCE(locked_by, ''), COALESCE(lock_token::text, ''),
COALESCE(locked_at::text, ''), COALESCE(last_error_code, ''),
COALESCE(last_error_message, ''), COALESCE(manual_review_reason, ''),
pricing_snapshot, payload, COALESCE(completed_at::text, ''), created_at, updated_at
FROM settlement_outbox
WHERE id = $1::uuid
FOR UPDATE`, id))
if err != nil {
return BillingSettlement{}, false, err
}
var recordedHash string
if err := tx.QueryRow(ctx, `
SELECT COALESCE(retry_idempotency_key_hash, '')
FROM settlement_outbox
WHERE id = $1::uuid`, id).Scan(&recordedHash); err != nil {
return BillingSettlement{}, false, err
}
if recordedHash != "" && recordedHash == idempotencyKeyHash {
return current, true, nil
}
if current.Status != "retryable_failed" && current.Status != "manual_review" {
return BillingSettlement{}, false, ErrBillingSettlementNotRetryable
}
item, err := scanBillingSettlement(tx.QueryRow(ctx, `
UPDATE settlement_outbox
SET status = 'pending',
next_attempt_at = now(),
locked_by = NULL,
lock_token = NULL,
locked_at = NULL,
manual_review_reason = NULL,
retry_idempotency_key_hash = $2,
retry_requested_at = now(),
updated_at = now()
WHERE id = $1::uuid
AND status IN ('retryable_failed', 'manual_review')
RETURNING id::text, task_id::text, action, amount::text, currency, status, attempts,
next_attempt_at, COALESCE(locked_by, ''), COALESCE(lock_token::text, ''),
COALESCE(locked_at::text, ''), COALESCE(last_error_code, ''),
COALESCE(last_error_message, ''), COALESCE(manual_review_reason, ''),
pricing_snapshot, payload, COALESCE(completed_at::text, ''), created_at, updated_at`, id, idempotencyKeyHash))
if err != nil {
return BillingSettlement{}, false, err
}
if _, err := tx.Exec(ctx, `
UPDATE gateway_tasks
SET billing_status = 'pending', billing_updated_at = now(), updated_at = now()
WHERE id = $1::uuid`, item.TaskID); err != nil {
return BillingSettlement{}, false, err
}
return item, false, nil
}
func scanBillingSettlement(scanner taskScanner) (BillingSettlement, error) {
var item BillingSettlement
var pricingSnapshotBytes []byte
var payloadBytes []byte
if err := scanner.Scan(
&item.ID, &item.TaskID, &item.Action, &item.amountExact, &item.Currency,
&item.Status, &item.Attempts, &item.NextAttemptAt, &item.LockedBy, &item.LockToken,
&item.LockedAt, &item.LastErrorCode, &item.LastErrorMessage, &item.ManualReviewReason,
&pricingSnapshotBytes, &payloadBytes, &item.CompletedAt, &item.CreatedAt, &item.UpdatedAt,
); err != nil {
return BillingSettlement{}, err
}
item.Amount, _ = strconv.ParseFloat(item.amountExact, 64)
item.PricingSnapshot = decodeObject(pricingSnapshotBytes)
item.Payload = decodeObject(payloadBytes)
return item, nil
}
@@ -0,0 +1,265 @@
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)
}
}
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{}
}
+191 -112
View File
@@ -47,16 +47,21 @@ func normalizeAPIKeyScopes(scopes []string) []string {
} }
var ( var (
ErrInvalidCredentials = errors.New("invalid account or password") ErrInvalidCredentials = errors.New("invalid account or password")
ErrInvalidInvitation = errors.New("invalid or expired invitation code") ErrInvalidInvitation = errors.New("invalid or expired invitation code")
ErrInvalidAPIKeyScopes = errors.New("api key scopes must not be empty") ErrInvalidAPIKeyScopes = errors.New("api key scopes must not be empty")
ErrAccessRuleResourceDenied = errors.New("access rule resource is not available") ErrAccessRuleResourceDenied = errors.New("access rule resource is not available")
ErrInsufficientWalletBalance = errors.New("insufficient wallet balance") ErrInsufficientWalletBalance = errors.New("insufficient wallet balance")
ErrLocalUserRequired = errors.New("local gateway user is required") ErrLocalUserRequired = errors.New("local gateway user is required")
ErrWalletBalanceUnchanged = errors.New("wallet balance unchanged") ErrWalletBalanceUnchanged = errors.New("wallet balance unchanged")
ErrProtectedDefault = errors.New("protected default resource cannot be deleted") ErrBalanceBelowFrozen = errors.New("wallet balance cannot be below frozen balance")
ErrUserAlreadyExists = errors.New("user already exists") ErrInvalidWalletAmount = errors.New("wallet amount must be a decimal with at most nine fractional digits")
ErrWeakPassword = errors.New("password must be at least 8 characters") ErrIdempotencyKeyReused = errors.New("idempotency key was reused for a different request")
ErrTaskExecutionLeaseUnavailable = errors.New("task execution lease is unavailable")
ErrTaskExecutionLeaseLost = errors.New("task execution lease was lost")
ErrProtectedDefault = errors.New("protected default resource cannot be deleted")
ErrUserAlreadyExists = errors.New("user already exists")
ErrWeakPassword = errors.New("password must be at least 8 characters")
) )
func Connect(ctx context.Context, databaseURL string) (*Store, error) { func Connect(ctx context.Context, databaseURL string) (*Store, error) {
@@ -414,64 +419,81 @@ type RateLimitWindow struct {
} }
type CreateTaskInput struct { type CreateTaskInput struct {
Kind string `json:"kind"` Kind string `json:"kind"`
Model string `json:"model"` Model string `json:"model"`
RunMode string `json:"runMode"` RunMode string `json:"runMode"`
Async bool `json:"async"` Async bool `json:"async"`
Request map[string]any `json:"request"` Request map[string]any `json:"request"`
ConversationID string `json:"conversationId"` ConversationID string `json:"conversationId"`
NewMessageCount int `json:"newMessageCount"` NewMessageCount int `json:"newMessageCount"`
MessageRefs []TaskMessageRefInput `json:"messageRefs"` MessageRefs []TaskMessageRefInput `json:"messageRefs"`
IdempotencyKeyHash string `json:"-"`
IdempotencyRequestHash string `json:"-"`
}
type CreateTaskResult struct {
Task GatewayTask
Replayed bool
} }
type GatewayTask struct { type GatewayTask struct {
ID string `json:"id"` ID string `json:"id"`
Kind string `json:"kind"` Kind string `json:"kind"`
RunMode string `json:"runMode"` RunMode string `json:"runMode"`
UserID string `json:"userId"` UserID string `json:"userId"`
GatewayUserID string `json:"gatewayUserId,omitempty"` GatewayUserID string `json:"gatewayUserId,omitempty"`
UserSource string `json:"userSource,omitempty"` UserSource string `json:"userSource,omitempty"`
GatewayTenantID string `json:"gatewayTenantId,omitempty"` GatewayTenantID string `json:"gatewayTenantId,omitempty"`
TenantID string `json:"tenantId,omitempty"` TenantID string `json:"tenantId,omitempty"`
TenantKey string `json:"tenantKey,omitempty"` TenantKey string `json:"tenantKey,omitempty"`
APIKeyID string `json:"apiKeyId,omitempty"` APIKeyID string `json:"apiKeyId,omitempty"`
APIKeyName string `json:"apiKeyName,omitempty"` APIKeyName string `json:"apiKeyName,omitempty"`
APIKeyPrefix string `json:"apiKeyPrefix,omitempty"` APIKeyPrefix string `json:"apiKeyPrefix,omitempty"`
UserGroupID string `json:"userGroupId,omitempty"` UserGroupID string `json:"userGroupId,omitempty"`
UserGroupKey string `json:"userGroupKey,omitempty"` UserGroupKey string `json:"userGroupKey,omitempty"`
Model string `json:"model"` Model string `json:"model"`
ModelType string `json:"modelType,omitempty"` ModelType string `json:"modelType,omitempty"`
RequestedModel string `json:"requestedModel,omitempty"` RequestedModel string `json:"requestedModel,omitempty"`
ResolvedModel string `json:"resolvedModel,omitempty"` ResolvedModel string `json:"resolvedModel,omitempty"`
RequestID string `json:"requestId,omitempty"` RequestID string `json:"requestId,omitempty"`
ConversationID string `json:"conversationId,omitempty"` ConversationID string `json:"conversationId,omitempty"`
NewMessageCount int `json:"newMessageCount,omitempty"` NewMessageCount int `json:"newMessageCount,omitempty"`
Request map[string]any `json:"request,omitempty"` Request map[string]any `json:"request,omitempty"`
AsyncMode bool `json:"asyncMode"` AsyncMode bool `json:"asyncMode"`
RiverJobID int64 `json:"riverJobId,omitempty"` RiverJobID int64 `json:"riverJobId,omitempty"`
Status string `json:"status"` Status string `json:"status"`
Cancellable *bool `json:"cancellable,omitempty"` Cancellable *bool `json:"cancellable,omitempty"`
Submitted *bool `json:"submitted,omitempty"` Submitted *bool `json:"submitted,omitempty"`
Message string `json:"message,omitempty"` Message string `json:"message,omitempty"`
AttemptCount int `json:"attemptCount"` AttemptCount int `json:"attemptCount"`
RemoteTaskID string `json:"remoteTaskId,omitempty"` RemoteTaskID string `json:"remoteTaskId,omitempty"`
RemoteTaskPayload map[string]any `json:"remoteTaskPayload,omitempty"` RemoteTaskPayload map[string]any `json:"remoteTaskPayload,omitempty"`
Result map[string]any `json:"result,omitempty"` Result map[string]any `json:"result,omitempty"`
Billings []any `json:"billings,omitempty"` Billings []any `json:"billings,omitempty"`
Usage map[string]any `json:"usage"` Usage map[string]any `json:"usage"`
Metrics map[string]any `json:"metrics"` Metrics map[string]any `json:"metrics"`
BillingSummary map[string]any `json:"billingSummary"` BillingSummary map[string]any `json:"billingSummary"`
FinalChargeAmount float64 `json:"finalChargeAmount"` FinalChargeAmount float64 `json:"finalChargeAmount"`
ResponseStartedAt string `json:"responseStartedAt,omitempty"` BillingVersion string `json:"billingVersion"`
ResponseFinishedAt string `json:"responseFinishedAt,omitempty"` BillingStatus string `json:"billingStatus"`
ResponseDurationMS int64 `json:"responseDurationMs"` BillingCurrency string `json:"billingCurrency"`
FinishedAt string `json:"finishedAt,omitempty"` PricingSnapshot map[string]any `json:"pricingSnapshot,omitempty"`
Error string `json:"error,omitempty"` RequestFingerprint string `json:"requestFingerprint,omitempty"`
ErrorCode string `json:"errorCode,omitempty"` ReservationAmount float64 `json:"reservationAmount"`
ErrorMessage string `json:"errorMessage,omitempty"` ExecutionToken string `json:"-"`
Attempts []TaskAttempt `json:"attempts,omitempty"` ExecutionLeaseUntil string `json:"executionLeaseExpiresAt,omitempty"`
CreatedAt time.Time `json:"createdAt"` BillingUpdatedAt string `json:"billingUpdatedAt,omitempty"`
UpdatedAt time.Time `json:"updatedAt"` BillingSettledAt string `json:"billingSettledAt,omitempty"`
ResponseStartedAt string `json:"responseStartedAt,omitempty"`
ResponseFinishedAt string `json:"responseFinishedAt,omitempty"`
ResponseDurationMS int64 `json:"responseDurationMs"`
FinishedAt string `json:"finishedAt,omitempty"`
Error string `json:"error,omitempty"`
ErrorCode string `json:"errorCode,omitempty"`
ErrorMessage string `json:"errorMessage,omitempty"`
Attempts []TaskAttempt `json:"attempts,omitempty"`
CreatedAt time.Time `json:"createdAt"`
UpdatedAt time.Time `json:"updatedAt"`
} }
const gatewayTaskColumns = ` const gatewayTaskColumns = `
@@ -485,7 +507,11 @@ request, COALESCE(async_mode, false), COALESCE(river_job_id, 0), status, COALESC
COALESCE(remote_task_id, ''), COALESCE(remote_task_payload, '{}'::jsonb), COALESCE(remote_task_id, ''), COALESCE(remote_task_payload, '{}'::jsonb),
COALESCE(result, '{}'::jsonb), COALESCE(billings, '[]'::jsonb), COALESCE(result, '{}'::jsonb), COALESCE(billings, '[]'::jsonb),
COALESCE(usage, '{}'::jsonb), COALESCE(metrics, '{}'::jsonb), COALESCE(billing_summary, '{}'::jsonb), COALESCE(usage, '{}'::jsonb), COALESCE(metrics, '{}'::jsonb), COALESCE(billing_summary, '{}'::jsonb),
COALESCE(final_charge_amount, 0)::float8, COALESCE(response_started_at::text, ''), COALESCE(final_charge_amount, 0)::float8, billing_version, billing_status, billing_currency,
COALESCE(pricing_snapshot, '{}'::jsonb), COALESCE(request_fingerprint, ''),
COALESCE(reservation_amount, 0)::float8, COALESCE(execution_token::text, ''),
COALESCE(execution_lease_expires_at::text, ''), COALESCE(billing_updated_at::text, ''),
COALESCE(billing_settled_at::text, ''), COALESCE(response_started_at::text, ''),
COALESCE(response_finished_at::text, ''), COALESCE(response_duration_ms, 0), COALESCE(error, ''), COALESCE(response_finished_at::text, ''), COALESCE(response_duration_ms, 0), COALESCE(error, ''),
COALESCE(error_code, ''), COALESCE(error_message, ''), COALESCE(error_code, ''), COALESCE(error_message, ''),
created_at, updated_at, COALESCE(finished_at::text, '')` created_at, updated_at, COALESCE(finished_at::text, '')`
@@ -505,35 +531,39 @@ type TaskEvent struct {
} }
type TaskAttempt struct { type TaskAttempt struct {
ID string `json:"id"` ID string `json:"id"`
TaskID string `json:"taskId"` TaskID string `json:"taskId"`
AttemptNo int `json:"attemptNo"` AttemptNo int `json:"attemptNo"`
PlatformID string `json:"platformId,omitempty"` PlatformID string `json:"platformId,omitempty"`
PlatformName string `json:"platformName,omitempty"` PlatformName string `json:"platformName,omitempty"`
Provider string `json:"provider,omitempty"` Provider string `json:"provider,omitempty"`
PlatformModelID string `json:"platformModelId,omitempty"` PlatformModelID string `json:"platformModelId,omitempty"`
ModelName string `json:"modelName,omitempty"` ModelName string `json:"modelName,omitempty"`
ProviderModelName string `json:"providerModelName,omitempty"` ProviderModelName string `json:"providerModelName,omitempty"`
ModelAlias string `json:"modelAlias,omitempty"` ModelAlias string `json:"modelAlias,omitempty"`
ModelType string `json:"modelType,omitempty"` ModelType string `json:"modelType,omitempty"`
ClientID string `json:"clientId,omitempty"` ClientID string `json:"clientId,omitempty"`
QueueKey string `json:"queueKey"` QueueKey string `json:"queueKey"`
Status string `json:"status"` Status string `json:"status"`
Retryable bool `json:"retryable"` Retryable bool `json:"retryable"`
Simulated bool `json:"simulated"` Simulated bool `json:"simulated"`
RequestID string `json:"requestId,omitempty"` RequestID string `json:"requestId,omitempty"`
StatusCode int `json:"statusCode,omitempty"` StatusCode int `json:"statusCode,omitempty"`
Usage map[string]any `json:"usage,omitempty"` Usage map[string]any `json:"usage,omitempty"`
Metrics map[string]any `json:"metrics,omitempty"` Metrics map[string]any `json:"metrics,omitempty"`
RequestSnapshot map[string]any `json:"requestSnapshot,omitempty"` RequestSnapshot map[string]any `json:"requestSnapshot,omitempty"`
ResponseSnapshot map[string]any `json:"responseSnapshot,omitempty"` ResponseSnapshot map[string]any `json:"responseSnapshot,omitempty"`
ResponseStartedAt string `json:"responseStartedAt,omitempty"` ResponseStartedAt string `json:"responseStartedAt,omitempty"`
ResponseFinishedAt string `json:"responseFinishedAt,omitempty"` ResponseFinishedAt string `json:"responseFinishedAt,omitempty"`
ResponseDurationMS int64 `json:"responseDurationMs"` ResponseDurationMS int64 `json:"responseDurationMs"`
ErrorCode string `json:"errorCode,omitempty"` PricingSnapshot map[string]any `json:"pricingSnapshot,omitempty"`
ErrorMessage string `json:"errorMessage,omitempty"` RequestFingerprint string `json:"requestFingerprint,omitempty"`
StartedAt time.Time `json:"startedAt"` UpstreamSubmissionStatus string `json:"upstreamSubmissionStatus"`
FinishedAt string `json:"finishedAt,omitempty"` UpstreamSubmissionUpdatedAt string `json:"upstreamSubmissionUpdatedAt,omitempty"`
ErrorCode string `json:"errorCode,omitempty"`
ErrorMessage string `json:"errorMessage,omitempty"`
StartedAt time.Time `json:"startedAt"`
FinishedAt string `json:"finishedAt,omitempty"`
} }
type TaskParamPreprocessingLog struct { type TaskParamPreprocessingLog struct {
@@ -1716,6 +1746,9 @@ ON CONFLICT (gateway_user_id, currency) DO NOTHING`,
); err != nil { ); err != nil {
return GatewayUser{}, err return GatewayUser{}, err
} }
if err := ensureWalletAccountAuditGuard(ctx, tx, user.ID, "resource"); err != nil {
return GatewayUser{}, err
}
if err := tx.Commit(ctx); err != nil { if err := tx.Commit(ctx); err != nil {
return GatewayUser{}, err return GatewayUser{}, err
} }
@@ -1826,6 +1859,11 @@ ORDER BY window_start DESC, scope_type ASC, scope_key ASC, metric ASC`)
} }
func (s *Store) CreateTask(ctx context.Context, input CreateTaskInput, user *auth.User) (GatewayTask, error) { func (s *Store) CreateTask(ctx context.Context, input CreateTaskInput, user *auth.User) (GatewayTask, error) {
result, err := s.CreateTaskIdempotent(ctx, input, user)
return result.Task, err
}
func (s *Store) CreateTaskIdempotent(ctx context.Context, input CreateTaskInput, user *auth.User) (CreateTaskResult, error) {
requestBody, _ := json.Marshal(input.Request) requestBody, _ := json.Marshal(input.Request)
runMode := normalizeRunMode(input.RunMode, input.Request) runMode := normalizeRunMode(input.RunMode, input.Request)
status := "queued" status := "queued"
@@ -1834,7 +1872,7 @@ func (s *Store) CreateTask(ctx context.Context, input CreateTaskInput, user *aut
tx, err := s.pool.Begin(ctx) tx, err := s.pool.Begin(ctx)
if err != nil { if err != nil {
return GatewayTask{}, err return CreateTaskResult{}, err
} }
defer tx.Rollback(ctx) defer tx.Rollback(ctx)
@@ -1842,33 +1880,62 @@ func (s *Store) CreateTask(ctx context.Context, input CreateTaskInput, user *aut
INSERT INTO gateway_tasks ( INSERT INTO gateway_tasks (
kind, run_mode, user_id, gateway_user_id, user_source, gateway_tenant_id, tenant_id, tenant_key, kind, run_mode, user_id, gateway_user_id, user_source, gateway_tenant_id, tenant_id, tenant_key,
api_key_id, api_key_name, api_key_prefix, user_group_id, user_group_key, api_key_id, api_key_name, api_key_prefix, user_group_id, user_group_key,
model, requested_model, request, async_mode, status, result, billings, conversation_id, new_message_count, finished_at model, requested_model, request, async_mode, status, result, billings, conversation_id, new_message_count,
idempotency_key_hash, idempotency_request_hash, finished_at
) )
VALUES ($1, $2, $3, NULLIF($4, '')::uuid, COALESCE(NULLIF($5, ''), 'gateway'), NULLIF($6, '')::uuid, NULLIF($7, ''), NULLIF($8, ''), NULLIF($9, ''), NULLIF($10, ''), NULLIF($11, ''), NULLIF($12, '')::uuid, NULLIF($13, ''), $14, $14, $15, $16, $17, $18::jsonb, $19::jsonb, NULLIF($20, '')::uuid, $21, CASE WHEN $22 THEN now() ELSE NULL END) VALUES ($1, $2, $3, NULLIF($4, '')::uuid, COALESCE(NULLIF($5, ''), 'gateway'), NULLIF($6, '')::uuid, NULLIF($7, ''), NULLIF($8, ''), NULLIF($9, ''), NULLIF($10, ''), NULLIF($11, ''), NULLIF($12, '')::uuid, NULLIF($13, ''), $14, $14, $15, $16, $17, $18::jsonb, $19::jsonb, NULLIF($20, '')::uuid, $21, NULLIF($22, ''), NULLIF($23, ''), NULL)
ON CONFLICT (user_id, idempotency_key_hash) WHERE idempotency_key_hash IS NOT NULL DO NOTHING
RETURNING `+gatewayTaskColumns, RETURNING `+gatewayTaskColumns,
input.Kind, runMode, user.ID, user.GatewayUserID, user.Source, user.GatewayTenantID, user.TenantID, user.TenantKey, user.APIKeyID, user.APIKeyName, user.APIKeyPrefix, user.UserGroupID, user.UserGroupKey, input.Model, requestBody, input.Async, status, resultBody, billingsBody, input.ConversationID, input.NewMessageCount, false, input.Kind, runMode, user.ID, user.GatewayUserID, user.Source, user.GatewayTenantID, user.TenantID, user.TenantKey, user.APIKeyID, user.APIKeyName, user.APIKeyPrefix, user.UserGroupID, user.UserGroupKey, input.Model, requestBody, input.Async, status, resultBody, billingsBody, input.ConversationID, input.NewMessageCount, strings.TrimSpace(input.IdempotencyKeyHash), strings.TrimSpace(input.IdempotencyRequestHash),
)) ))
replayed := false
if errors.Is(err, pgx.ErrNoRows) && strings.TrimSpace(input.IdempotencyKeyHash) != "" {
var existingRequestHash string
if err := tx.QueryRow(ctx, `
SELECT COALESCE(idempotency_request_hash, '')
FROM gateway_tasks
WHERE user_id = $1 AND idempotency_key_hash = $2
FOR UPDATE`, user.ID, strings.TrimSpace(input.IdempotencyKeyHash)).Scan(&existingRequestHash); err != nil {
return CreateTaskResult{}, err
}
if existingRequestHash != strings.TrimSpace(input.IdempotencyRequestHash) {
return CreateTaskResult{}, ErrIdempotencyKeyReused
}
task, err = scanGatewayTask(tx.QueryRow(ctx, `
SELECT `+gatewayTaskColumns+`
FROM gateway_tasks
WHERE user_id = $1 AND idempotency_key_hash = $2`, user.ID, strings.TrimSpace(input.IdempotencyKeyHash)))
replayed = true
}
if err != nil { if err != nil {
return GatewayTask{}, err return CreateTaskResult{}, err
} }
if err := insertTaskMessageRefs(ctx, tx, task.ID, input.MessageRefs); err != nil { if !replayed {
return GatewayTask{}, err if err := insertTaskMessageRefs(ctx, tx, task.ID, input.MessageRefs); err != nil {
} return CreateTaskResult{}, err
events := taskEventsForCreate(task.ID, runMode, status, nil) }
for _, event := range events { events := taskEventsForCreate(task.ID, runMode, status, nil)
payload, _ := json.Marshal(event.Payload) for _, event := range events {
if _, err := tx.Exec(ctx, ` payload, _ := json.Marshal(event.Payload)
if _, err := tx.Exec(ctx, `
INSERT INTO gateway_task_events (task_id, seq, event_type, status, phase, progress, message, payload, simulated) INSERT INTO gateway_task_events (task_id, seq, event_type, status, phase, progress, message, payload, simulated)
VALUES ($1::uuid, $2, $3::text, NULLIF($4::text, ''), NULLIF($5::text, ''), $6, NULLIF($7::text, ''), $8::jsonb, $9)`, VALUES ($1::uuid, $2, $3::text, NULLIF($4::text, ''), NULLIF($5::text, ''), $6, NULLIF($7::text, ''), $8::jsonb, $9)`,
task.ID, event.Seq, event.EventType, event.Status, event.Phase, event.Progress, event.Message, string(payload), event.Simulated, task.ID, event.Seq, event.EventType, event.Status, event.Phase, event.Progress, event.Message, string(payload), event.Simulated,
); err != nil { ); err != nil {
return GatewayTask{}, err return CreateTaskResult{}, err
}
} }
} }
if err := tx.Commit(ctx); err != nil { if err := tx.Commit(ctx); err != nil {
return GatewayTask{}, err return CreateTaskResult{}, err
} }
return task, nil if replayed {
task, err = s.GetTask(ctx, task.ID)
if err != nil {
return CreateTaskResult{}, err
}
}
return CreateTaskResult{Task: task, Replayed: replayed}, nil
} }
func (s *Store) GetTask(ctx context.Context, taskID string) (GatewayTask, error) { func (s *Store) GetTask(ctx context.Context, taskID string) (GatewayTask, error) {
@@ -1900,6 +1967,7 @@ func scanGatewayTask(scanner taskScanner) (GatewayTask, error) {
var usageBytes []byte var usageBytes []byte
var metricsBytes []byte var metricsBytes []byte
var billingSummaryBytes []byte var billingSummaryBytes []byte
var pricingSnapshotBytes []byte
var remoteTaskPayloadBytes []byte var remoteTaskPayloadBytes []byte
if err := scanner.Scan( if err := scanner.Scan(
&task.ID, &task.ID,
@@ -1936,6 +2004,16 @@ func scanGatewayTask(scanner taskScanner) (GatewayTask, error) {
&metricsBytes, &metricsBytes,
&billingSummaryBytes, &billingSummaryBytes,
&task.FinalChargeAmount, &task.FinalChargeAmount,
&task.BillingVersion,
&task.BillingStatus,
&task.BillingCurrency,
&pricingSnapshotBytes,
&task.RequestFingerprint,
&task.ReservationAmount,
&task.ExecutionToken,
&task.ExecutionLeaseUntil,
&task.BillingUpdatedAt,
&task.BillingSettledAt,
&task.ResponseStartedAt, &task.ResponseStartedAt,
&task.ResponseFinishedAt, &task.ResponseFinishedAt,
&task.ResponseDurationMS, &task.ResponseDurationMS,
@@ -1955,6 +2033,7 @@ func scanGatewayTask(scanner taskScanner) (GatewayTask, error) {
task.Usage = decodeObject(usageBytes) task.Usage = decodeObject(usageBytes)
task.Metrics = decodeObject(metricsBytes) task.Metrics = decodeObject(metricsBytes)
task.BillingSummary = decodeObject(billingSummaryBytes) task.BillingSummary = decodeObject(billingSummaryBytes)
task.PricingSnapshot = decodeObject(pricingSnapshotBytes)
return task, nil return task, nil
} }
+44 -18
View File
@@ -223,21 +223,24 @@ type RateLimitResult struct {
} }
type CreateTaskAttemptInput struct { type CreateTaskAttemptInput struct {
TaskID string TaskID string
AttemptNo int AttemptNo int
PlatformID string PlatformID string
PlatformModelID string PlatformModelID string
ClientID string ClientID string
QueueKey string QueueKey string
Status string Status string
Simulated bool Simulated bool
RequestSnapshot map[string]any RequestSnapshot map[string]any
Metrics map[string]any Metrics map[string]any
PricingSnapshot map[string]any
RequestFingerprint string
} }
type AsyncTaskQueueItem struct { type AsyncTaskQueueItem struct {
ID string ID string
Priority int Priority int
NextRunAt time.Time
} }
type FinishTaskAttemptInput struct { type FinishTaskAttemptInput struct {
@@ -256,15 +259,37 @@ type FinishTaskAttemptInput struct {
} }
type FinishTaskSuccessInput struct { type FinishTaskSuccessInput struct {
TaskID string
ExecutionToken string
AttemptID string
Result map[string]any
Billings []any
RequestID string
ResolvedModel string
Usage map[string]any
Metrics map[string]any
BillingSummary map[string]any
FinalChargeAmount float64
FinalChargeAmountText string
BillingCurrency string
PricingSnapshot map[string]any
RequestFingerprint string
ResponseStartedAt time.Time
ResponseFinishedAt time.Time
ResponseDurationMS int64
}
type FinishTaskManualReviewInput struct {
TaskID string TaskID string
ExecutionToken string
AttemptID string
TaskStatus string
Code string
Message string
Result map[string]any Result map[string]any
Billings []any
RequestID string RequestID string
ResolvedModel string PricingSnapshot map[string]any
Usage map[string]any RequestFingerprint string
Metrics map[string]any
BillingSummary map[string]any
FinalChargeAmount float64
ResponseStartedAt time.Time ResponseStartedAt time.Time
ResponseFinishedAt time.Time ResponseFinishedAt time.Time
ResponseDurationMS int64 ResponseDurationMS int64
@@ -272,6 +297,7 @@ type FinishTaskSuccessInput struct {
type FinishTaskFailureInput struct { type FinishTaskFailureInput struct {
TaskID string TaskID string
ExecutionToken string
Code string Code string
Message string Message string
Result map[string]any Result map[string]any
+330 -54
View File
@@ -3,6 +3,7 @@ package store
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"strconv" "strconv"
"strings" "strings"
@@ -152,18 +153,71 @@ func nullableTaskListTime(value *time.Time) any {
return *value return *value
} }
func (s *Store) MarkTaskRunning(ctx context.Context, taskID string, modelType string, normalizedRequest map[string]any) error { func (s *Store) ClaimTaskExecution(ctx context.Context, taskID string, executionToken string, leaseTTL time.Duration) (GatewayTask, error) {
normalizedJSON, _ := json.Marshal(emptyObjectIfNil(normalizedRequest)) if leaseTTL <= 0 {
_, err := s.pool.Exec(ctx, ` leaseTTL = 5 * time.Minute
}
task, err := scanGatewayTask(s.pool.QueryRow(ctx, `
UPDATE gateway_tasks UPDATE gateway_tasks
SET status = 'running', SET status = 'running',
model_type = NULLIF($2::text, ''), execution_token = $2::uuid,
normalized_request = $3::jsonb, execution_lease_expires_at = now() + ($3::int * interval '1 second'),
locked_at = now(), locked_at = now(),
heartbeat_at = now(), heartbeat_at = now(),
updated_at = now() updated_at = now()
WHERE id = $1::uuid`, taskID, modelType, string(normalizedJSON)) WHERE id = $1::uuid
return err AND (
(status = 'queued' AND next_run_at <= now())
OR (status = 'running' AND (execution_lease_expires_at IS NULL OR execution_lease_expires_at <= now()))
)
RETURNING `+gatewayTaskColumns, taskID, executionToken, int(leaseTTL/time.Second)))
if errors.Is(err, pgx.ErrNoRows) {
return GatewayTask{}, ErrTaskExecutionLeaseUnavailable
}
return task, err
}
func (s *Store) RenewTaskExecutionLease(ctx context.Context, taskID string, executionToken string, leaseTTL time.Duration) error {
if leaseTTL <= 0 {
leaseTTL = 5 * time.Minute
}
tag, err := s.pool.Exec(ctx, `
UPDATE gateway_tasks
SET execution_lease_expires_at = now() + ($3::int * interval '1 second'),
heartbeat_at = now(),
updated_at = now()
WHERE id = $1::uuid
AND status = 'running'
AND execution_token = $2::uuid`, taskID, executionToken, int(leaseTTL/time.Second))
if err != nil {
return err
}
if tag.RowsAffected() != 1 {
return ErrTaskExecutionLeaseLost
}
return nil
}
func (s *Store) MarkTaskRunning(ctx context.Context, taskID string, executionToken string, modelType string, normalizedRequest map[string]any) error {
normalizedJSON, _ := json.Marshal(emptyObjectIfNil(normalizedRequest))
tag, err := s.pool.Exec(ctx, `
UPDATE gateway_tasks
SET status = 'running',
model_type = NULLIF($3::text, ''),
normalized_request = $4::jsonb,
locked_at = now(),
heartbeat_at = now(),
updated_at = now()
WHERE id = $1::uuid
AND status = 'running'
AND execution_token = $2::uuid`, taskID, executionToken, modelType, string(normalizedJSON))
if err != nil {
return err
}
if tag.RowsAffected() != 1 {
return ErrTaskExecutionLeaseLost
}
return nil
} }
func (s *Store) ClaimAsyncQueuedTask(ctx context.Context, workerID string) (GatewayTask, error) { func (s *Store) ClaimAsyncQueuedTask(ctx context.Context, workerID string) (GatewayTask, error) {
@@ -189,7 +243,7 @@ WHERE t.id = picked.task_id
RETURNING `+gatewayTaskColumns, workerID)) RETURNING `+gatewayTaskColumns, workerID))
} }
func (s *Store) RequeueTask(ctx context.Context, taskID string, delay time.Duration, queueKey string) (GatewayTask, error) { func (s *Store) RequeueTask(ctx context.Context, taskID string, executionToken string, delay time.Duration, queueKey string) (GatewayTask, error) {
if delay < time.Second { if delay < time.Second {
delay = time.Second delay = time.Second
} }
@@ -203,14 +257,18 @@ SET status = 'queued',
locked_by = NULL, locked_by = NULL,
locked_at = NULL, locked_at = NULL,
heartbeat_at = NULL, heartbeat_at = NULL,
next_run_at = $2::timestamptz, execution_token = NULL,
queue_key = COALESCE(NULLIF($3::text, ''), queue_key), execution_lease_expires_at = NULL,
next_run_at = $3::timestamptz,
queue_key = COALESCE(NULLIF($4::text, ''), queue_key),
error = NULL, error = NULL,
error_code = NULL, error_code = NULL,
error_message = NULL, error_message = NULL,
updated_at = now() updated_at = now()
WHERE id = $1::uuid WHERE id = $1::uuid
RETURNING `+gatewayTaskColumns, taskID, nextRunAt, strings.TrimSpace(queueKey))) AND status = 'running'
AND execution_token = $2::uuid
RETURNING `+gatewayTaskColumns, taskID, executionToken, nextRunAt, strings.TrimSpace(queueKey)))
} }
func (s *Store) SetTaskRiverJobID(ctx context.Context, taskID string, riverJobID int64) error { func (s *Store) SetTaskRiverJobID(ctx context.Context, taskID string, riverJobID int64) error {
@@ -225,25 +283,32 @@ WHERE id = $1::uuid`, taskID, riverJobID)
return err return err
} }
func (s *Store) SetTaskRemoteTask(ctx context.Context, taskID string, attemptID string, remoteTaskID string, payload map[string]any) error { func (s *Store) SetTaskRemoteTask(ctx context.Context, taskID string, executionToken string, attemptID string, remoteTaskID string, payload map[string]any) error {
payloadJSON, _ := json.Marshal(emptyObjectIfNil(payload)) payloadJSON, _ := json.Marshal(emptyObjectIfNil(payload))
return pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { return pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error {
if _, err := tx.Exec(ctx, ` tag, err := tx.Exec(ctx, `
UPDATE gateway_tasks UPDATE gateway_tasks
SET remote_task_id = NULLIF($2::text, ''), SET remote_task_id = NULLIF($3::text, ''),
remote_task_payload = $3::jsonb, remote_task_payload = $4::jsonb,
updated_at = now() updated_at = now()
WHERE id = $1::uuid`, WHERE id = $1::uuid
AND status = 'running'
AND execution_token = $2::uuid`,
taskID, taskID,
executionToken,
remoteTaskID, remoteTaskID,
string(payloadJSON), string(payloadJSON),
); err != nil { )
if err != nil {
return err return err
} }
if tag.RowsAffected() != 1 {
return ErrTaskExecutionLeaseLost
}
if strings.TrimSpace(attemptID) == "" { if strings.TrimSpace(attemptID) == "" {
return nil return nil
} }
_, err := tx.Exec(ctx, ` _, err = tx.Exec(ctx, `
UPDATE gateway_task_attempts UPDATE gateway_task_attempts
SET remote_task_id = NULLIF($2::text, ''), SET remote_task_id = NULLIF($2::text, ''),
response_snapshot = COALESCE(response_snapshot, '{}'::jsonb) || jsonb_build_object('remote_task_payload', $3::jsonb) response_snapshot = COALESCE(response_snapshot, '{}'::jsonb) || jsonb_build_object('remote_task_payload', $3::jsonb)
@@ -256,6 +321,20 @@ WHERE id = $1::uuid`,
}) })
} }
func (s *Store) SetAttemptUpstreamSubmissionStatus(ctx context.Context, attemptID string, status string) error {
switch status {
case "not_submitted", "submitting", "response_received":
default:
return fmt.Errorf("invalid upstream submission status %q", status)
}
_, err := s.pool.Exec(ctx, `
UPDATE gateway_task_attempts
SET upstream_submission_status = $2,
upstream_submission_updated_at = now()
WHERE id = $1::uuid`, attemptID, status)
return err
}
func (s *Store) CancelQueuedTask(ctx context.Context, taskID string, message string) (GatewayTask, bool, error) { func (s *Store) CancelQueuedTask(ctx context.Context, taskID string, message string) (GatewayTask, bool, error) {
message = strings.TrimSpace(message) message = strings.TrimSpace(message)
if message == "" { if message == "" {
@@ -293,10 +372,13 @@ func (s *Store) ListRecoverableAsyncTasks(ctx context.Context, limit int) ([]Asy
limit = 500 limit = 500
} }
rows, err := s.pool.Query(ctx, ` rows, err := s.pool.Query(ctx, `
SELECT id::text, priority SELECT id::text, priority, next_run_at
FROM gateway_tasks FROM gateway_tasks
WHERE async_mode = true WHERE async_mode = true
AND status IN ('queued', 'running') AND (
status = 'queued'
OR (status = 'running' AND (execution_lease_expires_at IS NULL OR execution_lease_expires_at <= now()))
)
ORDER BY priority ASC, created_at ASC ORDER BY priority ASC, created_at ASC
LIMIT $1`, limit) LIMIT $1`, limit)
if err != nil { if err != nil {
@@ -306,7 +388,7 @@ LIMIT $1`, limit)
items := make([]AsyncTaskQueueItem, 0) items := make([]AsyncTaskQueueItem, 0)
for rows.Next() { for rows.Next() {
var item AsyncTaskQueueItem var item AsyncTaskQueueItem
if err := rows.Scan(&item.ID, &item.Priority); err != nil { if err := rows.Scan(&item.ID, &item.Priority, &item.NextRunAt); err != nil {
return nil, err return nil, err
} }
items = append(items, item) items = append(items, item)
@@ -320,6 +402,7 @@ LIMIT $1`, limit)
func (s *Store) CreateTaskAttempt(ctx context.Context, input CreateTaskAttemptInput) (string, error) { func (s *Store) CreateTaskAttempt(ctx context.Context, input CreateTaskAttemptInput) (string, error) {
requestJSON, _ := json.Marshal(emptyObjectIfNil(input.RequestSnapshot)) requestJSON, _ := json.Marshal(emptyObjectIfNil(input.RequestSnapshot))
metricsJSON, _ := json.Marshal(emptyObjectIfNil(input.Metrics)) metricsJSON, _ := json.Marshal(emptyObjectIfNil(input.Metrics))
pricingJSON, _ := json.Marshal(emptyObjectIfNil(input.PricingSnapshot))
tx, err := s.pool.Begin(ctx) tx, err := s.pool.Begin(ctx)
if err != nil { if err != nil {
return "", err return "", err
@@ -330,11 +413,13 @@ func (s *Store) CreateTaskAttempt(ctx context.Context, input CreateTaskAttemptIn
err = tx.QueryRow(ctx, ` err = tx.QueryRow(ctx, `
INSERT INTO gateway_task_attempts ( INSERT INTO gateway_task_attempts (
task_id, attempt_no, platform_id, platform_model_id, client_id, queue_key, task_id, attempt_no, platform_id, platform_model_id, client_id, queue_key,
status, simulated, request_snapshot, metrics status, simulated, request_snapshot, metrics, pricing_snapshot, request_fingerprint,
upstream_submission_status, upstream_submission_updated_at
) )
VALUES ( VALUES (
$1::uuid, $2::int, NULLIF($3::text, '')::uuid, NULLIF($4::text, '')::uuid, NULLIF($5::text, ''), $6, $1::uuid, $2::int, NULLIF($3::text, '')::uuid, NULLIF($4::text, '')::uuid, NULLIF($5::text, ''), $6,
$7, $8, $9::jsonb, $10::jsonb $7, $8, $9::jsonb, $10::jsonb, $11::jsonb, NULLIF($12, ''),
'not_submitted', now()
) )
RETURNING id::text`, RETURNING id::text`,
input.TaskID, input.TaskID,
@@ -347,6 +432,8 @@ RETURNING id::text`,
input.Simulated, input.Simulated,
string(requestJSON), string(requestJSON),
string(metricsJSON), string(metricsJSON),
string(pricingJSON),
input.RequestFingerprint,
).Scan(&attemptID) ).Scan(&attemptID)
if err != nil { if err != nil {
return "", err return "", err
@@ -489,6 +576,8 @@ SELECT a.id::text, a.task_id::text, a.attempt_no,
a.request_snapshot, COALESCE(a.response_snapshot, '{}'::jsonb), a.request_snapshot, COALESCE(a.response_snapshot, '{}'::jsonb),
COALESCE(a.response_started_at::text, ''), COALESCE(a.response_finished_at::text, ''), COALESCE(a.response_started_at::text, ''), COALESCE(a.response_finished_at::text, ''),
COALESCE(a.response_duration_ms, 0), COALESCE(a.error_code, ''), COALESCE(a.error_message, ''), COALESCE(a.response_duration_ms, 0), COALESCE(a.error_code, ''), COALESCE(a.error_message, ''),
COALESCE(a.pricing_snapshot, '{}'::jsonb), COALESCE(a.request_fingerprint, ''),
a.upstream_submission_status, COALESCE(a.upstream_submission_updated_at::text, ''),
a.started_at, COALESCE(a.finished_at::text, '') a.started_at, COALESCE(a.finished_at::text, '')
FROM gateway_task_attempts a FROM gateway_task_attempts a
LEFT JOIN integration_platforms p ON p.id = a.platform_id LEFT JOIN integration_platforms p ON p.id = a.platform_id
@@ -518,6 +607,7 @@ func scanTaskAttempt(scanner taskScanner) (TaskAttempt, error) {
var metricsBytes []byte var metricsBytes []byte
var requestBytes []byte var requestBytes []byte
var responseBytes []byte var responseBytes []byte
var pricingSnapshotBytes []byte
if err := scanner.Scan( if err := scanner.Scan(
&item.ID, &item.ID,
&item.TaskID, &item.TaskID,
@@ -544,6 +634,10 @@ func scanTaskAttempt(scanner taskScanner) (TaskAttempt, error) {
&item.ResponseDurationMS, &item.ResponseDurationMS,
&item.ErrorCode, &item.ErrorCode,
&item.ErrorMessage, &item.ErrorMessage,
&pricingSnapshotBytes,
&item.RequestFingerprint,
&item.UpstreamSubmissionStatus,
&item.UpstreamSubmissionUpdatedAt,
&item.StartedAt, &item.StartedAt,
&item.FinishedAt, &item.FinishedAt,
); err != nil { ); err != nil {
@@ -553,6 +647,7 @@ func scanTaskAttempt(scanner taskScanner) (TaskAttempt, error) {
item.Metrics = decodeObject(metricsBytes) item.Metrics = decodeObject(metricsBytes)
item.RequestSnapshot = decodeObject(requestBytes) item.RequestSnapshot = decodeObject(requestBytes)
item.ResponseSnapshot = decodeObject(responseBytes) item.ResponseSnapshot = decodeObject(responseBytes)
item.PricingSnapshot = decodeObject(pricingSnapshotBytes)
enrichTaskAttemptFromMetrics(&item) enrichTaskAttemptFromMetrics(&item)
return item, nil return item, nil
} }
@@ -670,7 +765,41 @@ func (s *Store) FinishTaskSuccess(ctx context.Context, input FinishTaskSuccessIn
usageJSON, _ := json.Marshal(emptyObjectIfNil(input.Usage)) usageJSON, _ := json.Marshal(emptyObjectIfNil(input.Usage))
metricsJSON, _ := json.Marshal(emptyObjectIfNil(input.Metrics)) metricsJSON, _ := json.Marshal(emptyObjectIfNil(input.Metrics))
billingSummaryJSON, _ := json.Marshal(emptyObjectIfNil(input.BillingSummary)) billingSummaryJSON, _ := json.Marshal(emptyObjectIfNil(input.BillingSummary))
if _, err := s.pool.Exec(ctx, ` pricingSnapshotJSON, _ := json.Marshal(emptyObjectIfNil(input.PricingSnapshot))
finalChargeAmount := strings.TrimSpace(input.FinalChargeAmountText)
if finalChargeAmount == "" {
finalChargeAmount = strconv.FormatFloat(input.FinalChargeAmount, 'f', 9, 64)
}
currency := normalizeWalletCurrency(input.BillingCurrency)
err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error {
if strings.TrimSpace(input.AttemptID) != "" {
if _, err := tx.Exec(ctx, `
UPDATE gateway_task_attempts
SET status = 'succeeded',
retryable = false,
request_id = NULLIF($2, ''),
usage = $3::jsonb,
metrics = $4::jsonb,
response_snapshot = $5::jsonb,
pricing_snapshot = $6::jsonb,
request_fingerprint = NULLIF($7, ''),
upstream_submission_status = 'response_received',
upstream_submission_updated_at = now(),
response_started_at = $8::timestamptz,
response_finished_at = $9::timestamptz,
response_duration_ms = $10,
error_code = NULL,
error_message = NULL,
finished_at = now()
WHERE id = $1::uuid`,
input.AttemptID, input.RequestID, string(usageJSON), string(metricsJSON),
string(resultJSON), string(pricingSnapshotJSON), input.RequestFingerprint,
nullableTime(input.ResponseStartedAt), nullableTime(input.ResponseFinishedAt), input.ResponseDurationMS,
); err != nil {
return err
}
}
tag, err := tx.Exec(ctx, `
UPDATE gateway_tasks UPDATE gateway_tasks
SET status = 'succeeded', SET status = 'succeeded',
result = $2::jsonb, result = $2::jsonb,
@@ -680,32 +809,142 @@ SET status = 'succeeded',
usage = $6::jsonb, usage = $6::jsonb,
metrics = $7::jsonb, metrics = $7::jsonb,
billing_summary = $8::jsonb, billing_summary = $8::jsonb,
final_charge_amount = $9, final_charge_amount = $9::numeric,
response_started_at = $10::timestamptz, billing_version = 'effective-pricing-v2',
response_finished_at = $11::timestamptz, billing_status = CASE
response_duration_ms = $12, WHEN run_mode = 'production' AND gateway_user_id IS NOT NULL THEN 'pending'
ELSE 'not_required'
END,
billing_currency = $10,
pricing_snapshot = $11::jsonb,
request_fingerprint = NULLIF($12, ''),
billing_updated_at = now(),
response_started_at = $13::timestamptz,
response_finished_at = $14::timestamptz,
response_duration_ms = $15,
error = NULL, error = NULL,
error_code = NULL, error_code = NULL,
error_message = NULL, error_message = NULL,
locked_by = NULL, locked_by = NULL,
locked_at = NULL, locked_at = NULL,
heartbeat_at = NULL, heartbeat_at = NULL,
execution_token = NULL,
execution_lease_expires_at = NULL,
finished_at = now(), finished_at = now(),
updated_at = now() updated_at = now()
WHERE id = $1::uuid`, WHERE id = $1::uuid
input.TaskID, AND status = 'running'
string(resultJSON), AND execution_token = $16::uuid`,
string(billingsJSON), input.TaskID,
input.RequestID, string(resultJSON),
input.ResolvedModel, string(billingsJSON),
string(usageJSON), input.RequestID,
string(metricsJSON), input.ResolvedModel,
string(billingSummaryJSON), string(usageJSON),
input.FinalChargeAmount, string(metricsJSON),
nullableTime(input.ResponseStartedAt), string(billingSummaryJSON),
nullableTime(input.ResponseFinishedAt), finalChargeAmount,
input.ResponseDurationMS, currency,
); err != nil { string(pricingSnapshotJSON),
input.RequestFingerprint,
nullableTime(input.ResponseStartedAt),
nullableTime(input.ResponseFinishedAt),
input.ResponseDurationMS,
input.ExecutionToken,
)
if err != nil {
return err
}
if tag.RowsAffected() != 1 {
return ErrTaskExecutionLeaseLost
}
payloadJSON, _ := json.Marshal(map[string]any{"taskId": input.TaskID, "pricingVersion": "effective-pricing-v2"})
_, err = tx.Exec(ctx, `
INSERT INTO settlement_outbox (
task_id, event_type, action, amount, currency, pricing_snapshot, payload, status, next_attempt_at
)
SELECT id, 'task.billing.settle', 'settle', $2::numeric, $3, $4::jsonb, $5::jsonb, 'pending', now()
FROM gateway_tasks
WHERE id = $1::uuid
AND run_mode = 'production'
AND gateway_user_id IS NOT NULL
ON CONFLICT (task_id, event_type) DO NOTHING`,
input.TaskID, finalChargeAmount, currency, string(pricingSnapshotJSON), string(payloadJSON))
return err
})
if err != nil {
return GatewayTask{}, err
}
return s.GetTask(ctx, input.TaskID)
}
func (s *Store) FinishTaskManualReview(ctx context.Context, input FinishTaskManualReviewInput) (GatewayTask, error) {
status := strings.TrimSpace(input.TaskStatus)
if status != "succeeded" && status != "failed" {
return GatewayTask{}, fmt.Errorf("manual review task status must be succeeded or failed")
}
resultJSON, _ := json.Marshal(emptyObjectIfNil(input.Result))
pricingSnapshotJSON, _ := json.Marshal(emptyObjectIfNil(input.PricingSnapshot))
err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error {
if strings.TrimSpace(input.AttemptID) != "" {
attemptStatus := "failed"
if status == "succeeded" {
attemptStatus = "succeeded"
}
if _, err := tx.Exec(ctx, `
UPDATE gateway_task_attempts
SET status = $2,
request_id = NULLIF($3, ''),
error_code = CASE WHEN $2 = 'failed' THEN NULLIF($4, '') ELSE NULL END,
error_message = CASE WHEN $2 = 'failed' THEN NULLIF($5, '') ELSE NULL END,
upstream_submission_status = CASE WHEN $2 = 'succeeded' THEN 'response_received' ELSE upstream_submission_status END,
upstream_submission_updated_at = now(),
response_started_at = $6::timestamptz,
response_finished_at = $7::timestamptz,
response_duration_ms = $8,
finished_at = now(),
updated_at = now()
WHERE id = $1::uuid`, input.AttemptID, attemptStatus, input.RequestID, input.Code, input.Message,
nullableTime(input.ResponseStartedAt), nullableTime(input.ResponseFinishedAt), input.ResponseDurationMS); err != nil {
return err
}
}
tag, err := tx.Exec(ctx, `
UPDATE gateway_tasks
SET status = $2,
result = $3::jsonb,
request_id = NULLIF($4, ''),
error = CASE WHEN $2 = 'failed' THEN NULLIF($6, '') ELSE NULL END,
error_code = NULLIF($5, ''),
error_message = NULLIF($6, ''),
billing_status = 'manual_review',
pricing_snapshot = CASE WHEN $7::jsonb = '{}'::jsonb THEN pricing_snapshot ELSE $7::jsonb END,
request_fingerprint = COALESCE(NULLIF($8, ''), request_fingerprint),
billing_updated_at = now(),
response_started_at = $9::timestamptz,
response_finished_at = $10::timestamptz,
response_duration_ms = $11,
locked_by = NULL,
locked_at = NULL,
heartbeat_at = NULL,
execution_token = NULL,
execution_lease_expires_at = NULL,
finished_at = now(),
updated_at = now()
WHERE id = $1::uuid
AND status = 'running'
AND execution_token = $12::uuid`, input.TaskID, status, string(resultJSON), input.RequestID, input.Code, input.Message,
string(pricingSnapshotJSON), input.RequestFingerprint, nullableTime(input.ResponseStartedAt),
nullableTime(input.ResponseFinishedAt), input.ResponseDurationMS, input.ExecutionToken)
if err != nil {
return err
}
if tag.RowsAffected() != 1 {
return ErrTaskExecutionLeaseLost
}
return nil
})
if err != nil {
return GatewayTask{}, err return GatewayTask{}, err
} }
return s.GetTask(ctx, input.TaskID) return s.GetTask(ctx, input.TaskID)
@@ -743,6 +982,9 @@ ON CONFLICT (gateway_user_id, currency) DO NOTHING`,
task.GatewayTenantID, task.GatewayUserID, task.TenantID, task.TenantKey, task.UserID, currency); err != nil { task.GatewayTenantID, task.GatewayUserID, task.TenantID, task.TenantKey, task.UserID, currency); err != nil {
return err return err
} }
if err := ensureWalletAccountAuditGuard(ctx, tx, task.GatewayUserID, currency); err != nil {
return err
}
var exists bool var exists bool
var accountID string var accountID string
var balanceBefore float64 var balanceBefore float64
@@ -859,7 +1101,8 @@ func taskBillingString(value any) string {
func (s *Store) FinishTaskFailure(ctx context.Context, input FinishTaskFailureInput) (GatewayTask, error) { func (s *Store) FinishTaskFailure(ctx context.Context, input FinishTaskFailureInput) (GatewayTask, error) {
metricsJSON, _ := json.Marshal(emptyObjectIfNil(input.Metrics)) metricsJSON, _ := json.Marshal(emptyObjectIfNil(input.Metrics))
resultJSON, _ := json.Marshal(emptyObjectIfNil(input.Result)) resultJSON, _ := json.Marshal(emptyObjectIfNil(input.Result))
if _, err := s.pool.Exec(ctx, ` err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error {
tag, err := tx.Exec(ctx, `
UPDATE gateway_tasks UPDATE gateway_tasks
SET status = 'failed', SET status = 'failed',
error = NULLIF($2::text, ''), error = NULLIF($2::text, ''),
@@ -871,22 +1114,55 @@ func (s *Store) FinishTaskFailure(ctx context.Context, input FinishTaskFailureIn
response_finished_at = $7::timestamptz, response_finished_at = $7::timestamptz,
response_duration_ms = $8, response_duration_ms = $8,
result = $9::jsonb, result = $9::jsonb,
billing_status = CASE
WHEN run_mode <> 'production' OR gateway_user_id IS NULL THEN 'not_required'
WHEN reservation_amount > 0 THEN 'pending'
ELSE 'released'
END,
billing_updated_at = now(),
locked_by = NULL, locked_by = NULL,
locked_at = NULL, locked_at = NULL,
heartbeat_at = NULL, heartbeat_at = NULL,
execution_token = NULL,
execution_lease_expires_at = NULL,
finished_at = now(), finished_at = now(),
updated_at = now() updated_at = now()
WHERE id = $1::uuid`, WHERE id = $1::uuid
input.TaskID, AND status = 'running'
input.Message, AND execution_token = $10::uuid`,
input.Code, input.TaskID,
input.RequestID, input.Message,
string(metricsJSON), input.Code,
nullableTime(input.ResponseStartedAt), input.RequestID,
nullableTime(input.ResponseFinishedAt), string(metricsJSON),
input.ResponseDurationMS, nullableTime(input.ResponseStartedAt),
string(resultJSON), nullableTime(input.ResponseFinishedAt),
); err != nil { input.ResponseDurationMS,
string(resultJSON),
input.ExecutionToken,
)
if err != nil {
return err
}
if tag.RowsAffected() != 1 {
return ErrTaskExecutionLeaseLost
}
payloadJSON, _ := json.Marshal(map[string]any{"taskId": input.TaskID, "reason": input.Code})
_, err = tx.Exec(ctx, `
INSERT INTO settlement_outbox (
task_id, event_type, action, amount, currency, pricing_snapshot, payload, status, next_attempt_at
)
SELECT id, 'task.billing.release', 'release', reservation_amount, billing_currency,
pricing_snapshot, $2::jsonb, 'pending', now()
FROM gateway_tasks
WHERE id = $1::uuid
AND run_mode = 'production'
AND gateway_user_id IS NOT NULL
AND reservation_amount > 0
ON CONFLICT (task_id, event_type) DO NOTHING`, input.TaskID, string(payloadJSON))
return err
})
if err != nil {
return GatewayTask{}, err return GatewayTask{}, err
} }
return s.GetTask(ctx, input.TaskID) return s.GetTask(ctx, input.TaskID)
+334 -60
View File
@@ -3,6 +3,7 @@ package store
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"strconv" "strconv"
"strings" "strings"
@@ -82,6 +83,7 @@ type WalletBalanceAdjustmentInput struct {
GatewayUserID string `json:"gatewayUserId"` GatewayUserID string `json:"gatewayUserId"`
Currency string `json:"currency"` Currency string `json:"currency"`
Balance float64 `json:"balance"` Balance float64 `json:"balance"`
BalanceText string `json:"-"`
Reason string `json:"reason"` Reason string `json:"reason"`
IdempotencyKey string `json:"idempotencyKey"` IdempotencyKey string `json:"idempotencyKey"`
Metadata map[string]any `json:"metadata"` Metadata map[string]any `json:"metadata"`
@@ -91,6 +93,7 @@ type WalletRechargeInput struct {
GatewayUserID string `json:"gatewayUserId"` GatewayUserID string `json:"gatewayUserId"`
Currency string `json:"currency"` Currency string `json:"currency"`
Amount float64 `json:"amount"` Amount float64 `json:"amount"`
AmountText string `json:"-"`
Reason string `json:"reason"` Reason string `json:"reason"`
IdempotencyKey string `json:"idempotencyKey"` IdempotencyKey string `json:"idempotencyKey"`
Metadata map[string]any `json:"metadata"` Metadata map[string]any `json:"metadata"`
@@ -135,7 +138,7 @@ func (s *Store) WalletAvailability(ctx context.Context, user *auth.User, currenc
return result, nil return result, nil
} }
func (s *Store) ReserveTaskBilling(ctx context.Context, task GatewayTask, user *auth.User, billings []any) ([]WalletBillingReservation, error) { func (s *Store) ReserveTaskBilling(ctx context.Context, task GatewayTask, user *auth.User, billings []any, pricingSnapshots ...map[string]any) ([]WalletBillingReservation, error) {
gatewayUserID := taskGatewayUserID(task, user) gatewayUserID := taskGatewayUserID(task, user)
if gatewayUserID == "" { if gatewayUserID == "" {
return nil, nil return nil, nil
@@ -145,12 +148,21 @@ func (s *Store) ReserveTaskBilling(ctx context.Context, task GatewayTask, user *
return nil, fmt.Errorf("task id is required for wallet reservation") return nil, fmt.Errorf("task id is required for wallet reservation")
} }
pricingSnapshot := map[string]any{}
if len(pricingSnapshots) > 0 {
pricingSnapshot = emptyObjectIfNil(pricingSnapshots[0])
}
if exactAmount := walletString(pricingSnapshot["reservationAmount"]); exactAmount != "" {
return s.reserveTaskBillingExact(ctx, task, gatewayUserID, exactAmount, pricingSnapshot)
}
amounts := walletBillingAmounts(billings) amounts := walletBillingAmounts(billings)
if len(amounts) == 0 { if len(amounts) == 0 {
return nil, nil return nil, nil
} }
reservations := make([]WalletBillingReservation, 0, len(amounts)) reservations := make([]WalletBillingReservation, 0, len(amounts))
pricingSnapshotJSON, _ := json.Marshal(pricingSnapshot)
requestFingerprint := walletString(pricingSnapshot["requestFingerprint"])
err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error {
for currency, rawAmount := range amounts { for currency, rawAmount := range amounts {
amount := roundMoney(rawAmount) amount := roundMoney(rawAmount)
@@ -257,7 +269,30 @@ VALUES (
} }
reservations = append(reservations, reservation) reservations = append(reservations, reservation)
} }
return nil _, err := tx.Exec(ctx, `
UPDATE gateway_tasks task
SET reservation_amount = COALESCE((
SELECT SUM(reserve.amount)
FROM gateway_wallet_transactions reserve
WHERE reserve.reference_type = 'gateway_task'
AND reserve.reference_id = task.id::text
AND reserve.transaction_type = 'reserve'
AND NOT EXISTS (
SELECT 1
FROM gateway_wallet_transactions release
WHERE release.account_id = reserve.account_id
AND release.transaction_type = 'release'
AND release.idempotency_key = reserve.idempotency_key || ':release'
)
), 0),
billing_status = 'pending',
billing_currency = 'resource',
pricing_snapshot = $2::jsonb,
request_fingerprint = NULLIF($3, ''),
billing_updated_at = now(),
updated_at = now()
WHERE task.id = $1::uuid`, taskID, string(pricingSnapshotJSON), requestFingerprint)
return err
}) })
if err != nil { if err != nil {
return nil, err return nil, err
@@ -265,6 +300,106 @@ VALUES (
return reservations, err return reservations, err
} }
func (s *Store) reserveTaskBillingExact(ctx context.Context, task GatewayTask, gatewayUserID string, amount string, pricingSnapshot map[string]any) ([]WalletBillingReservation, error) {
taskID := strings.TrimSpace(task.ID)
currency := normalizeWalletCurrency(walletString(pricingSnapshot["currency"]))
if currency != "resource" {
return nil, fmt.Errorf("unsupported billing currency %q", currency)
}
pricingSnapshotJSON, _ := json.Marshal(pricingSnapshot)
requestFingerprint := walletString(pricingSnapshot["requestFingerprint"])
var reservations []WalletBillingReservation
err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error {
var positive bool
if err := tx.QueryRow(ctx, `SELECT $1::numeric(38, 9) > 0`, amount).Scan(&positive); err != nil {
return fmt.Errorf("invalid exact reservation amount: %w", err)
}
if !positive {
return nil
}
account, err := s.ensureWalletAccount(ctx, tx, gatewayUserID, currency)
if err != nil {
return err
}
locked, err := lockWalletAccount(ctx, tx, account.ID)
if err != nil {
return err
}
activeKey, activeAmount, err := activeWalletReservationExact(ctx, tx, locked.ID, taskID)
if err != nil {
return err
}
if activeAmount != "" {
presentationAmount, _ := strconv.ParseFloat(activeAmount, 64)
reservations = []WalletBillingReservation{{
TaskID: taskID, AccountID: locked.ID, GatewayUserID: gatewayUserID,
GatewayTenantID: firstNonEmpty(locked.GatewayTenantID, task.GatewayTenantID),
Currency: locked.Currency, Amount: presentationAmount, IdempotencyKey: activeKey,
}}
return nil
}
sequence, err := nextWalletReservationSequence(ctx, tx, locked.ID, taskID)
if err != nil {
return err
}
key := billingReservationIdempotencyKey(taskID, locked.Currency, sequence)
var balanceBefore string
var frozenBefore string
if err := tx.QueryRow(ctx, `
SELECT balance::text, frozen_balance::text
FROM gateway_wallet_accounts
WHERE id = $1::uuid`, locked.ID).Scan(&balanceBefore, &frozenBefore); err != nil {
return err
}
var frozenAfter string
err = tx.QueryRow(ctx, `
UPDATE gateway_wallet_accounts
SET frozen_balance = frozen_balance + $2::numeric(38, 9), updated_at = now()
WHERE id = $1::uuid
AND balance - frozen_balance >= $2::numeric(38, 9)
RETURNING frozen_balance::text`, locked.ID, amount).Scan(&frozenAfter)
if errors.Is(err, pgx.ErrNoRows) {
return ErrInsufficientWalletBalance
}
if err != nil {
return err
}
metadata, _ := json.Marshal(map[string]any{
"taskId": taskID, "reserved": amount, "balance": balanceBefore,
"frozenBefore": frozenBefore, "frozenAfter": frozenAfter,
"pricingVersion": walletString(pricingSnapshot["pricingVersion"]),
})
if _, err := tx.Exec(ctx, `
INSERT INTO gateway_wallet_transactions (
account_id, gateway_tenant_id, gateway_user_id, direction, transaction_type,
amount, balance_before, balance_after, idempotency_key, reference_type, reference_id, metadata
)
VALUES (
$1::uuid, NULLIF($2, '')::uuid, $3::uuid, 'debit', 'reserve',
$4::numeric(38, 9), $5::numeric, $5::numeric, $6, 'gateway_task', $7, $8::jsonb
)`, locked.ID, firstNonEmpty(locked.GatewayTenantID, task.GatewayTenantID), gatewayUserID,
amount, balanceBefore, key, taskID, string(metadata)); err != nil {
return err
}
if _, err := tx.Exec(ctx, `
UPDATE gateway_tasks
SET reservation_amount = $2::numeric(38, 9), billing_status = 'pending',
billing_currency = $3, pricing_snapshot = $4::jsonb,
request_fingerprint = NULLIF($5, ''), billing_updated_at = now(), updated_at = now()
WHERE id = $1::uuid`, taskID, amount, currency, string(pricingSnapshotJSON), requestFingerprint); err != nil {
return err
}
presentationAmount, _ := strconv.ParseFloat(amount, 64)
reservations = []WalletBillingReservation{{
TaskID: taskID, AccountID: locked.ID, GatewayUserID: gatewayUserID,
GatewayTenantID: firstNonEmpty(locked.GatewayTenantID, task.GatewayTenantID),
Currency: locked.Currency, Amount: presentationAmount, IdempotencyKey: key,
}}
return nil
})
return reservations, err
}
func (s *Store) ReleaseTaskBillingReservations(ctx context.Context, reservations []WalletBillingReservation, reason string) error { func (s *Store) ReleaseTaskBillingReservations(ctx context.Context, reservations []WalletBillingReservation, reason string) error {
if len(reservations) == 0 { if len(reservations) == 0 {
return nil return nil
@@ -274,10 +409,14 @@ func (s *Store) ReleaseTaskBillingReservations(ctx context.Context, reservations
reason = "task_not_settled" reason = "task_not_settled"
} }
return pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { return pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error {
taskIDs := map[string]struct{}{}
for _, reservation := range reservations { for _, reservation := range reservations {
if reservation.Amount <= 0 || strings.TrimSpace(reservation.AccountID) == "" { if strings.TrimSpace(reservation.AccountID) == "" {
continue continue
} }
if taskID := strings.TrimSpace(reservation.TaskID); taskID != "" {
taskIDs[taskID] = struct{}{}
}
reserveKey := strings.TrimSpace(reservation.IdempotencyKey) reserveKey := strings.TrimSpace(reservation.IdempotencyKey)
if reserveKey == "" { if reserveKey == "" {
reserveKey = billingReservationIdempotencyKey(reservation.TaskID, reservation.Currency, 1) reserveKey = billingReservationIdempotencyKey(reservation.TaskID, reservation.Currency, 1)
@@ -303,39 +442,44 @@ SELECT EXISTS (
if alreadyReleased { if alreadyReleased {
continue continue
} }
var storedReservedAmount float64 var storedReservedAmount string
if err := tx.QueryRow(ctx, ` if err := tx.QueryRow(ctx, `
SELECT COALESCE(( SELECT COALESCE((
SELECT amount::float8 SELECT amount::text
FROM gateway_wallet_transactions FROM gateway_wallet_transactions
WHERE account_id = $1::uuid WHERE account_id = $1::uuid
AND idempotency_key = $2 AND idempotency_key = $2
AND transaction_type = 'reserve' AND transaction_type = 'reserve'
LIMIT 1 LIMIT 1
), 0)::float8`, reservation.AccountID, reserveKey).Scan(&storedReservedAmount); err != nil { ), '0')`, reservation.AccountID, reserveKey).Scan(&storedReservedAmount); err != nil {
return err return err
} }
if storedReservedAmount <= 0 { var positive bool
if err := tx.QueryRow(ctx, `SELECT $1::numeric > 0`, storedReservedAmount).Scan(&positive); err != nil {
return err
}
if !positive {
continue continue
} }
amount := roundMoney(storedReservedAmount) var balanceBefore string
frozenAfter := roundMoney(locked.FrozenBalance - amount) var frozenBefore string
if frozenAfter < 0 { var frozenAfter string
frozenAfter = 0 if err := tx.QueryRow(ctx, `
}
if _, err := tx.Exec(ctx, `
UPDATE gateway_wallet_accounts UPDATE gateway_wallet_accounts
SET frozen_balance = $2, SET frozen_balance = frozen_balance - $2::numeric,
updated_at = now() updated_at = now()
WHERE id = $1::uuid`, locked.ID, frozenAfter); err != nil { WHERE id = $1::uuid
AND frozen_balance >= $2::numeric
RETURNING balance::text, (frozen_balance + $2::numeric)::text, frozen_balance::text`,
locked.ID, storedReservedAmount).Scan(&balanceBefore, &frozenBefore, &frozenAfter); err != nil {
return err return err
} }
metadata, _ := json.Marshal(map[string]any{ metadata, _ := json.Marshal(map[string]any{
"taskId": reservation.TaskID, "taskId": reservation.TaskID,
"reason": reason, "reason": reason,
"reserved": amount, "reserved": storedReservedAmount,
"frozenBefore": roundMoney(locked.FrozenBalance), "frozenBefore": frozenBefore,
"frozenAfter": frozenAfter, "frozenAfter": frozenAfter,
}) })
if _, err := tx.Exec(ctx, ` if _, err := tx.Exec(ctx, `
@@ -345,15 +489,14 @@ INSERT INTO gateway_wallet_transactions (
) )
VALUES ( VALUES (
$1::uuid, NULLIF($2, '')::uuid, $3::uuid, 'credit', 'release', $1::uuid, NULLIF($2, '')::uuid, $3::uuid, 'credit', 'release',
$4, $5, $6, $7, 'gateway_task', $8, $9::jsonb $4::numeric, $5::numeric, $5::numeric, $6, 'gateway_task', $7, $8::jsonb
) )
ON CONFLICT (account_id, idempotency_key) WHERE idempotency_key IS NOT NULL DO NOTHING`, ON CONFLICT (account_id, idempotency_key) WHERE idempotency_key IS NOT NULL DO NOTHING`,
locked.ID, locked.ID,
locked.GatewayTenantID, locked.GatewayTenantID,
locked.GatewayUserID, locked.GatewayUserID,
amount, storedReservedAmount,
roundMoney(locked.Balance), balanceBefore,
roundMoney(locked.Balance),
releaseKey, releaseKey,
reservation.TaskID, reservation.TaskID,
string(metadata), string(metadata),
@@ -361,6 +504,33 @@ ON CONFLICT (account_id, idempotency_key) WHERE idempotency_key IS NOT NULL DO N
return err return err
} }
} }
for taskID := range taskIDs {
if _, err := tx.Exec(ctx, `
UPDATE gateway_tasks task
SET reservation_amount = COALESCE((
SELECT SUM(reserve.amount)
FROM gateway_wallet_transactions reserve
WHERE reserve.reference_type = 'gateway_task'
AND reserve.reference_id = task.id::text
AND reserve.transaction_type = 'reserve'
AND NOT EXISTS (
SELECT 1 FROM gateway_wallet_transactions release
WHERE release.account_id = reserve.account_id
AND release.transaction_type = 'release'
AND release.idempotency_key = reserve.idempotency_key || ':release'
)
), 0),
billing_status = CASE
WHEN status IN ('failed', 'cancelled') THEN 'released'
WHEN status = 'succeeded' THEN billing_status
ELSE 'not_started'
END,
billing_updated_at = now(),
updated_at = now()
WHERE id = $1::uuid`, taskID); err != nil {
return err
}
}
return nil return nil
}) })
} }
@@ -588,8 +758,20 @@ func (s *Store) SetUserWalletBalanceTx(ctx context.Context, tx Tx, input WalletB
if input.GatewayUserID == "" { if input.GatewayUserID == "" {
return WalletAdjustmentResult{}, ErrLocalUserRequired return WalletAdjustmentResult{}, ErrLocalUserRequired
} }
if input.Balance < 0 { targetInput := strings.TrimSpace(input.BalanceText)
return WalletAdjustmentResult{}, fmt.Errorf("wallet balance cannot be negative") if targetInput == "" {
targetInput = strconv.FormatFloat(input.Balance, 'f', 9, 64)
}
targetBalance, err := canonicalWalletAmount(ctx, tx, targetInput)
if err != nil {
return WalletAdjustmentResult{}, err
}
var nonnegative bool
if err := tx.QueryRow(ctx, `SELECT $1::numeric >= 0`, targetBalance).Scan(&nonnegative); err != nil {
return WalletAdjustmentResult{}, err
}
if !nonnegative {
return WalletAdjustmentResult{}, ErrInvalidWalletAmount
} }
account, err := s.ensureWalletAccount(ctx, tx, input.GatewayUserID, input.Currency) account, err := s.ensureWalletAccount(ctx, tx, input.GatewayUserID, input.Currency)
if err != nil { if err != nil {
@@ -608,38 +790,68 @@ FOR UPDATE`, account.ID))
return WalletAdjustmentResult{}, err return WalletAdjustmentResult{}, err
} }
before := locked before := locked
nextBalance := roundMoney(input.Balance) var balanceBefore string
delta := roundMoney(nextBalance - locked.Balance) var frozenBalance string
if delta == 0 { if err := tx.QueryRow(ctx, `
SELECT balance::text, frozen_balance::text
FROM gateway_wallet_accounts
WHERE id = $1::uuid`, locked.ID).Scan(&balanceBefore, &frozenBalance); err != nil {
return WalletAdjustmentResult{}, err
}
var unchanged bool
var belowFrozen bool
if err := tx.QueryRow(ctx, `
SELECT $1::numeric = $2::numeric, $1::numeric < $3::numeric`,
targetBalance, balanceBefore, frozenBalance).Scan(&unchanged, &belowFrozen); err != nil {
return WalletAdjustmentResult{}, err
}
if unchanged {
return WalletAdjustmentResult{}, ErrWalletBalanceUnchanged return WalletAdjustmentResult{}, ErrWalletBalanceUnchanged
} }
if belowFrozen {
return WalletAdjustmentResult{}, ErrBalanceBelowFrozen
}
direction := "credit" direction := "credit"
amount := delta if strings.HasPrefix(balanceBefore, "-") {
if delta < 0 { return WalletAdjustmentResult{}, ErrInvalidWalletAmount
}
var targetGreater bool
if err := tx.QueryRow(ctx, `SELECT $1::numeric > $2::numeric`, targetBalance, balanceBefore).Scan(&targetGreater); err != nil {
return WalletAdjustmentResult{}, err
}
if !targetGreater {
direction = "debit" direction = "debit"
amount = -delta }
var amount string
if err := tx.QueryRow(ctx, `SELECT abs($1::numeric - $2::numeric)::text`, targetBalance, balanceBefore).Scan(&amount); err != nil {
return WalletAdjustmentResult{}, err
} }
reason := strings.TrimSpace(input.Reason) reason := strings.TrimSpace(input.Reason)
if reason == "" { if reason == "" {
reason = "后台余额调整" reason = "后台余额调整"
} }
if _, err := tx.Exec(ctx, ` tag, err := tx.Exec(ctx, `
UPDATE gateway_wallet_accounts UPDATE gateway_wallet_accounts
SET balance = $2, SET balance = $2::numeric,
total_recharged = total_recharged + CASE WHEN $3 = 'credit' THEN $4 ELSE 0 END, total_recharged = total_recharged + CASE WHEN $3 = 'credit' THEN $4::numeric ELSE 0 END,
updated_at = now() updated_at = now()
WHERE id = $1::uuid`, WHERE id = $1::uuid
AND $2::numeric >= frozen_balance`,
locked.ID, locked.ID,
nextBalance, targetBalance,
direction, direction,
amount, amount,
); err != nil { )
if err != nil {
return WalletAdjustmentResult{}, err return WalletAdjustmentResult{}, err
} }
if tag.RowsAffected() != 1 {
return WalletAdjustmentResult{}, ErrBalanceBelowFrozen
}
metadata := mergeObjects(input.Metadata, map[string]any{ metadata := mergeObjects(input.Metadata, map[string]any{
"reason": reason, "reason": reason,
"previousBalance": roundMoney(before.Balance), "previousBalance": balanceBefore,
"targetBalance": nextBalance, "targetBalance": targetBalance,
}) })
metadataJSON, _ := json.Marshal(emptyObjectIfNil(metadata)) metadataJSON, _ := json.Marshal(emptyObjectIfNil(metadata))
transaction, err := scanWalletTransaction(tx.QueryRow(ctx, ` transaction, err := scanWalletTransaction(tx.QueryRow(ctx, `
@@ -649,7 +861,7 @@ INSERT INTO gateway_wallet_transactions (
) )
VALUES ( VALUES (
$1::uuid, NULLIF($2, '')::uuid, $3::uuid, $4, 'admin_adjust', $1::uuid, NULLIF($2, '')::uuid, $3::uuid, $4, 'admin_adjust',
$5, $6, $7, NULLIF($8, ''), 'gateway_user', $9, $10::jsonb $5::numeric, $6::numeric, $7::numeric, NULLIF($8, ''), 'gateway_user', $9, $10::jsonb
) )
RETURNING id::text, account_id::text, COALESCE(gateway_tenant_id::text, ''), COALESCE(gateway_user_id::text, ''), RETURNING id::text, account_id::text, COALESCE(gateway_tenant_id::text, ''), COALESCE(gateway_user_id::text, ''),
direction, transaction_type, amount::float8, balance_before::float8, balance_after::float8, direction, transaction_type, amount::float8, balance_before::float8, balance_after::float8,
@@ -660,8 +872,8 @@ RETURNING id::text, account_id::text, COALESCE(gateway_tenant_id::text, ''), COA
locked.GatewayUserID, locked.GatewayUserID,
direction, direction,
amount, amount,
roundMoney(before.Balance), balanceBefore,
nextBalance, targetBalance,
strings.TrimSpace(input.IdempotencyKey), strings.TrimSpace(input.IdempotencyKey),
locked.GatewayUserID, locked.GatewayUserID,
string(metadataJSON), string(metadataJSON),
@@ -669,11 +881,10 @@ RETURNING id::text, account_id::text, COALESCE(gateway_tenant_id::text, ''), COA
if err != nil { if err != nil {
return WalletAdjustmentResult{}, err return WalletAdjustmentResult{}, err
} }
locked.Balance = nextBalance locked, err = s.ensureWalletAccount(ctx, tx, input.GatewayUserID, input.Currency)
if direction == "credit" { if err != nil {
locked.TotalRecharged = roundMoney(locked.TotalRecharged + amount) return WalletAdjustmentResult{}, err
} }
locked.UpdatedAt = time.Now()
return WalletAdjustmentResult{Account: locked, Before: before, Transaction: transaction}, nil return WalletAdjustmentResult{Account: locked, Before: before, Transaction: transaction}, nil
} }
@@ -695,9 +906,20 @@ func (s *Store) RechargeUserWalletBalanceTx(ctx context.Context, tx Tx, input Wa
if input.GatewayUserID == "" { if input.GatewayUserID == "" {
return WalletAdjustmentResult{}, ErrLocalUserRequired return WalletAdjustmentResult{}, ErrLocalUserRequired
} }
amount := roundMoney(input.Amount) amountInput := strings.TrimSpace(input.AmountText)
if amount <= 0 { if amountInput == "" {
return WalletAdjustmentResult{}, fmt.Errorf("wallet recharge amount must be positive") amountInput = strconv.FormatFloat(input.Amount, 'f', 9, 64)
}
amount, err := canonicalWalletAmount(ctx, tx, amountInput)
if err != nil {
return WalletAdjustmentResult{}, err
}
var positive bool
if err := tx.QueryRow(ctx, `SELECT $1::numeric > 0`, amount).Scan(&positive); err != nil {
return WalletAdjustmentResult{}, err
}
if !positive {
return WalletAdjustmentResult{}, ErrInvalidWalletAmount
} }
account, err := s.ensureWalletAccount(ctx, tx, input.GatewayUserID, input.Currency) account, err := s.ensureWalletAccount(ctx, tx, input.GatewayUserID, input.Currency)
if err != nil { if err != nil {
@@ -715,26 +937,30 @@ FOR UPDATE`, account.ID))
return WalletAdjustmentResult{}, err return WalletAdjustmentResult{}, err
} }
before := locked before := locked
nextBalance := roundMoney(locked.Balance + amount) var balanceBefore string
if err := tx.QueryRow(ctx, `SELECT balance::text FROM gateway_wallet_accounts WHERE id = $1::uuid`, locked.ID).Scan(&balanceBefore); err != nil {
return WalletAdjustmentResult{}, err
}
reason := strings.TrimSpace(input.Reason) reason := strings.TrimSpace(input.Reason)
if reason == "" { if reason == "" {
reason = "后台余额充值" reason = "后台余额充值"
} }
if _, err := tx.Exec(ctx, ` var nextBalance string
if err := tx.QueryRow(ctx, `
UPDATE gateway_wallet_accounts UPDATE gateway_wallet_accounts
SET balance = $2, SET balance = balance + $2::numeric,
total_recharged = total_recharged + $3, total_recharged = total_recharged + $2::numeric,
updated_at = now() updated_at = now()
WHERE id = $1::uuid`, WHERE id = $1::uuid
RETURNING balance::text`,
locked.ID, locked.ID,
nextBalance,
amount, amount,
); err != nil { ).Scan(&nextBalance); err != nil {
return WalletAdjustmentResult{}, err return WalletAdjustmentResult{}, err
} }
metadata := mergeObjects(input.Metadata, map[string]any{ metadata := mergeObjects(input.Metadata, map[string]any{
"reason": reason, "reason": reason,
"previousBalance": roundMoney(before.Balance), "previousBalance": balanceBefore,
"rechargeAmount": amount, "rechargeAmount": amount,
"targetBalance": nextBalance, "targetBalance": nextBalance,
}) })
@@ -746,7 +972,7 @@ INSERT INTO gateway_wallet_transactions (
) )
VALUES ( VALUES (
$1::uuid, NULLIF($2, '')::uuid, $3::uuid, 'credit', 'recharge', $1::uuid, NULLIF($2, '')::uuid, $3::uuid, 'credit', 'recharge',
$4, $5, $6, NULLIF($7, ''), 'gateway_user', $8, $9::jsonb $4::numeric, $5::numeric, $6::numeric, NULLIF($7, ''), 'gateway_user', $8, $9::jsonb
) )
RETURNING id::text, account_id::text, COALESCE(gateway_tenant_id::text, ''), COALESCE(gateway_user_id::text, ''), RETURNING id::text, account_id::text, COALESCE(gateway_tenant_id::text, ''), COALESCE(gateway_user_id::text, ''),
direction, transaction_type, amount::float8, balance_before::float8, balance_after::float8, direction, transaction_type, amount::float8, balance_before::float8, balance_after::float8,
@@ -756,7 +982,7 @@ RETURNING id::text, account_id::text, COALESCE(gateway_tenant_id::text, ''), COA
locked.GatewayTenantID, locked.GatewayTenantID,
locked.GatewayUserID, locked.GatewayUserID,
amount, amount,
roundMoney(before.Balance), balanceBefore,
nextBalance, nextBalance,
strings.TrimSpace(input.IdempotencyKey), strings.TrimSpace(input.IdempotencyKey),
locked.GatewayUserID, locked.GatewayUserID,
@@ -765,12 +991,21 @@ RETURNING id::text, account_id::text, COALESCE(gateway_tenant_id::text, ''), COA
if err != nil { if err != nil {
return WalletAdjustmentResult{}, err return WalletAdjustmentResult{}, err
} }
locked.Balance = nextBalance locked, err = s.ensureWalletAccount(ctx, tx, input.GatewayUserID, input.Currency)
locked.TotalRecharged = roundMoney(locked.TotalRecharged + amount) if err != nil {
locked.UpdatedAt = time.Now() return WalletAdjustmentResult{}, err
}
return WalletAdjustmentResult{Account: locked, Before: before, Transaction: transaction}, nil return WalletAdjustmentResult{Account: locked, Before: before, Transaction: transaction}, nil
} }
func canonicalWalletAmount(ctx context.Context, q Tx, value string) (string, error) {
var canonical string
if err := q.QueryRow(ctx, `SELECT $1::numeric(38, 9)::text`, strings.TrimSpace(value)).Scan(&canonical); err != nil {
return "", fmt.Errorf("%w: %v", ErrInvalidWalletAmount, err)
}
return canonical, nil
}
func (s *Store) ensureWalletAccount(ctx context.Context, q Tx, gatewayUserID string, currency string) (GatewayWalletAccount, error) { func (s *Store) ensureWalletAccount(ctx context.Context, q Tx, gatewayUserID string, currency string) (GatewayWalletAccount, error) {
currency = normalizeWalletCurrency(currency) currency = normalizeWalletCurrency(currency)
if _, err := q.Exec(ctx, ` if _, err := q.Exec(ctx, `
@@ -794,6 +1029,9 @@ WHERE gateway_wallet_accounts.gateway_tenant_id IS NULL
OR COALESCE(gateway_wallet_accounts.user_id, '') = ''`, gatewayUserID, currency); err != nil { OR COALESCE(gateway_wallet_accounts.user_id, '') = ''`, gatewayUserID, currency); err != nil {
return GatewayWalletAccount{}, err return GatewayWalletAccount{}, err
} }
if err := ensureWalletAccountAuditGuard(ctx, q, gatewayUserID, currency); err != nil {
return GatewayWalletAccount{}, err
}
account, err := scanWalletAccount(q.QueryRow(ctx, ` account, err := scanWalletAccount(q.QueryRow(ctx, `
SELECT id::text, COALESCE(gateway_tenant_id::text, ''), gateway_user_id::text, SELECT id::text, COALESCE(gateway_tenant_id::text, ''), gateway_user_id::text,
COALESCE(tenant_id, ''), COALESCE(tenant_key, ''), COALESCE(user_id, ''), COALESCE(tenant_id, ''), COALESCE(tenant_key, ''), COALESCE(user_id, ''),
@@ -811,6 +1049,17 @@ WHERE gateway_user_id = $1::uuid
return account, nil return account, nil
} }
func ensureWalletAccountAuditGuard(ctx context.Context, q Tx, gatewayUserID string, currency string) error {
_, err := q.Exec(ctx, `
INSERT INTO gateway_wallet_account_audit_guards (account_id)
SELECT id
FROM gateway_wallet_accounts
WHERE gateway_user_id = $1::uuid
AND currency = $2
ON CONFLICT (account_id) DO NOTHING`, gatewayUserID, normalizeWalletCurrency(currency))
return err
}
func lockWalletAccount(ctx context.Context, tx pgx.Tx, accountID string) (GatewayWalletAccount, error) { func lockWalletAccount(ctx context.Context, tx pgx.Tx, accountID string) (GatewayWalletAccount, error) {
return scanWalletAccount(tx.QueryRow(ctx, ` return scanWalletAccount(tx.QueryRow(ctx, `
SELECT id::text, COALESCE(gateway_tenant_id::text, ''), gateway_user_id::text, SELECT id::text, COALESCE(gateway_tenant_id::text, ''), gateway_user_id::text,
@@ -851,6 +1100,31 @@ LIMIT 1`, accountID, taskID).Scan(&key, &amount)
return key, roundMoney(amount), nil return key, roundMoney(amount), nil
} }
func activeWalletReservationExact(ctx context.Context, tx pgx.Tx, accountID string, taskID string) (string, string, error) {
var key string
var amount string
err := tx.QueryRow(ctx, `
SELECT COALESCE(t.idempotency_key, ''), t.amount::text
FROM gateway_wallet_transactions t
WHERE t.account_id = $1::uuid
AND t.reference_type = 'gateway_task'
AND t.reference_id = $2
AND t.transaction_type = 'reserve'
AND COALESCE(t.idempotency_key, '') <> ''
AND NOT EXISTS (
SELECT 1 FROM gateway_wallet_transactions release
WHERE release.account_id = t.account_id
AND release.transaction_type = 'release'
AND release.idempotency_key = t.idempotency_key || ':release'
)
ORDER BY t.created_at DESC
LIMIT 1`, accountID, taskID).Scan(&key, &amount)
if errors.Is(err, pgx.ErrNoRows) {
return "", "", nil
}
return key, amount, err
}
func nextWalletReservationSequence(ctx context.Context, tx pgx.Tx, accountID string, taskID string) (int, error) { func nextWalletReservationSequence(ctx context.Context, tx pgx.Tx, accountID string, taskID string) (int, error) {
var count int var count int
if err := tx.QueryRow(ctx, ` if err := tx.QueryRow(ctx, `
@@ -132,6 +132,10 @@ RETURNING id::text`, "wallet-reservation-user-"+suffix, "wallet_reservation_"+su
} }
t.Cleanup(func() { t.Cleanup(func() {
cleanupCtx := context.Background() cleanupCtx := context.Background()
_, _ = db.pool.Exec(cleanupCtx, `DELETE FROM gateway_wallet_transactions WHERE gateway_user_id = $1::uuid`, userID)
_, _ = db.pool.Exec(cleanupCtx, `
DELETE FROM gateway_wallet_account_audit_guards
WHERE account_id IN (SELECT id FROM gateway_wallet_accounts WHERE gateway_user_id = $1::uuid)`, userID)
_, _ = db.pool.Exec(cleanupCtx, `DELETE FROM gateway_users WHERE id = $1::uuid`, userID) _, _ = db.pool.Exec(cleanupCtx, `DELETE FROM gateway_users WHERE id = $1::uuid`, userID)
_, _ = db.pool.Exec(cleanupCtx, `DELETE FROM gateway_tenants WHERE id = $1::uuid`, tenantID) _, _ = db.pool.Exec(cleanupCtx, `DELETE FROM gateway_tenants WHERE id = $1::uuid`, tenantID)
}) })
@@ -1,16 +1,24 @@
ALTER TABLE gateway_tasks ALTER TABLE gateway_tasks
ADD COLUMN IF NOT EXISTS billing_version text NOT NULL DEFAULT 'effective-pricing-v2', ADD COLUMN IF NOT EXISTS billing_version text DEFAULT 'effective-pricing-v2',
ADD COLUMN IF NOT EXISTS billing_status text NOT NULL DEFAULT 'not_started', ADD COLUMN IF NOT EXISTS billing_status text DEFAULT 'not_started',
ADD COLUMN IF NOT EXISTS billing_currency text NOT NULL DEFAULT 'resource', ADD COLUMN IF NOT EXISTS billing_currency text DEFAULT 'resource',
ADD COLUMN IF NOT EXISTS pricing_snapshot jsonb NOT NULL DEFAULT '{}'::jsonb, ADD COLUMN IF NOT EXISTS pricing_snapshot jsonb DEFAULT '{}'::jsonb,
ADD COLUMN IF NOT EXISTS request_fingerprint text, ADD COLUMN IF NOT EXISTS request_fingerprint text,
ADD COLUMN IF NOT EXISTS idempotency_key_hash text, ADD COLUMN IF NOT EXISTS idempotency_key_hash text,
ADD COLUMN IF NOT EXISTS reservation_amount numeric(38, 9) NOT NULL DEFAULT 0, ADD COLUMN IF NOT EXISTS idempotency_request_hash text,
ADD COLUMN IF NOT EXISTS reservation_amount numeric(38, 9) DEFAULT 0,
ADD COLUMN IF NOT EXISTS execution_token uuid, ADD COLUMN IF NOT EXISTS execution_token uuid,
ADD COLUMN IF NOT EXISTS execution_lease_expires_at timestamptz, ADD COLUMN IF NOT EXISTS execution_lease_expires_at timestamptz,
ADD COLUMN IF NOT EXISTS billing_updated_at timestamptz, ADD COLUMN IF NOT EXISTS billing_updated_at timestamptz,
ADD COLUMN IF NOT EXISTS billing_settled_at timestamptz; ADD COLUMN IF NOT EXISTS billing_settled_at timestamptz;
UPDATE gateway_tasks
SET billing_version = COALESCE(billing_version, 'effective-pricing-v2'),
billing_status = COALESCE(billing_status, 'not_started'),
billing_currency = COALESCE(billing_currency, 'resource'),
pricing_snapshot = COALESCE(pricing_snapshot, '{}'::jsonb),
reservation_amount = COALESCE(reservation_amount, 0);
ALTER TABLE gateway_tasks ALTER TABLE gateway_tasks
ADD CONSTRAINT gateway_tasks_billing_status_v2_check CHECK ( ADD CONSTRAINT gateway_tasks_billing_status_v2_check CHECK (
billing_status IN ( billing_status IN (
@@ -32,28 +40,40 @@ CREATE INDEX IF NOT EXISTS idx_gateway_tasks_billing_status_v2
ON gateway_tasks(billing_status, billing_updated_at, created_at); ON gateway_tasks(billing_status, billing_updated_at, created_at);
ALTER TABLE gateway_task_attempts ALTER TABLE gateway_task_attempts
ADD COLUMN IF NOT EXISTS pricing_snapshot jsonb NOT NULL DEFAULT '{}'::jsonb, ADD COLUMN IF NOT EXISTS pricing_snapshot jsonb DEFAULT '{}'::jsonb,
ADD COLUMN IF NOT EXISTS request_fingerprint text, ADD COLUMN IF NOT EXISTS request_fingerprint text,
ADD COLUMN IF NOT EXISTS upstream_submission_status text NOT NULL DEFAULT 'not_submitted', ADD COLUMN IF NOT EXISTS upstream_submission_status text DEFAULT 'not_submitted',
ADD COLUMN IF NOT EXISTS upstream_submission_updated_at timestamptz; ADD COLUMN IF NOT EXISTS upstream_submission_updated_at timestamptz;
UPDATE gateway_task_attempts
SET pricing_snapshot = COALESCE(pricing_snapshot, '{}'::jsonb),
upstream_submission_status = COALESCE(upstream_submission_status, 'not_submitted');
ALTER TABLE gateway_task_attempts ALTER TABLE gateway_task_attempts
ADD CONSTRAINT gateway_task_attempts_submission_v2_check CHECK ( ADD CONSTRAINT gateway_task_attempts_submission_v2_check CHECK (
upstream_submission_status IN ('not_submitted', 'submitting', 'response_received') upstream_submission_status IN ('not_submitted', 'submitting', 'response_received')
) NOT VALID; ) NOT VALID;
ALTER TABLE settlement_outbox ALTER TABLE settlement_outbox
ADD COLUMN IF NOT EXISTS action text NOT NULL DEFAULT 'settle', ADD COLUMN IF NOT EXISTS action text DEFAULT 'settle',
ADD COLUMN IF NOT EXISTS amount numeric(38, 9) NOT NULL DEFAULT 0, ADD COLUMN IF NOT EXISTS amount numeric(38, 9) DEFAULT 0,
ADD COLUMN IF NOT EXISTS currency text NOT NULL DEFAULT 'resource', ADD COLUMN IF NOT EXISTS currency text DEFAULT 'resource',
ADD COLUMN IF NOT EXISTS pricing_snapshot jsonb NOT NULL DEFAULT '{}'::jsonb, ADD COLUMN IF NOT EXISTS pricing_snapshot jsonb DEFAULT '{}'::jsonb,
ADD COLUMN IF NOT EXISTS locked_by text, ADD COLUMN IF NOT EXISTS locked_by text,
ADD COLUMN IF NOT EXISTS lock_token uuid, ADD COLUMN IF NOT EXISTS lock_token uuid,
ADD COLUMN IF NOT EXISTS locked_at timestamptz, ADD COLUMN IF NOT EXISTS locked_at timestamptz,
ADD COLUMN IF NOT EXISTS last_error_code text, ADD COLUMN IF NOT EXISTS last_error_code text,
ADD COLUMN IF NOT EXISTS last_error_message text, ADD COLUMN IF NOT EXISTS last_error_message text,
ADD COLUMN IF NOT EXISTS completed_at timestamptz, ADD COLUMN IF NOT EXISTS completed_at timestamptz,
ADD COLUMN IF NOT EXISTS manual_review_reason text; ADD COLUMN IF NOT EXISTS manual_review_reason text,
ADD COLUMN IF NOT EXISTS retry_idempotency_key_hash text,
ADD COLUMN IF NOT EXISTS retry_requested_at timestamptz;
UPDATE settlement_outbox
SET action = COALESCE(action, 'settle'),
amount = COALESCE(amount, 0),
currency = COALESCE(currency, 'resource'),
pricing_snapshot = COALESCE(pricing_snapshot, '{}'::jsonb);
ALTER TABLE settlement_outbox ALTER TABLE settlement_outbox
ADD CONSTRAINT settlement_outbox_action_v2_check CHECK (action IN ('settle', 'release')) NOT VALID, ADD CONSTRAINT settlement_outbox_action_v2_check CHECK (action IN ('settle', 'release')) NOT VALID,
@@ -71,7 +91,9 @@ CREATE INDEX IF NOT EXISTS idx_settlement_outbox_stale_v2
WHERE status = 'processing'; WHERE status = 'processing';
ALTER TABLE model_pricing_rules ALTER TABLE model_pricing_rules
ADD COLUMN IF NOT EXISTS is_free boolean NOT NULL DEFAULT false; ADD COLUMN IF NOT EXISTS is_free boolean DEFAULT false;
UPDATE model_pricing_rules SET is_free = false WHERE is_free IS NULL;
ALTER TABLE model_pricing_rules ALTER TABLE model_pricing_rules
ADD CONSTRAINT model_pricing_rules_explicit_free_v2_check CHECK ( ADD CONSTRAINT model_pricing_rules_explicit_free_v2_check CHECK (
@@ -82,9 +104,18 @@ ALTER TABLE gateway_wallet_accounts
ADD CONSTRAINT gateway_wallet_nonnegative_frozen_v2_check CHECK (frozen_balance >= 0) NOT VALID, ADD CONSTRAINT gateway_wallet_nonnegative_frozen_v2_check CHECK (frozen_balance >= 0) NOT VALID,
ADD CONSTRAINT gateway_wallet_balance_covers_frozen_v2_check CHECK (balance >= frozen_balance) NOT VALID; ADD CONSTRAINT gateway_wallet_balance_covers_frozen_v2_check CHECK (balance >= frozen_balance) NOT VALID;
CREATE TABLE IF NOT EXISTS gateway_wallet_account_audit_guards (
account_id uuid PRIMARY KEY REFERENCES gateway_wallet_accounts(id) ON DELETE RESTRICT,
created_at timestamptz NOT NULL DEFAULT now()
);
INSERT INTO gateway_wallet_account_audit_guards (account_id)
SELECT id FROM gateway_wallet_accounts
ON CONFLICT (account_id) DO NOTHING;
ALTER TABLE gateway_wallet_transactions ALTER TABLE gateway_wallet_transactions
ADD CONSTRAINT gateway_wallet_transactions_account_restrict_v2 ADD CONSTRAINT gateway_wallet_transactions_audit_guard_v2
FOREIGN KEY (account_id) REFERENCES gateway_wallet_accounts(id) ON DELETE RESTRICT NOT VALID; FOREIGN KEY (account_id) REFERENCES gateway_wallet_account_audit_guards(account_id) ON DELETE RESTRICT NOT VALID;
UPDATE gateway_tasks task UPDATE gateway_tasks task
SET billing_status = 'not_required', SET billing_status = 'not_required',