From 80a229bce5d386cacaa8ea3dfb3d82495644db9d Mon Sep 17 00:00:00 2001 From: shaw Date: Tue, 7 Jul 2026 18:53:56 +0800 Subject: [PATCH 1/2] =?UTF-8?q?fix(batch-image):=20=E4=BF=AE=E5=A4=8D?= =?UTF-8?q?=E5=AE=A1=E8=AE=A1=E5=8F=91=E7=8E=B0=E7=9A=84=E8=AE=A1=E8=B4=B9?= =?UTF-8?q?=E6=AD=BB=E9=94=81=E3=80=81=E7=8A=B6=E6=80=81=E6=9C=BA=E4=B8=8E?= =?UTF-8?q?=E9=98=9F=E5=88=97=E5=8E=9F=E5=AD=90=E6=80=A7=E7=BC=BA=E9=99=B7?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 修复 PR #3768 批量图像 MVP 合并后审计报告中的全部问题: 结算与计费(高危): - 所有 SETTLEMENT_* 失败(超冻结/计数非法/manifest 冲突/定价缺失/扣费失败) 统一计入 retry_count 并在耗尽时释放冻结转 failed,消灭 settling 无限 requeue 导致的冻结余额永久锁死 - 耗尽出口的释放指纹统一为 RequestHash,与 processor/Cancel/recovery 一致; release 遇同 request id 指纹冲突视为幂等成功,治愈历史毒消息 - 管理端校验 hold_multiplier >= discount_multiplier,定价快照对存量脏数据钳制 - 释放前校验 per-job hold claim(dedup+归档表),杜绝幻影释放 索引对账(高危): - provider 输出与提交 custom_id 集对账:未知条目丢弃并记事件, 漏项补 PROVIDER_RESULT_MISSING 失败行,保证 success+fail == item_count 提交与恢复(高危): - 提交前转 uploading 并在 provider.Submit 期间心跳刷新 updated_at; 恢复扫描改为原子复核(FailStaleUnsubmittedBatchImageJob), 消灭慢提交被误杀退款而上游任务照常计费的孤儿场景 - 上游任务创建成功但本地状态推进失败时,尽力取消上游并清理输入 - recovery 释放失败时入队交由 worker releaseTerminalHold 兜底重试 队列与并发(中危): - Enqueue(SetNX+LPush)与 Reserve(BRPop+ZAdd)均改为 Lua 原子脚本, 消灭崩溃窗口导致 job 脱离队列、被 7 天 inflight 键锁死 - 锁冲突按 LockConflictDelay 重新入队(原直接丢弃需等 10 分钟 stale 恢复) - 处理期间心跳:active zset 续期(ZAddXX 防幽灵成员)+ 锁 TTL 续期 - ReplaceBatchImageItemsForJob 增加 indexing 状态守卫,防掉队 worker 重写账目 存量回归(中危): - image-only 定价条目(仅图片价无 token 价)恢复 token 计费 fail-closed, 不再按 $0 计费;图片计费路径不受影响 - 鉴权余额门槛恢复 balance <= 0 语义,MinimumBalanceReserve 不再作硬 403 加固: - ZIP max_items 钳制到管理员上限;Submit 补齐 Platform==Gemini 校验; gemini downloadUri 跟随前做 host 白名单校验 - 批量客户端改用共享 httpclient(拨号/TLS/响应头超时有界) - 审计点名的忽略错误(MarkDownloaded/SettlementFailed/AppendEvent 等)改为记日志 --- .../internal/handler/batch_image_handler.go | 17 ++- .../internal/repository/batch_image_queue.go | 116 ++++++++++++--- .../repository/batch_image_queue_test.go | 76 ++++++++++ .../internal/repository/batch_image_repo.go | 47 +++++- .../internal/repository/usage_billing_repo.go | 39 +++++ .../usage_billing_repo_unit_test.go | 32 +++- .../server/middleware/api_key_auth.go | 13 +- .../middleware/api_key_auth_google_test.go | 32 ++++ .../server/middleware/api_key_auth_test.go | 45 +++++- backend/internal/service/admin_service.go | 11 ++ .../service/admin_service_group_test.go | 23 ++- backend/internal/service/batch_image.go | 8 + .../service/batch_image_billing_hold.go | 12 ++ .../service/batch_image_billing_recovery.go | 61 ++++++-- .../batch_image_billing_recovery_test.go | 65 ++++++++ .../internal/service/batch_image_cleanup.go | 28 +++- .../internal/service/batch_image_download.go | 5 +- .../internal/service/batch_image_processor.go | 86 +++++++++++ .../service/batch_image_processor_test.go | 88 ++++++++++- .../service/batch_image_provider_gemini.go | 41 ++++- .../service/batch_image_provider_vertex.go | 4 +- .../internal/service/batch_image_public.go | 140 +++++++++++++++++- .../service/batch_image_public_test.go | 8 +- .../service/batch_image_settlement.go | 86 ++++++++--- .../service/batch_image_settlement_test.go | 67 +++++++++ .../internal/service/batch_image_worker.go | 58 +++++++- .../service/batch_image_worker_runtime.go | 1 + .../service/batch_image_worker_test.go | 9 +- backend/internal/service/billing_service.go | 8 + backend/internal/service/pricing_service.go | 6 + .../internal/service/pricing_service_test.go | 38 +++++ 31 files changed, 1180 insertions(+), 90 deletions(-) diff --git a/backend/internal/handler/batch_image_handler.go b/backend/internal/handler/batch_image_handler.go index 3a36634618..d739cba998 100644 --- a/backend/internal/handler/batch_image_handler.go +++ b/backend/internal/handler/batch_image_handler.go @@ -8,10 +8,12 @@ import ( "strings" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/Wei-Shaw/sub2api/internal/server/middleware" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/gin-gonic/gin" + "go.uber.org/zap" ) type BatchImageHandler struct { @@ -160,7 +162,18 @@ func (h *BatchImageHandler) ItemContent(c *gin.Context) { if _, err := io.Copy(c.Writer, stream.Reader); err != nil { return } - _ = h.service.MarkDownloaded(c.Request.Context(), owner, c.Param("id")) + h.markDownloadedBestEffort(c, owner) +} + +// markDownloadedBestEffort 在响应体已写出后标记下载状态; +// 此时无法再向客户端返回错误,失败只能记日志(不能静默丢弃)。 +func (h *BatchImageHandler) markDownloadedBestEffort(c *gin.Context, owner service.BatchImageOwner) { + if err := h.service.MarkDownloaded(c.Request.Context(), owner, c.Param("id")); err != nil { + logger.L().Warn("batch_image.mark_downloaded_failed", + zap.String("batch_id", c.Param("id")), + zap.Error(err), + ) + } } func (h *BatchImageHandler) Download(c *gin.Context) { @@ -186,7 +199,7 @@ func (h *BatchImageHandler) Download(c *gin.Context) { } return } - _ = h.service.MarkDownloaded(c.Request.Context(), owner, c.Param("id")) + h.markDownloadedBestEffort(c, owner) } func (h *BatchImageHandler) DeleteRecord(c *gin.Context) { diff --git a/backend/internal/repository/batch_image_queue.go b/backend/internal/repository/batch_image_queue.go index 16de6652e3..e10bc1584a 100644 --- a/backend/internal/repository/batch_image_queue.go +++ b/backend/internal/repository/batch_image_queue.go @@ -20,6 +20,10 @@ const ( defaultBatchImageLockPrefix = "batch_image:queue:lock:" defaultBatchImageInflightTTL = 7 * 24 * time.Hour defaultBatchImageJobLockTTL = 5 * time.Minute + + // batchImageReservePollInterval 是原子 Reserve 脚本空轮询的间隔。 + // 用轮询替代 BRPop 是为了保证 "弹出 + 写 active" 的原子性。 + batchImageReservePollInterval = time.Second ) var batchImageMoveDueDelayedScript = redis.NewScript(` @@ -47,6 +51,36 @@ end return 0 `) +var batchImageRefreshLockScript = redis.NewScript(` +if redis.call("GET", KEYS[1]) == ARGV[1] then + return redis.call("PEXPIRE", KEYS[1], ARGV[2]) +end +return 0 +`) + +// batchImageReserveScript 原子地从 ready 弹出并写入 active zset。 +// BRPop + ZAdd 两步方案在两步之间进程崩溃时 job 会脱离所有队列结构, +// 且 inflight 去重键(默认 7 天)会挡住所有重新入队。 +var batchImageReserveScript = redis.NewScript(` +local job = redis.call("RPOP", KEYS[1]) +if not job then + return nil +end +redis.call("ZADD", KEYS[2], ARGV[1], job) +return job +`) + +// batchImageEnqueueScript 原子地设置 inflight 去重键并推入 ready。 +// SetNX + LPush 两步方案在两步之间进程崩溃时,inflight 键(默认 7 天) +// 会挡住所有后续入队,而 job 从未进入 ready。 +var batchImageEnqueueScript = redis.NewScript(` +if redis.call("SET", KEYS[1], ARGV[1], "NX", "PX", ARGV[2]) then + redis.call("LPUSH", KEYS[2], ARGV[1]) + return 1 +end +return 0 +`) + type batchImageQueue struct { rdb *redis.Client readyKey string @@ -131,40 +165,65 @@ func (q *batchImageQueue) Enqueue(ctx context.Context, batchID string) error { return service.ErrInvalidBatchImageQueuePayload } - ok, err := q.rdb.SetNX(ctx, q.inflightKey(batchID), batchID, q.inflightTTL).Result() + applied, err := batchImageEnqueueScript.Run(ctx, q.rdb, + []string{q.inflightKey(batchID), q.readyKey}, + batchID, q.inflightTTL.Milliseconds(), + ).Int() if err != nil { return err } - if !ok { + if applied == 0 { return service.ErrBatchImageAlreadyQueued } - if err := q.rdb.LPush(ctx, q.readyKey, batchID).Err(); err != nil { - _ = q.rdb.Del(ctx, q.inflightKey(batchID)).Err() - return err - } return nil } func (q *batchImageQueue) Reserve(ctx context.Context, blockTimeout time.Duration) (service.ReservedBatchImageJob, error) { - result, err := q.rdb.BRPop(ctx, blockTimeout, q.readyKey).Result() + deadline := time.Now().Add(blockTimeout) + for { + batchID, err := q.reserveOnce(ctx) + if err == nil { + return service.ReservedBatchImageJob{BatchID: batchID}, nil + } + if !errors.Is(err, service.ErrBatchImageQueueEmpty) { + return service.ReservedBatchImageJob{}, err + } + remaining := time.Until(deadline) + if remaining <= 0 { + return service.ReservedBatchImageJob{}, service.ErrBatchImageQueueEmpty + } + wait := batchImageReservePollInterval + if remaining < wait { + wait = remaining + } + timer := time.NewTimer(wait) + select { + case <-ctx.Done(): + timer.Stop() + return service.ReservedBatchImageJob{}, ctx.Err() + case <-timer.C: + } + } +} + +func (q *batchImageQueue) reserveOnce(ctx context.Context) (string, error) { + raw, err := batchImageReserveScript.Run(ctx, q.rdb, []string{q.readyKey, q.activeKey}, time.Now().UnixMilli()).Result() if errors.Is(err, redis.Nil) { - return service.ReservedBatchImageJob{}, service.ErrBatchImageQueueEmpty + return "", service.ErrBatchImageQueueEmpty } if err != nil { - return service.ReservedBatchImageJob{}, err + return "", err } - if len(result) != 2 || !service.IsValidBatchImageID(result[1]) { - return service.ReservedBatchImageJob{}, service.ErrInvalidBatchImageQueuePayload + batchID, ok := raw.(string) + if !ok || !service.IsValidBatchImageID(batchID) { + // 非法 payload 已被脚本写入 active,必须移除,否则 stale 恢复会把它 + // 无限重投回 ready。 + if ok && batchID != "" { + _ = q.rdb.ZRem(ctx, q.activeKey, batchID).Err() + } + return "", service.ErrInvalidBatchImageQueuePayload } - - batchID := result[1] - if err := q.rdb.ZAdd(ctx, q.activeKey, redis.Z{ - Score: float64(time.Now().UnixMilli()), - Member: batchID, - }).Err(); err != nil { - return service.ReservedBatchImageJob{}, err - } - return service.ReservedBatchImageJob{BatchID: batchID}, nil + return batchID, nil } func (q *batchImageQueue) RequeueAfter(ctx context.Context, batchID string, delay time.Duration) error { @@ -202,7 +261,9 @@ func (q *batchImageQueue) Heartbeat(ctx context.Context, batchID string) error { if !service.IsValidBatchImageID(batchID) { return service.ErrInvalidBatchImageQueuePayload } - return q.rdb.ZAdd(ctx, q.activeKey, redis.Z{ + // XX:只刷新已存在的 active 成员。无条件 ZAdd 会在 Ack/Requeue 之后的 + // 竞态心跳里把幽灵成员塞回 active zset。 + return q.rdb.ZAddXX(ctx, q.activeKey, redis.Z{ Score: float64(time.Now().UnixMilli()), Member: batchID, }).Err() @@ -269,6 +330,19 @@ func (l *batchImageRedisJobLock) Release(ctx context.Context) error { return batchImageReleaseLockScript.Run(ctx, l.rdb, []string{l.key}, l.token).Err() } +// Refresh 在仍持有锁(token 匹配)时续期 TTL,供长处理任务的心跳调用。 +func (l *batchImageRedisJobLock) Refresh(ctx context.Context, ttl time.Duration) error { + if l == nil || l.rdb == nil || l.key == "" || l.token == "" { + return nil + } + if ttl <= 0 { + ttl = defaultBatchImageJobLockTTL + } + return batchImageRefreshLockScript.Run(ctx, l.rdb, []string{l.key}, l.token, ttl.Milliseconds()).Err() +} + +var _ service.BatchImageJobLockRefresher = (*batchImageRedisJobLock)(nil) + func newBatchImageLockToken() (string, error) { var b [16]byte if _, err := rand.Read(b[:]); err != nil { diff --git a/backend/internal/repository/batch_image_queue_test.go b/backend/internal/repository/batch_image_queue_test.go index 5188e67d07..63412ce54e 100644 --- a/backend/internal/repository/batch_image_queue_test.go +++ b/backend/internal/repository/batch_image_queue_test.go @@ -108,6 +108,82 @@ func TestBatchImageQueue_JobLockReleaseOnlyDeletesMatchingToken(t *testing.T) { require.ErrorIs(t, queue.rdb.Get(ctx, queue.lockKey(batchID)).Err(), redis.Nil) } +func TestBatchImageQueue_ReserveAtomicallyMovesJobToActive(t *testing.T) { + ctx := context.Background() + queue, _ := newBatchImageQueueTest(t) + batchID := "imgbatch_reserve" + require.NoError(t, queue.Enqueue(ctx, batchID)) + + reserved, err := queue.Reserve(ctx, time.Second) + require.NoError(t, err) + require.Equal(t, batchID, reserved.BatchID) + + // 弹出与写入 active 必须原子完成:ready 已空,active 中有该 job。 + require.Equal(t, int64(0), queue.rdb.LLen(ctx, queue.readyKey).Val()) + score, err := queue.rdb.ZScore(ctx, queue.activeKey, batchID).Result() + require.NoError(t, err) + require.Positive(t, score) +} + +func TestBatchImageQueue_ReserveReturnsEmptyAfterTimeout(t *testing.T) { + ctx := context.Background() + queue, _ := newBatchImageQueueTest(t) + + start := time.Now() + _, err := queue.Reserve(ctx, 50*time.Millisecond) + require.ErrorIs(t, err, service.ErrBatchImageQueueEmpty) + require.Less(t, time.Since(start), 5*time.Second) +} + +func TestBatchImageQueue_ReserveDropsInvalidPayload(t *testing.T) { + ctx := context.Background() + queue, _ := newBatchImageQueueTest(t) + require.NoError(t, queue.rdb.LPush(ctx, queue.readyKey, "not-a-batch-id").Err()) + + _, err := queue.Reserve(ctx, 10*time.Millisecond) + require.ErrorIs(t, err, service.ErrInvalidBatchImageQueuePayload) + // 非法 payload 不得残留在 active zset,否则 stale 恢复会无限重投。 + require.ErrorIs(t, queue.rdb.ZScore(ctx, queue.activeKey, "not-a-batch-id").Err(), redis.Nil) +} + +func TestBatchImageQueue_HeartbeatOnlyRefreshesExistingActiveMember(t *testing.T) { + ctx := context.Background() + queue, _ := newBatchImageQueueTest(t) + batchID := "imgbatch_heartbeat" + + // 不在 active 中:心跳不得创建幽灵成员。 + require.NoError(t, queue.Heartbeat(ctx, batchID)) + require.ErrorIs(t, queue.rdb.ZScore(ctx, queue.activeKey, batchID).Err(), redis.Nil) + + require.NoError(t, queue.rdb.ZAdd(ctx, queue.activeKey, redis.Z{Score: 1, Member: batchID}).Err()) + require.NoError(t, queue.Heartbeat(ctx, batchID)) + score, err := queue.rdb.ZScore(ctx, queue.activeKey, batchID).Result() + require.NoError(t, err) + require.Greater(t, score, float64(1)) +} + +func TestBatchImageQueue_JobLockRefreshExtendsTTLOnlyForHolder(t *testing.T) { + ctx := context.Background() + queue, mr := newBatchImageQueueTest(t) + batchID := "imgbatch_lock_refresh" + + lock, ok, err := queue.TryAcquireJobLock(ctx, batchID, time.Minute) + require.NoError(t, err) + require.True(t, ok) + refresher, isRefresher := lock.(service.BatchImageJobLockRefresher) + require.True(t, isRefresher) + + require.NoError(t, refresher.Refresh(ctx, 10*time.Minute)) + ttl := mr.TTL(queue.lockKey(batchID)) + require.Greater(t, ttl, 5*time.Minute) + + // token 不匹配时不得续期他人持有的锁。 + require.NoError(t, queue.rdb.Set(ctx, queue.lockKey(batchID), "other-token", time.Minute).Err()) + require.NoError(t, refresher.Refresh(ctx, 10*time.Minute)) + ttl = mr.TTL(queue.lockKey(batchID)) + require.LessOrEqual(t, ttl, time.Minute) +} + func newBatchImageQueueTest(t *testing.T) (*batchImageQueue, *miniredis.Miniredis) { t.Helper() mr := miniredis.RunT(t) diff --git a/backend/internal/repository/batch_image_repo.go b/backend/internal/repository/batch_image_repo.go index cd4bf70339..6531f75fd5 100644 --- a/backend/internal/repository/batch_image_repo.go +++ b/backend/internal/repository/batch_image_repo.go @@ -152,6 +152,45 @@ func (r *batchImageRepository) TransitionBatchImageJobStatus(ctx context.Context return tx.Commit() } +func (r *batchImageRepository) TouchBatchImageJobSubmitting(ctx context.Context, batchID string) error { + _, err := r.sql.ExecContext(ctx, ` +UPDATE batch_image_jobs +SET updated_at = $2 +WHERE batch_id = $1 + AND status IN ('created', 'uploading')`, batchID, time.Now()) + return err +} + +func (r *batchImageRepository) FailStaleUnsubmittedBatchImageJob(ctx context.Context, batchID string, cutoff time.Time, code, message string) (bool, error) { + now := time.Now() + res, err := r.sql.ExecContext(ctx, ` +UPDATE batch_image_jobs +SET status = 'failed', + last_error_code = $2, + last_error_message = $3, + finished_at = CASE WHEN finished_at IS NULL THEN $4 ELSE finished_at END, + updated_at = $4, + version = version + 1 +WHERE batch_id = $1 + AND status IN ('created', 'uploading') + AND provider_job_name IS NULL + AND updated_at <= $5`, batchID, code, message, now, cutoff) + if err != nil { + return false, err + } + affected, err := res.RowsAffected() + if err != nil { + return false, err + } + if affected == 0 { + return false, nil + } + return true, appendBatchImageEventWithSQL(ctx, r.sql, batchID, "billing_hold_recovery_failed_unsubmitted", map[string]any{ + "batch_id": batchID, + "error_code": code, + }) +} + func (r *batchImageRepository) UpdateBatchImageJobProviderOutputRef(ctx context.Context, batchID, providerOutputRef string) error { res, err := r.sql.ExecContext(ctx, ` UPDATE batch_image_jobs @@ -410,9 +449,15 @@ func (r *batchImageRepository) ReplaceBatchImageItemsForJob(ctx context.Context, func (r *batchImageRepository) replaceBatchImageItemsForJobWithSQL(ctx context.Context, sqlq batchImageSQLExecutor, batchID string, items []service.CreateBatchImageItemParams, counts service.BatchImageCounts) error { var id int64 - if err := sqlq.QueryRowContext(ctx, `SELECT id FROM batch_image_jobs WHERE batch_id = $1 FOR UPDATE`, batchID).Scan(&id); err != nil { + var status string + if err := sqlq.QueryRowContext(ctx, `SELECT id, status FROM batch_image_jobs WHERE batch_id = $1 FOR UPDATE`, batchID).Scan(&id, &status); err != nil { return translatePersistenceError(err, service.ErrBatchImageJobNotFound, nil) } + // 仅允许 indexing 状态重建 item 表:防止锁过期后掉队的 worker + // 重写已完成/已结算 job 的条目,造成账目与结果漂移。 + if status != service.BatchImageJobStatusIndexing { + return service.ErrBatchImageIndexStateConflict + } promptPreviews, err := r.batchImageItemPromptPreviews(ctx, sqlq, batchID) if err != nil { return err diff --git a/backend/internal/repository/usage_billing_repo.go b/backend/internal/repository/usage_billing_repo.go index f7e675439f..ce0adc2772 100644 --- a/backend/internal/repository/usage_billing_repo.go +++ b/backend/internal/repository/usage_billing_repo.go @@ -335,6 +335,16 @@ func releaseUsageBillingBatchImageBalance(ctx context.Context, tx *sql.Tx, cmd * if cmd.HoldAmount <= 0 { return &service.BatchImageBalanceHoldResult{}, nil } + // 释放前校验该 job 确实预留过 hold(hold request id 已被 claim), + // 防止从未成功冻结的 job 触发"幻影释放",从其他用户的冻结资金池中凭空生成余额。 + held, heldErr := batchImageHoldClaimExists(ctx, tx, service.BatchImageHoldRequestID(cmd.BatchID), cmd.APIKeyID) + if heldErr != nil { + return nil, heldErr + } + if !held { + logger.LegacyPrintf("repository.usage_billing", "[BatchImage] release skipped, hold was never reserved: batch=%s", cmd.BatchID) + return &service.BatchImageBalanceHoldResult{}, nil + } var balance, frozen float64 err := tx.QueryRowContext(ctx, ` UPDATE users @@ -358,6 +368,35 @@ func releaseUsageBillingBatchImageBalance(ctx context.Context, tx *sql.Tx, cmd * return nil, errors.New("batch image frozen balance is insufficient") } +// batchImageHoldClaimExists 检查 hold request id 是否已在 dedup(或归档)表中被 claim, +// 即该 batch 的冻结操作确实成功提交过。 +func batchImageHoldClaimExists(ctx context.Context, tx *sql.Tx, holdRequestID string, apiKeyID int64) (bool, error) { + var exists int + err := tx.QueryRowContext(ctx, ` + SELECT 1 + FROM usage_billing_dedup + WHERE request_id = $1 AND api_key_id = $2 + `, holdRequestID, apiKeyID).Scan(&exists) + if err == nil { + return true, nil + } + if !errors.Is(err, sql.ErrNoRows) { + return false, err + } + err = tx.QueryRowContext(ctx, ` + SELECT 1 + FROM usage_billing_dedup_archive + WHERE request_id = $1 AND api_key_id = $2 + `, holdRequestID, apiKeyID).Scan(&exists) + if err == nil { + return true, nil + } + if errors.Is(err, sql.ErrNoRows) { + return false, nil + } + return false, err +} + func userExistsForBilling(ctx context.Context, tx *sql.Tx, userID int64) (bool, error) { var exists int err := tx.QueryRowContext(ctx, ` diff --git a/backend/internal/repository/usage_billing_repo_unit_test.go b/backend/internal/repository/usage_billing_repo_unit_test.go index 0c469db899..6fe4d4fc26 100644 --- a/backend/internal/repository/usage_billing_repo_unit_test.go +++ b/backend/internal/repository/usage_billing_repo_unit_test.go @@ -217,15 +217,45 @@ func TestReleaseUsageBillingBatchImageBalance_ReturnsFrozenToAvailable(t *testin mock.ExpectBegin() tx, err := db.BeginTx(ctx, nil) require.NoError(t, err) + mock.ExpectQuery(`SELECT 1\s+FROM usage_billing_dedup\s+WHERE request_id = \$1 AND api_key_id = \$2`). + WithArgs(service.BatchImageHoldRequestID("imgbatch_release"), int64(7)). + WillReturnRows(sqlmock.NewRows([]string{"?column?"}).AddRow(1)) mock.ExpectQuery(releaseBatchImageHoldSQL). WithArgs(1.0, int64(42)). WillReturnRows(sqlmock.NewRows([]string{"balance", "frozen_balance"}).AddRow(10.0, 0.0)) mock.ExpectCommit() - result, err := releaseUsageBillingBatchImageBalance(ctx, tx, &service.BatchImageBalanceHoldCommand{UserID: 42, HoldAmount: 1}) + result, err := releaseUsageBillingBatchImageBalance(ctx, tx, &service.BatchImageBalanceHoldCommand{UserID: 42, APIKeyID: 7, BatchID: "imgbatch_release", HoldAmount: 1}) require.NoError(t, err) require.InDelta(t, 10.0, *result.NewBalance, 0.000001) require.InDelta(t, 0.0, *result.FrozenBalance, 0.000001) require.NoError(t, tx.Commit()) require.NoError(t, mock.ExpectationsWereMet()) } + +func TestReleaseUsageBillingBatchImageBalance_SkipsWhenHoldNeverReserved(t *testing.T) { + ctx := context.Background() + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer func() { _ = db.Close() }() + + mock.ExpectBegin() + tx, err := db.BeginTx(ctx, nil) + require.NoError(t, err) + // dedup 与归档表均无 hold claim:说明该 job 从未成功冻结, + // 释放必须跳过,不得从他人冻结资金池中凭空生成余额。 + mock.ExpectQuery(`SELECT 1\s+FROM usage_billing_dedup\s+WHERE request_id = \$1 AND api_key_id = \$2`). + WithArgs(service.BatchImageHoldRequestID("imgbatch_phantom"), int64(7)). + WillReturnError(sql.ErrNoRows) + mock.ExpectQuery(`SELECT 1\s+FROM usage_billing_dedup_archive\s+WHERE request_id = \$1 AND api_key_id = \$2`). + WithArgs(service.BatchImageHoldRequestID("imgbatch_phantom"), int64(7)). + WillReturnError(sql.ErrNoRows) + mock.ExpectCommit() + + result, err := releaseUsageBillingBatchImageBalance(ctx, tx, &service.BatchImageBalanceHoldCommand{UserID: 42, APIKeyID: 7, BatchID: "imgbatch_phantom", HoldAmount: 1}) + require.NoError(t, err) + require.Nil(t, result.NewBalance) + require.Nil(t, result.FrozenBalance) + require.NoError(t, tx.Commit()) + require.NoError(t, mock.ExpectationsWereMet()) +} diff --git a/backend/internal/server/middleware/api_key_auth.go b/backend/internal/server/middleware/api_key_auth.go index 9bcb56d7fa..1610390440 100644 --- a/backend/internal/server/middleware/api_key_auth.go +++ b/backend/internal/server/middleware/api_key_auth.go @@ -289,14 +289,11 @@ func setGroupContext(c *gin.Context, group *service.Group) { c.Request = c.Request.WithContext(ctx) } -func apiKeyBalanceBelowAuthThreshold(balance float64, cfg *config.Config) bool { - if balance <= 0 { - return true - } - if cfg == nil || cfg.Billing.MinimumBalanceReserve <= 0 { - return false - } - return balance < cfg.Billing.MinimumBalanceReserve +// apiKeyBalanceBelowAuthThreshold 保持鉴权层的历史语义:仅在余额耗尽(<=0)时拒绝。 +// MinimumBalanceReserve 只作为 billing-cache 预检的保守下限,不得复用为鉴权硬门槛, +// 否则已配置该值的存量部署升级后,0 < balance < reserve 的用户会在所有端点被静默 403。 +func apiKeyBalanceBelowAuthThreshold(balance float64, _ *config.Config) bool { + return balance <= 0 } func abortIfAPIKeyGroupUnavailable(c *gin.Context, apiKey *service.APIKey) bool { diff --git a/backend/internal/server/middleware/api_key_auth_google_test.go b/backend/internal/server/middleware/api_key_auth_google_test.go index 899cd8bbe9..45ddb0bf93 100644 --- a/backend/internal/server/middleware/api_key_auth_google_test.go +++ b/backend/internal/server/middleware/api_key_auth_google_test.go @@ -542,6 +542,8 @@ func TestApiKeyAuthWithSubscriptionGoogle_InsufficientBalance(t *testing.T) { func TestApiKeyAuthWithSubscriptionGoogle_BalanceBelowMinimumReserve(t *testing.T) { gin.SetMode(gin.TestMode) + // 鉴权层保持历史语义:MinimumBalanceReserve 只用于 billing-cache 预检, + // 0 < balance < reserve 的用户不得在鉴权中间件被硬 403。 r := gin.New() apiKeyService := newTestAPIKeyService(fakeAPIKeyRepo{ getByKey: func(ctx context.Context, key string) (*service.APIKey, error) { @@ -567,6 +569,36 @@ func TestApiKeyAuthWithSubscriptionGoogle_BalanceBelowMinimumReserve(t *testing. rec := httptest.NewRecorder() r.ServeHTTP(rec, req) + require.Equal(t, http.StatusOK, rec.Code) +} + +func TestApiKeyAuthWithSubscriptionGoogle_RejectsExhaustedBalance(t *testing.T) { + gin.SetMode(gin.TestMode) + + r := gin.New() + apiKeyService := newTestAPIKeyService(fakeAPIKeyRepo{ + getByKey: func(ctx context.Context, key string) (*service.APIKey, error) { + return &service.APIKey{ + ID: 1, + Key: key, + Status: service.StatusActive, + User: &service.User{ + ID: 123, + Status: service.StatusActive, + Balance: 0, + }, + }, nil + }, + }) + cfg := &config.Config{} + r.Use(APIKeyAuthWithSubscriptionGoogle(apiKeyService, nil, cfg)) + r.GET("/v1beta/test", func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) }) + + req := httptest.NewRequest(http.MethodGet, "/v1beta/test", nil) + req.Header.Set("Authorization", "Bearer ok") + rec := httptest.NewRecorder() + r.ServeHTTP(rec, req) + require.Equal(t, http.StatusForbidden, rec.Code) var resp googleErrorResponse require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp)) diff --git a/backend/internal/server/middleware/api_key_auth_test.go b/backend/internal/server/middleware/api_key_auth_test.go index 04ab9410ac..d5fbc46098 100644 --- a/backend/internal/server/middleware/api_key_auth_test.go +++ b/backend/internal/server/middleware/api_key_auth_test.go @@ -1000,7 +1000,7 @@ func TestAPIKeyAuthTouchesLastUsedInStandardMode(t *testing.T) { require.Equal(t, 1, touchCalls) } -func TestAPIKeyAuthRejectsBalanceBelowMinimumReserve(t *testing.T) { +func TestAPIKeyAuthAllowsBalanceBelowMinimumReserve(t *testing.T) { gin.SetMode(gin.TestMode) user := &service.User{ @@ -1039,6 +1039,49 @@ func TestAPIKeyAuthRejectsBalanceBelowMinimumReserve(t *testing.T) { req.Header.Set("x-api-key", apiKey.Key) router.ServeHTTP(w, req) + // 鉴权层保持历史语义:MinimumBalanceReserve 只用于 billing-cache 预检, + // 0 < balance < reserve 不得被鉴权中间件硬 403(存量部署静默行为变更)。 + require.Equal(t, http.StatusOK, w.Code) +} + +func TestAPIKeyAuthRejectsExhaustedBalance(t *testing.T) { + gin.SetMode(gin.TestMode) + + user := &service.User{ + ID: 10, + Role: service.RoleUser, + Status: service.StatusActive, + Balance: 0, + Concurrency: 3, + } + apiKey := &service.APIKey{ + ID: 104, + UserID: user.ID, + Key: "held-balance-zero", + Status: service.StatusActive, + User: user, + } + apiKeyRepo := &stubApiKeyRepo{ + getByKey: func(ctx context.Context, key string) (*service.APIKey, error) { + if key != apiKey.Key { + return nil, service.ErrAPIKeyNotFound + } + clone := *apiKey + userClone := *user + clone.User = &userClone + return &clone, nil + }, + } + + cfg := &config.Config{RunMode: config.RunModeStandard} + apiKeyService := service.NewAPIKeyService(apiKeyRepo, nil, nil, nil, nil, nil, cfg) + router := newAuthTestRouter(apiKeyService, nil, cfg) + + w := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/t", nil) + req.Header.Set("x-api-key", apiKey.Key) + router.ServeHTTP(w, req) + require.Equal(t, http.StatusForbidden, w.Code) requireAPIKeyAuthError(t, w, "INSUFFICIENT_BALANCE", "Insufficient account balance") } diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go index 8c67ae987e..f3b6fd555a 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -1877,6 +1877,11 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn } batchImageHoldMultiplier = *input.BatchImageHoldMultiplier } + // 不变式:hold 比例 >= discount 比例。否则批量任务成功率足够高时 + // 实际成本会超过冻结额,结算永远失败、用户冻结余额无法解冻。 + if batchImageHoldMultiplier < batchImageDiscountMultiplier { + return nil, errors.New("batch_image_hold_multiplier must be >= batch_image_discount_multiplier") + } peakRateMultiplier := 1.0 if input.PeakRateMultiplier != nil { @@ -2174,6 +2179,12 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd } group.BatchImageHoldMultiplier = *input.BatchImageHoldMultiplier } + // 仅在本次更新显式触碰任一比例时校验合并后的不变式(hold >= discount), + // 避免存量脏数据阻塞其他字段的正常更新(提交侧另有钳制兜底)。 + if (input.BatchImageDiscountMultiplier != nil || input.BatchImageHoldMultiplier != nil) && + group.BatchImageHoldMultiplier < group.BatchImageDiscountMultiplier { + return nil, errors.New("batch_image_hold_multiplier must be >= batch_image_discount_multiplier") + } if input.PeakRateEnabled != nil { group.PeakRateEnabled = *input.PeakRateEnabled } diff --git a/backend/internal/service/admin_service_group_test.go b/backend/internal/service/admin_service_group_test.go index 0c229c389e..d2e3c49996 100644 --- a/backend/internal/service/admin_service_group_test.go +++ b/backend/internal/service/admin_service_group_test.go @@ -475,7 +475,7 @@ func TestAdminService_CreateGroup_BatchImagePricingSettings(t *testing.T) { repo := &groupRepoStubForAdmin{} svc := &adminServiceImpl{groupRepo: repo} discount := 0.8 - hold := 0.6 + hold := 0.9 group, err := svc.CreateGroup(context.Background(), &CreateGroupInput{ Name: "batch-image-pricing", @@ -488,7 +488,26 @@ func TestAdminService_CreateGroup_BatchImagePricingSettings(t *testing.T) { require.NotNil(t, group) require.NotNil(t, repo.created) require.InDelta(t, 0.8, repo.created.BatchImageDiscountMultiplier, 1e-12) - require.InDelta(t, 0.6, repo.created.BatchImageHoldMultiplier, 1e-12) + require.InDelta(t, 0.9, repo.created.BatchImageHoldMultiplier, 1e-12) +} + +func TestAdminService_CreateGroup_RejectsHoldBelowDiscount(t *testing.T) { + repo := &groupRepoStubForAdmin{} + svc := &adminServiceImpl{groupRepo: repo} + discount := 0.8 + hold := 0.6 + + // hold < discount 时,成功率足够高的批量任务实际成本会超过冻结额, + // 结算永远失败,必须在配置入口拒绝。 + _, err := svc.CreateGroup(context.Background(), &CreateGroupInput{ + Name: "batch-image-pricing-invalid", + Platform: PlatformGemini, + RateMultiplier: 1, + BatchImageDiscountMultiplier: &discount, + BatchImageHoldMultiplier: &hold, + }) + require.Error(t, err) + require.Nil(t, repo.created) } func TestAdminService_GroupBatchImagePricingValidation(t *testing.T) { diff --git a/backend/internal/service/batch_image.go b/backend/internal/service/batch_image.go index 4fad11dd5e..af0206b8b0 100644 --- a/backend/internal/service/batch_image.go +++ b/backend/internal/service/batch_image.go @@ -50,6 +50,7 @@ var ( ErrBatchImageIndexParseFailed = infraerrors.New(http.StatusBadGateway, "BATCH_IMAGE_INDEX_PARSE_FAILED", "batch image provider output parse failed") ErrBatchImageIndexNoResultLines = infraerrors.New(http.StatusBadGateway, "BATCH_IMAGE_INDEX_NO_RESULT_LINES", "batch image provider output has no result lines") ErrBatchImageDuplicateCustomID = infraerrors.New(http.StatusBadGateway, "DUPLICATE_CUSTOM_ID_IN_OUTPUT", "batch image provider output contains duplicate custom id") + ErrBatchImageIndexStateConflict = infraerrors.New(http.StatusConflict, "BATCH_IMAGE_INDEX_STATE_CONFLICT", "batch image job is no longer in indexing state") ErrBatchImageSettlementInvalidStatus = infraerrors.New(http.StatusBadRequest, "BATCH_IMAGE_SETTLEMENT_INVALID_STATUS", "batch image job is not ready for settlement") ErrBatchImageSettlementManifestConflict = infraerrors.New(http.StatusConflict, "BATCH_IMAGE_SETTLEMENT_MANIFEST_CONFLICT", "batch image settlement manifest hash conflict") @@ -307,6 +308,13 @@ type BatchImageRepository interface { GetBatchImageJobByID(ctx context.Context, id int64) (*BatchImageJob, error) ListBatchImageJobsForOwner(ctx context.Context, userID, apiKeyID int64, filter BatchImageJobFilter) ([]*BatchImageJob, error) TransitionBatchImageJobStatus(ctx context.Context, batchID, toStatus string, opts BatchImageTransitionOptions) error + // TouchBatchImageJobSubmitting 刷新未提交(created/uploading)job 的 updated_at, + // 作为慢提交期间的心跳,防止被 stale 恢复扫描误杀。 + TouchBatchImageJobSubmitting(ctx context.Context, batchID string) error + // FailStaleUnsubmittedBatchImageJob 原子地将仍处于 created/uploading 且 + // provider_job_name 为空、updated_at 早于 cutoff 的 job 转为 failed。 + // 返回 false 表示 job 已被并发推进(如已提交成功),调用方不得释放冻结。 + FailStaleUnsubmittedBatchImageJob(ctx context.Context, batchID string, cutoff time.Time, code, message string) (bool, error) UpdateBatchImageJobProviderOutputRef(ctx context.Context, batchID, providerOutputRef string) error UpdateBatchImageJobProviderSubmit(ctx context.Context, params UpdateBatchImageJobProviderSubmitParams) error RecordBatchImageJobSubmitFailure(ctx context.Context, batchID, code, message string, markFailed bool) error diff --git a/backend/internal/service/batch_image_billing_hold.go b/backend/internal/service/batch_image_billing_hold.go index af61117db0..80f77208ba 100644 --- a/backend/internal/service/batch_image_billing_hold.go +++ b/backend/internal/service/batch_image_billing_hold.go @@ -4,6 +4,9 @@ import ( "context" "errors" "strings" + + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "go.uber.org/zap" ) const ( @@ -98,6 +101,15 @@ func releaseBatchImageBalanceHold(ctx context.Context, repo UsageBillingReposito return nil } if _, err := repo.ReleaseBatchImageBalance(ctx, cmd); err != nil { + // 同一 release request id 出现指纹冲突,说明此前已有一次携带不同 + // payloadHash 的释放成功提交(资金已归还)。视为幂等成功, + // 避免历史指纹不一致的 job 永远卡在释放失败的毒消息循环里。 + if errors.Is(err, ErrUsageBillingRequestConflict) { + logger.L().Warn("batch_image.release_fingerprint_conflict_treated_as_released", + zap.String("batch_id", job.BatchID), + ) + return nil + } return ErrBatchImageBillingHoldFailed.WithCause(err) } return nil diff --git a/backend/internal/service/batch_image_billing_recovery.go b/backend/internal/service/batch_image_billing_recovery.go index d89f508f47..5987fec48e 100644 --- a/backend/internal/service/batch_image_billing_recovery.go +++ b/backend/internal/service/batch_image_billing_recovery.go @@ -4,6 +4,9 @@ import ( "context" "errors" "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "go.uber.org/zap" ) const ( @@ -15,6 +18,7 @@ type BatchImageBillingRecoveryService struct { Repo BatchImageRepository Billing UsageBillingRepository AuthCache APIKeyAuthCacheInvalidator + Queue BatchImageQueue StaleAfter time.Duration Limit int } @@ -31,32 +35,69 @@ func (s *BatchImageBillingRecoveryService) ReleaseStaleUnsubmittedOnce(ctx conte if limit <= 0 { limit = defaultBatchImageBillingRecoveryLimit } - jobs, err := s.Repo.ListStaleUnsubmittedBatchImageJobs(ctx, time.Now().Add(-staleAfter), limit) + cutoff := time.Now().Add(-staleAfter) + jobs, err := s.Repo.ListStaleUnsubmittedBatchImageJobs(ctx, cutoff, limit) if err != nil { return 0, err } released := 0 + var lastErr error for _, job := range jobs { if job == nil { continue } - msg := "batch image submission did not reach provider before recovery cutoff" - if err := s.Repo.TransitionBatchImageJobStatus(ctx, job.BatchID, BatchImageJobStatusFailed, BatchImageTransitionOptions{ - EventType: "billing_hold_recovery_failed_unsubmitted", - EventPayload: map[string]any{"batch_id": job.BatchID}, - ErrorCode: batchImageStringPtr("SUBMIT_STALE_BEFORE_PROVIDER"), - ErrorMessage: batchImageStringPtr(msg), - }); err != nil && !errors.Is(err, ErrBatchImageInvalidTransition) { + if err := ctx.Err(); err != nil { return released, err } + msg := "batch image submission did not reach provider before recovery cutoff" + // 原子转 failed 并复核 stale 条件:List 与转态之间 job 可能已被慢提交 + // 心跳续期或提交成功(provider_job_name 已写入),此时绝不能退款, + // 否则上游任务照常产生成本而用户已拿回冻结余额。 + applied, err := s.Repo.FailStaleUnsubmittedBatchImageJob(ctx, job.BatchID, cutoff, "SUBMIT_STALE_BEFORE_PROVIDER", msg) + if err != nil { + // applied=true 时 UPDATE 已提交(仅审计事件写入失败):必须继续释放, + // 否则 job 已转 failed、不再出现在 stale 列表,冻结余额会永久泄漏。 + if !applied { + lastErr = err + continue + } + logger.L().Warn("batch_image.recovery_fail_event_append_failed", + zap.String("batch_id", job.BatchID), + zap.Error(err), + ) + } + if !applied { + continue + } job.Status = BatchImageJobStatusFailed if err := releaseBatchImageBalanceHold(ctx, s.Billing, job, batchImageDerefString(job.RequestHash)); err != nil { - return released, err + // job 已转 failed、不会再进入 stale 列表:必须给释放失败留下 + // 自动重试路径(入队后由 worker 的 releaseTerminalHold 兜底), + // 否则冻结余额永久泄漏。 + logger.L().Warn("batch_image.recovery_release_failed", + zap.String("batch_id", job.BatchID), + zap.Error(err), + ) + s.enqueueReleaseRetry(ctx, job.BatchID) + lastErr = err + continue } if s.AuthCache != nil && job.UserID > 0 { s.AuthCache.InvalidateAuthCacheByUserID(ctx, job.UserID) } released++ } - return released, nil + return released, lastErr +} + +func (s *BatchImageBillingRecoveryService) enqueueReleaseRetry(ctx context.Context, batchID string) { + if s == nil || s.Queue == nil { + return + } + if err := s.Queue.Enqueue(ctx, batchID); err != nil && !errors.Is(err, ErrBatchImageAlreadyQueued) { + logger.L().Warn("batch_image.recovery_release_retry_enqueue_failed", + zap.String("batch_id", batchID), + zap.Error(err), + ) + } } diff --git a/backend/internal/service/batch_image_billing_recovery_test.go b/backend/internal/service/batch_image_billing_recovery_test.go index 2ab2783f83..8ddadfff29 100644 --- a/backend/internal/service/batch_image_billing_recovery_test.go +++ b/backend/internal/service/batch_image_billing_recovery_test.go @@ -4,12 +4,23 @@ package service import ( "context" + "errors" "testing" "time" "github.com/stretchr/testify/require" ) +type recordingBatchImageQueue struct { + *fakeBatchImageQueue + enqueued []string +} + +func (q *recordingBatchImageQueue) Enqueue(_ context.Context, batchID string) error { + q.enqueued = append(q.enqueued, batchID) + return nil +} + func TestBatchImageBillingRecoveryService_ReleasesStaleUnsubmittedHold(t *testing.T) { repo := newFakeBatchImageRepository() apiKeyID := int64(22) @@ -50,3 +61,57 @@ func TestBatchImageBillingRecoveryService_ReleasesStaleUnsubmittedHold(t *testin require.Equal(t, BatchImageReleaseRequestID(stale.BatchID), billing.releases[0].RequestID) require.Equal(t, BatchImageJobStatusSubmitted, repo.jobs[active.BatchID].Status) } + +func TestBatchImageBillingRecoveryService_SkipsJobRefreshedByHeartbeat(t *testing.T) { + repo := newFakeBatchImageRepository() + apiKeyID := int64(22) + holdAmount := 0.5 + // updated_at 在 cutoff 之后(慢提交心跳持续续期):不得误杀退款。 + fresh := &BatchImageJob{ + BatchID: "imgbatch_fresh_uploading", + UserID: 11, + APIKeyID: &apiKeyID, + Status: BatchImageJobStatusUploading, + EstimatedCost: holdAmount, + HoldAmount: &holdAmount, + CreatedAt: time.Now().Add(-time.Hour), + UpdatedAt: time.Now(), + } + repo.jobs[fresh.BatchID] = fresh + billing := &fakeBatchImageBillingRepo{} + svc := &BatchImageBillingRecoveryService{Repo: repo, Billing: billing, StaleAfter: time.Minute, Limit: 10} + + released, err := svc.ReleaseStaleUnsubmittedOnce(context.Background()) + require.NoError(t, err) + require.Equal(t, 0, released) + require.Equal(t, BatchImageJobStatusUploading, repo.jobs[fresh.BatchID].Status) + require.Empty(t, billing.releases) +} + +func TestBatchImageBillingRecoveryService_EnqueuesRetryWhenReleaseFails(t *testing.T) { + repo := newFakeBatchImageRepository() + apiKeyID := int64(22) + holdAmount := 0.5 + stale := &BatchImageJob{ + BatchID: "imgbatch_stale_release_fail", + UserID: 11, + APIKeyID: &apiKeyID, + Status: BatchImageJobStatusCreated, + EstimatedCost: holdAmount, + HoldAmount: &holdAmount, + CreatedAt: time.Now().Add(-time.Hour), + UpdatedAt: time.Now().Add(-time.Hour), + } + repo.jobs[stale.BatchID] = stale + billing := &fakeBatchImageBillingRepo{releaseErr: errors.New("billing db down")} + queue := &recordingBatchImageQueue{fakeBatchImageQueue: newFakeBatchImageQueue("")} + svc := &BatchImageBillingRecoveryService{Repo: repo, Billing: billing, Queue: queue, StaleAfter: time.Minute, Limit: 10} + + released, err := svc.ReleaseStaleUnsubmittedOnce(context.Background()) + // job 已转 failed、不会再出现在 stale 列表:释放失败必须入队重试 + //(由 worker 的 releaseTerminalHold 兜底),否则冻结余额永久泄漏。 + require.Error(t, err) + require.Equal(t, 0, released) + require.Equal(t, BatchImageJobStatusFailed, repo.jobs[stale.BatchID].Status) + require.Equal(t, []string{stale.BatchID}, queue.enqueued) +} diff --git a/backend/internal/service/batch_image_cleanup.go b/backend/internal/service/batch_image_cleanup.go index 7b2b527080..364a7e4795 100644 --- a/backend/internal/service/batch_image_cleanup.go +++ b/backend/internal/service/batch_image_cleanup.go @@ -9,6 +9,8 @@ import ( "github.com/Wei-Shaw/sub2api/internal/config" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "go.uber.org/zap" ) const ( @@ -38,6 +40,17 @@ func NewBatchImageCleanupService(repo BatchImageRepository, accountRepo AccountR } } +// appendCleanupEvent 追加清理审计事件;事件写入失败不阻断清理流程,但必须留痕。 +func (s *BatchImageCleanupService) appendCleanupEvent(ctx context.Context, batchID, eventType string, payload any) { + if err := s.Repo.AppendBatchImageEvent(ctx, batchID, eventType, payload); err != nil { + logger.L().Warn("batch_image.cleanup_event_failed", + zap.String("batch_id", batchID), + zap.String("event_type", eventType), + zap.Error(err), + ) + } +} + func (s *BatchImageCleanupService) DeleteOutputsForOwner(ctx context.Context, owner BatchImageOwner, batchID string) (*BatchImagePublicBatch, error) { job, err := s.Repo.GetBatchImageJobByBatchIDForOwner(ctx, owner.UserID, owner.APIKeyID, batchID) if err != nil { @@ -49,7 +62,7 @@ func (s *BatchImageCleanupService) DeleteOutputsForOwner(ctx context.Context, ow if job.Status != BatchImageJobStatusCompleted { return nil, ErrBatchImageOutputDeleteNotReady } - _ = s.Repo.AppendBatchImageEvent(ctx, job.BatchID, "manual_output_delete_requested", map[string]any{ + s.appendCleanupEvent(ctx, job.BatchID, "manual_output_delete_requested", map[string]any{ "batch_id": job.BatchID, "cleanup_target": "output", "reason": "manual", @@ -178,7 +191,7 @@ func (s *BatchImageCleanupService) cleanupJob(ctx context.Context, job *BatchIma if !IsTerminalBatchImageJobStatus(job.Status) { return ErrBatchImageCleanupFailed } - _ = s.Repo.AppendBatchImageEvent(ctx, job.BatchID, "input_cleanup_started", cleanupEventPayload(job.BatchID, target, reason, nil)) + s.appendCleanupEvent(ctx, job.BatchID, "input_cleanup_started", cleanupEventPayload(job.BatchID, target, reason, nil)) case CleanupTargetOutput: if job.OutputDeletedAt != nil || job.Status == BatchImageJobStatusOutputDeleted { return nil @@ -186,7 +199,7 @@ func (s *BatchImageCleanupService) cleanupJob(ctx context.Context, job *BatchIma if job.Status != BatchImageJobStatusCompleted && job.Status != BatchImageJobStatusFailed && job.Status != BatchImageJobStatusCancelled { return ErrBatchImageOutputDeleteNotReady } - _ = s.Repo.AppendBatchImageEvent(ctx, job.BatchID, "output_cleanup_started", cleanupEventPayload(job.BatchID, target, reason, nil)) + s.appendCleanupEvent(ctx, job.BatchID, "output_cleanup_started", cleanupEventPayload(job.BatchID, target, reason, nil)) default: return ErrUnsupportedCleanupTarget } @@ -194,9 +207,14 @@ func (s *BatchImageCleanupService) cleanupJob(ctx context.Context, job *BatchIma if err := s.callProviderCleanup(ctx, job, target); err != nil { code := cleanupFailureCode(err) msg := sanitizeBatchImagePublicMessage(err.Error()) - _ = s.Repo.RecordBatchImageCleanupFailure(ctx, job.BatchID, code, msg) + if recordErr := s.Repo.RecordBatchImageCleanupFailure(ctx, job.BatchID, code, msg); recordErr != nil { + logger.L().Warn("batch_image.cleanup_failure_record_failed", + zap.String("batch_id", job.BatchID), + zap.Error(recordErr), + ) + } event := string(target) + "_cleanup_failed" - _ = s.Repo.AppendBatchImageEvent(ctx, job.BatchID, event, map[string]any{"batch_id": job.BatchID, "cleanup_target": string(target), "reason": reason, "error_code": code}) + s.appendCleanupEvent(ctx, job.BatchID, event, map[string]any{"batch_id": job.BatchID, "cleanup_target": string(target), "reason": reason, "error_code": code}) if errors.Is(err, ErrBatchImageProviderUnsafeCleanupPath) { return ErrBatchImageCleanupUnsafePath } diff --git a/backend/internal/service/batch_image_download.go b/backend/internal/service/batch_image_download.go index b08b99c75f..7f303d22c5 100644 --- a/backend/internal/service/batch_image_download.go +++ b/backend/internal/service/batch_image_download.go @@ -184,8 +184,9 @@ func (s *BatchImageDownloadService) StreamZip(ctx context.Context, owner BatchIm return nil, err } maxItems := opts.MaxItems - if maxItems <= 0 { - maxItems = s.maxZipItems() + if cap := s.maxZipItems(); maxItems <= 0 || maxItems > cap { + // 客户端传入的 max_items 不得放大管理员配置的 ZIP 上限。 + maxItems = cap } if job.SuccessCount > maxItems { return nil, ErrBatchImageZipTooManyItems diff --git a/backend/internal/service/batch_image_processor.go b/backend/internal/service/batch_image_processor.go index 24496f6bfa..414890bd88 100644 --- a/backend/internal/service/batch_image_processor.go +++ b/backend/internal/service/batch_image_processor.go @@ -8,6 +8,7 @@ import ( "fmt" "io" "net/http" + "sort" "strconv" "strings" "time" @@ -188,6 +189,11 @@ func (p *BatchImageProviderProcessor) indexAndSettle(ctx context.Context, job *B if errors.Is(err, ErrBatchImageIndexOutputMissing) { return BatchImageProcessResult{}, err } + // job 状态已被并发方推进(如已进入 settling/终态):不是索引数据问题, + // 短延迟 requeue 让下一轮按最新状态处理,不能误转 failed。 + if errors.Is(err, ErrBatchImageIndexStateConflict) { + return BatchImageProcessResult{RequeueAfter: time.Millisecond}, nil + } code := "INDEX_PARSE_FAILED" if errors.Is(err, ErrBatchImageDuplicateCustomID) { code = "DUPLICATE_CUSTOM_ID_IN_OUTPUT" @@ -281,6 +287,11 @@ func (i *BatchImageResultIndexer) Index(ctx context.Context, job *BatchImageJob, if i == nil || i.Repo == nil || job == nil || provider == nil { return nil, ErrBatchImageIndexOutputMissing } + expected, err := i.listExpectedCustomIDs(ctx, job.BatchID) + if err != nil { + return nil, err + } + r, _, err := provider.OpenResult(ctx, job, account) if err != nil { return nil, ErrBatchImageIndexOutputMissing.WithCause(err) @@ -291,6 +302,7 @@ func (i *BatchImageResultIndexer) Index(ctx context.Context, job *BatchImageJob, scanner.Buffer(make([]byte, 0, 64*1024), 16*1024*1024) seen := make(map[string]int) + unknownCount := 0 var items []CreateBatchImageItemParams result := &BatchImageIndexResult{} lineNumber := 0 @@ -310,6 +322,14 @@ func (i *BatchImageResultIndexer) Index(ctx context.Context, job *BatchImageJob, if err != nil { return nil, err } + // 与提交时的 custom_id 集对账:provider 输出中未知/多余的行不能进入 item 表, + // 否则 success+fail > item_count 会让结算永远校验失败。 + if len(expected) > 0 { + if _, ok := expected[parsed.CustomID]; !ok { + unknownCount++ + continue + } + } if firstLine, ok := seen[parsed.CustomID]; ok { return nil, ErrBatchImageDuplicateCustomID.WithCause(fmt.Errorf("custom id %q duplicated at lines %d and %d", parsed.CustomID, firstLine, lineNumber)) } @@ -344,9 +364,52 @@ func (i *BatchImageResultIndexer) Index(ctx context.Context, job *BatchImageJob, } return nil, err } + // 输出中漏掉的已提交项必须补失败记录,而不是静默消失: + // 否则用户看不到该项,且只按成功数计费会掩盖 provider 的丢单。 + missingCount := 0 + if len(expected) > 0 { + missingIDs := make([]string, 0) + for customID := range expected { + if _, ok := seen[customID]; !ok { + missingIDs = append(missingIDs, customID) + } + } + sort.Strings(missingIDs) + for _, customID := range missingIDs { + items = append(items, CreateBatchImageItemParams{ + JobID: job.BatchID, + CustomID: customID, + Status: BatchImageItemStatusFailed, + ProviderSourceObject: batchImageOptionalStringPtr(sourceObject), + ErrorCode: batchImageStringPtr("PROVIDER_RESULT_MISSING"), + ErrorMessage: batchImageStringPtr("provider output did not include a result for this item"), + IndexedAt: &now, + }) + result.FailCount++ + result.TotalCount++ + } + missingCount = len(missingIDs) + } if result.TotalCount == 0 { return nil, ErrBatchImageIndexNoResultLines } + if unknownCount > 0 || missingCount > 0 { + logger.L().Warn("batch_image.index_reconciled", + zap.String("batch_id", job.BatchID), + zap.Int("unknown_custom_ids", unknownCount), + zap.Int("missing_custom_ids", missingCount), + ) + if err := i.Repo.AppendBatchImageEvent(ctx, job.BatchID, "index_reconciled", map[string]any{ + "batch_id": job.BatchID, + "unknown_custom_ids": unknownCount, + "missing_custom_ids": missingCount, + }); err != nil { + logger.L().Warn("batch_image.index_reconcile_event_failed", + zap.String("batch_id", job.BatchID), + zap.Error(err), + ) + } + } if err := i.Repo.ReplaceBatchImageItemsForJob(ctx, job.BatchID, items, BatchImageCounts{ SuccessCount: result.SuccessCount, FailCount: result.FailCount, @@ -356,6 +419,29 @@ func (i *BatchImageResultIndexer) Index(ctx context.Context, job *BatchImageJob, return result, nil } +// listExpectedCustomIDs 返回该 job 当前 item 表中的全部 custom_id 集合, +// 即提交时预创建(或上一轮索引重建)的完整条目清单,用于与 provider 输出对账。 +func (i *BatchImageResultIndexer) listExpectedCustomIDs(ctx context.Context, batchID string) (map[string]struct{}, error) { + const pageSize = 500 + expected := make(map[string]struct{}) + offset := 0 + for { + page, err := i.Repo.ListBatchImageItems(ctx, batchID, BatchImageItemFilter{Limit: pageSize, Offset: offset}) + if err != nil { + return nil, err + } + for _, item := range page { + if item != nil { + expected[item.CustomID] = struct{}{} + } + } + if len(page) < pageSize { + return expected, nil + } + offset += len(page) + } +} + type ParsedBatchImageResult struct { CustomID string Status string diff --git a/backend/internal/service/batch_image_processor_test.go b/backend/internal/service/batch_image_processor_test.go index 8a37126161..f54a1a99ec 100644 --- a/backend/internal/service/batch_image_processor_test.go +++ b/backend/internal/service/batch_image_processor_test.go @@ -114,12 +114,58 @@ func TestBatchImageResultIndexer_WritesCountsAndReplacesItems(t *testing.T) { require.Equal(t, BatchImageCounts{SuccessCount: 1, FailCount: 1}, repo.counts[job.BatchID]) require.NotContains(t, fmt.Sprintf("%+v", repo.items[job.BatchID]), batchImageTestData) + // 重新索引时与现有 custom_id 集对账:未知的 "ok2" 被丢弃, + // 输出中缺失的 ok/bad 补为 PROVIDER_RESULT_MISSING 失败记录。 provider.result = `{"key":"ok2","response":{"candidates":[{"content":{"parts":[{"inlineData":{"mimeType":"image/webp","data":"` + batchImageTestData + `"}}]}}]}}` + "\n" result, err = (&BatchImageResultIndexer{Repo: repo}).Index(context.Background(), job, provider, &Account{}) require.NoError(t, err) - require.Equal(t, 1, result.TotalCount) - require.Len(t, repo.items[job.BatchID], 1) - require.Equal(t, "ok2", repo.items[job.BatchID][0].CustomID) + require.Equal(t, 2, result.TotalCount) + require.Equal(t, 0, result.SuccessCount) + require.Equal(t, 2, result.FailCount) + require.Len(t, repo.items[job.BatchID], 2) + gotIDs := []string{repo.items[job.BatchID][0].CustomID, repo.items[job.BatchID][1].CustomID} + require.ElementsMatch(t, []string{"ok", "bad"}, gotIDs) + for _, item := range repo.items[job.BatchID] { + require.Equal(t, BatchImageItemStatusFailed, item.Status) + require.Equal(t, "PROVIDER_RESULT_MISSING", batchImageDerefString(item.ErrorCode)) + } +} + +func TestBatchImageResultIndexer_ReconcilesMissingAndUnknownCustomIDs(t *testing.T) { + repo := newFakeBatchImageRepository() + outputRef := "files/output" + job := &BatchImageJob{BatchID: "imgbatch_reconcile", ProviderOutputRef: &outputRef, ItemCount: 3} + // 预创建提交时的 pending 条目(提交流程的行为)。 + require.NoError(t, repo.BulkCreateBatchImageItems(context.Background(), []CreateBatchImageItemParams{ + {JobID: job.BatchID, CustomID: "a", Status: BatchImageItemStatusPending}, + {JobID: job.BatchID, CustomID: "b", Status: BatchImageItemStatusPending}, + {JobID: job.BatchID, CustomID: "c", Status: BatchImageItemStatusPending}, + })) + // provider 输出:a 成功,b 失败,c 漏掉,多出未知的 x。 + output := strings.Join([]string{ + `{"key":"a","response":{"candidates":[{"content":{"parts":[{"inlineData":{"mimeType":"image/png","data":"` + batchImageTestData + `"}}]}}]}}`, + `{"key":"b","error":{"code":"SAFETY","message":"blocked"}}`, + `{"key":"x","error":{"code":"UNKNOWN","message":"not ours"}}`, + }, "\n") + "\n" + provider := &fakeProcessorProvider{result: output} + + result, err := (&BatchImageResultIndexer{Repo: repo}).Index(context.Background(), job, provider, &Account{}) + require.NoError(t, err) + require.Equal(t, 3, result.TotalCount) + require.Equal(t, 1, result.SuccessCount) + require.Equal(t, 2, result.FailCount) + require.Len(t, repo.items[job.BatchID], 3) + byID := make(map[string]CreateBatchImageItemParams) + for _, item := range repo.items[job.BatchID] { + byID[item.CustomID] = item + } + require.NotContains(t, byID, "x") + require.Equal(t, BatchImageItemStatusSuccess, byID["a"].Status) + require.Equal(t, BatchImageItemStatusFailed, byID["b"].Status) + require.Equal(t, BatchImageItemStatusFailed, byID["c"].Status) + require.Equal(t, "PROVIDER_RESULT_MISSING", batchImageDerefString(byID["c"].ErrorCode)) + // 对账后 success+fail == item_count,结算计数校验可通过。 + require.Equal(t, job.ItemCount, result.SuccessCount+result.FailCount) } func TestBatchImageResultIndexer_EmptyInvalidAndDuplicateOutput(t *testing.T) { @@ -487,6 +533,37 @@ func (r *fakeBatchImageRepository) TransitionBatchImageJobStatus(_ context.Conte return nil } +func (r *fakeBatchImageRepository) TouchBatchImageJobSubmitting(_ context.Context, batchID string) error { + job, ok := r.jobs[batchID] + if !ok { + return ErrBatchImageJobNotFound + } + if job.Status == BatchImageJobStatusCreated || job.Status == BatchImageJobStatusUploading { + job.UpdatedAt = time.Now() + } + return nil +} + +func (r *fakeBatchImageRepository) FailStaleUnsubmittedBatchImageJob(_ context.Context, batchID string, cutoff time.Time, code, message string) (bool, error) { + job, ok := r.jobs[batchID] + if !ok { + return false, ErrBatchImageJobNotFound + } + if job.Status != BatchImageJobStatusCreated && job.Status != BatchImageJobStatusUploading { + return false, nil + } + if batchImageDerefString(job.ProviderJobName) != "" || job.UpdatedAt.After(cutoff) { + return false, nil + } + job.Status = BatchImageJobStatusFailed + job.LastErrorCode = batchImageStringPtr(code) + job.LastErrorMessage = batchImageStringPtr(message) + job.UpdatedAt = time.Now() + r.transitions[batchID] = append(r.transitions[batchID], BatchImageJobStatusFailed) + r.events[batchID] = append(r.events[batchID], "billing_hold_recovery_failed_unsubmitted") + return true, nil +} + func (r *fakeBatchImageRepository) UpdateBatchImageJobProviderOutputRef(_ context.Context, batchID, providerOutputRef string) error { job, ok := r.jobs[batchID] if !ok { @@ -589,6 +666,11 @@ func (r *fakeBatchImageRepository) BulkCreateBatchImageItems(ctx context.Context } func (r *fakeBatchImageRepository) ReplaceBatchImageItemsForJob(_ context.Context, batchID string, items []CreateBatchImageItemParams, counts BatchImageCounts) error { + // 与真实实现一致:仅 indexing 状态允许重建 item 表(未注册的 job 保持宽松, + // 供直接构造 job 的单测使用)。 + if job, ok := r.jobs[batchID]; ok && job.Status != BatchImageJobStatusIndexing { + return ErrBatchImageIndexStateConflict + } r.replaceCalls++ copied := append([]CreateBatchImageItemParams(nil), items...) for idx := range copied { diff --git a/backend/internal/service/batch_image_provider_gemini.go b/backend/internal/service/batch_image_provider_gemini.go index ccf91de714..64efb0b35d 100644 --- a/backend/internal/service/batch_image_provider_gemini.go +++ b/backend/internal/service/batch_image_provider_gemini.go @@ -17,6 +17,7 @@ import ( infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/geminicli" + "github.com/Wei-Shaw/sub2api/internal/pkg/httpclient" ) const defaultGeminiBatchRequeueAfter = 30 * time.Second @@ -460,11 +461,24 @@ func NewGeminiBatchHTTPClient(baseURL string, client *http.Client) *GeminiBatchH baseURL = geminicli.AIStudioBaseURL } if client == nil { - client = http.DefaultClient + client = batchImageDefaultHTTPClient() } return &GeminiBatchHTTPClient{baseURL: baseURL, client: client} } +// batchImageDefaultHTTPClient 返回带连接/握手/响应头超时的共享客户端。 +// 不设整体 Timeout:大文件上传与结果流式下载耗时不可预估, +// 但拨号、TLS、等待响应头必须有界,否则挂死的连接会无限占用提交路径。 +func batchImageDefaultHTTPClient() *http.Client { + client, err := httpclient.GetClient(httpclient.Options{ + ResponseHeaderTimeout: 60 * time.Second, + }) + if err != nil { + return http.DefaultClient + } + return client +} + func (c *GeminiBatchHTTPClient) UploadJSONL(ctx context.Context, apiKey string, displayName string, r io.Reader) (*GeminiUploadedFile, error) { var body bytes.Buffer writer := multipart.NewWriter(&body) @@ -567,6 +581,11 @@ func (c *GeminiBatchHTTPClient) DownloadFile(ctx context.Context, apiKey string, if downloadURL == "" { downloadURL = c.baseURL + "/v1beta/" + strings.TrimLeft(fileName, "/") + ":download" } + // 纵深加固:downloadUri 来自上游响应,跟随前校验目标 host, + // 防止异常/被劫持的响应把带 api key 的请求带到任意主机。 + if err := validateGeminiDownloadHost(downloadURL, c.baseURL); err != nil { + return nil, "", err + } req, err := http.NewRequestWithContext(ctx, http.MethodGet, downloadURL, nil) if err != nil { return nil, "", err @@ -643,6 +662,26 @@ func (c *GeminiBatchHTTPClient) newRequest(ctx context.Context, method, path, ap return req, nil } +// validateGeminiDownloadHost 只允许跟随到 googleapis.com(含子域) +// 或与配置的 baseURL 同 host 的下载地址。 +func validateGeminiDownloadHost(downloadURL, baseURL string) error { + parsed, err := url.Parse(downloadURL) + if err != nil { + return geminiProviderError("GEMINI_INVALID_RESPONSE", "Gemini download uri is invalid", err) + } + if parsed.Scheme != "https" { + return geminiProviderError("GEMINI_INVALID_RESPONSE", "Gemini download uri must use https", nil) + } + host := strings.ToLower(parsed.Hostname()) + if host == "googleapis.com" || strings.HasSuffix(host, ".googleapis.com") { + return nil + } + if base, err := url.Parse(baseURL); err == nil && strings.EqualFold(base.Hostname(), host) { + return nil + } + return geminiProviderError("GEMINI_INVALID_RESPONSE", "Gemini download uri host is not allowed", nil) +} + type GeminiAPIError struct { StatusCode int Code string diff --git a/backend/internal/service/batch_image_provider_vertex.go b/backend/internal/service/batch_image_provider_vertex.go index 9289c1484c..bb3363dfe3 100644 --- a/backend/internal/service/batch_image_provider_vertex.go +++ b/backend/internal/service/batch_image_provider_vertex.go @@ -736,7 +736,7 @@ type VertexBatchHTTPClient struct { func NewVertexBatchHTTPClient(baseURL string, client *http.Client) *VertexBatchHTTPClient { if client == nil { - client = http.DefaultClient + client = batchImageDefaultHTTPClient() } return &VertexBatchHTTPClient{baseURL: strings.TrimRight(strings.TrimSpace(baseURL), "/"), client: client} } @@ -794,7 +794,7 @@ type VertexGCSObjectStore struct { func NewVertexGCSObjectStore(baseURL string, client *http.Client) *VertexGCSObjectStore { if client == nil { - client = http.DefaultClient + client = batchImageDefaultHTTPClient() } baseURL = strings.TrimRight(strings.TrimSpace(baseURL), "/") if baseURL == "" { diff --git a/backend/internal/service/batch_image_public.go b/backend/internal/service/batch_image_public.go index 0b792aacb2..2960d4cd9f 100644 --- a/backend/internal/service/batch_image_public.go +++ b/backend/internal/service/batch_image_public.go @@ -14,6 +14,8 @@ import ( "github.com/Wei-Shaw/sub2api/internal/config" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "go.uber.org/zap" ) const ( @@ -202,6 +204,11 @@ func (s *BatchImagePublicService) Submit(ctx context.Context, owner BatchImageOw if err != nil { return nil, err } + // 与 ListModels 使用同一鉴权谓词(AllowBatchImageGeneration + Platform==Gemini), + // 避免两个入口校验口径不一致留下防御纵深缺口。 + if err := s.ensureGroupAllowsBatchImage(ctx, owner.GroupID); err != nil { + return nil, err + } requestHash := HashBatchImageSubmitRequest(normalized) idempotencyKey = strings.TrimSpace(idempotencyKey) if idempotencyKey != "" { @@ -319,7 +326,31 @@ func (s *BatchImagePublicService) Submit(ctx context.Context, owner BatchImageOw }) } + // 上游提交(上传参考图 + 创建批任务)可能长达数分钟且不刷新 updated_at, + // 会被 stale 恢复扫描误判为滞留并退款。提交前转入 uploading 刷新时间戳, + // 提交期间用心跳持续续期。 + if err := s.Repo.TransitionBatchImageJobStatus(ctx, job.BatchID, BatchImageJobStatusUploading, BatchImageTransitionOptions{ + EventType: "upload_started", + EventPayload: map[string]any{"batch_id": job.BatchID}, + }); err != nil { + if releaseErr := s.releaseFailedSubmitHold(ctx, job, requestHash); releaseErr != nil { + return nil, releaseErr + } + // 并发 Cancel 等导致的非法转换:job 已处于终态,不再覆盖其状态。 + if !errors.Is(err, ErrBatchImageInvalidTransition) { + _ = s.Repo.RecordBatchImageJobSubmitFailure(ctx, job.BatchID, "UPLOAD_TRANSITION_FAILED", sanitizeBatchImagePublicMessage(err.Error()), true) + s.hidePreUpstreamSubmitFailure(ctx, owner, job) + } + return nil, err + } + job.Status = BatchImageJobStatusUploading + + hbCtx, hbCancel := context.WithCancel(ctx) + hbDone := make(chan struct{}) + go s.runSubmitHeartbeat(hbCtx, job.BatchID, hbDone) providerJob, err := provider.Submit(ctx, job, account, input) + hbCancel() + <-hbDone if err != nil { if releaseErr := s.releaseFailedSubmitHold(ctx, job, requestHash); releaseErr != nil { return nil, releaseErr @@ -348,6 +379,9 @@ func (s *BatchImagePublicService) Submit(ctx context.Context, owner BatchImageOw GCSOutputURI: batchImageGCSRef(provider.Name(), providerJob.ProviderOutputRef), EventPayload: map[string]any{"provider": provider.Name()}, }); err != nil { + // job 可能已被恢复扫描转 failed 并退款:上游批任务已创建成功, + // 必须尽力取消并清理输入,否则上游照常产生成本(孤儿任务)。 + s.abortOrphanProviderJob(ctx, provider, job, account, providerJob) return nil, err } @@ -375,6 +409,75 @@ func (s *BatchImagePublicService) releaseFailedSubmitHold(ctx context.Context, j return nil } +// runSubmitHeartbeat 在 provider.Submit 期间周期性刷新 job 的 updated_at, +// 使 stale 恢复扫描能区分"仍在慢提交"与"进程死亡后的滞留"。 +func (s *BatchImagePublicService) runSubmitHeartbeat(ctx context.Context, batchID string, done chan<- struct{}) { + defer close(done) + interval := s.submitHeartbeatInterval() + ticker := time.NewTicker(interval) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + if err := s.Repo.TouchBatchImageJobSubmitting(ctx, batchID); err != nil && ctx.Err() == nil { + logger.L().Warn("batch_image.submit_heartbeat_failed", + zap.String("batch_id", batchID), + zap.Error(err), + ) + } + } + } +} + +func (s *BatchImagePublicService) submitHeartbeatInterval() time.Duration { + staleAfter := 10 * time.Minute + if s != nil && s.Config != nil && s.Config.BatchImage.StaleActiveAfterSeconds > 0 { + staleAfter = time.Duration(s.Config.BatchImage.StaleActiveAfterSeconds) * time.Second + } + interval := staleAfter / 3 + if interval < 15*time.Second { + interval = 15 * time.Second + } + return interval +} + +// abortOrphanProviderJob 在上游任务创建成功但本地状态推进失败时, +// 尽力取消上游批任务并清理已上传的输入文件,避免孤儿任务持续产生成本。 +func (s *BatchImagePublicService) abortOrphanProviderJob(ctx context.Context, provider BatchImageProvider, job *BatchImageJob, account *Account, providerJob *BatchProviderJob) { + if s == nil || provider == nil || job == nil || providerJob == nil { + return + } + orphan := *job + orphan.ProviderJobName = batchImageOptionalStringPtr(providerJob.ProviderJobName) + orphan.ProviderInputRef = batchImageOptionalStringPtr(providerJob.ProviderInputRef) + orphan.GCSInputURI = batchImageOptionalStringPtr(batchImageGCSRef(provider.Name(), providerJob.ProviderInputRef)) + if err := provider.Cancel(ctx, &orphan, account); err != nil { + logger.L().Warn("batch_image.orphan_provider_job_cancel_failed", + zap.String("batch_id", job.BatchID), + zap.String("provider", provider.Name()), + zap.Error(err), + ) + } + if err := provider.Cleanup(ctx, &orphan, account, CleanupTargetInput); err != nil { + logger.L().Warn("batch_image.orphan_provider_job_cleanup_failed", + zap.String("batch_id", job.BatchID), + zap.String("provider", provider.Name()), + zap.Error(err), + ) + } + if err := s.Repo.AppendBatchImageEvent(ctx, job.BatchID, "provider_job_aborted_after_submit", map[string]any{ + "batch_id": job.BatchID, + "provider": provider.Name(), + }); err != nil { + logger.L().Warn("batch_image.orphan_provider_job_event_failed", + zap.String("batch_id", job.BatchID), + zap.Error(err), + ) + } +} + func (s *BatchImagePublicService) createPendingItems(ctx context.Context, batchID, requestHash string, items []BatchImageSubmitItem) error { if s == nil || s.Repo == nil || len(items) == 0 { return nil @@ -399,10 +502,19 @@ func (s *BatchImagePublicService) enqueueBillingRetry(ctx context.Context, batch return } if err := s.Queue.Enqueue(ctx, batchID); err != nil && !errors.Is(err, ErrBatchImageAlreadyQueued) { - _ = s.Repo.AppendBatchImageEvent(ctx, batchID, "billing_retry_enqueue_failed", map[string]any{ + logger.L().Warn("batch_image.billing_retry_enqueue_failed", + zap.String("batch_id", batchID), + zap.Error(err), + ) + if eventErr := s.Repo.AppendBatchImageEvent(ctx, batchID, "billing_retry_enqueue_failed", map[string]any{ "batch_id": batchID, "error": sanitizeBatchImagePublicMessage(err.Error()), - }) + }); eventErr != nil { + logger.L().Warn("batch_image.billing_retry_event_failed", + zap.String("batch_id", batchID), + zap.Error(eventErr), + ) + } } } @@ -410,7 +522,12 @@ func (s *BatchImagePublicService) hidePreUpstreamSubmitFailure(ctx context.Conte if s == nil || s.Repo == nil || job == nil || job.ProviderJobName != nil { return } - _ = s.Repo.MarkBatchImageJobUserDeleted(ctx, owner.UserID, owner.APIKeyID, job.BatchID, time.Now()) + if err := s.Repo.MarkBatchImageJobUserDeleted(ctx, owner.UserID, owner.APIKeyID, job.BatchID, time.Now()); err != nil { + logger.L().Warn("batch_image.hide_pre_upstream_failure_failed", + zap.String("batch_id", job.BatchID), + zap.Error(err), + ) + } } func (s *BatchImagePublicService) Get(ctx context.Context, owner BatchImageOwner, batchID string) (*BatchImagePublicBatch, error) { @@ -615,7 +732,12 @@ func (s *BatchImagePublicService) Cancel(ctx context.Context, owner BatchImageOw if err := provider.Cancel(ctx, job, account); err != nil { return nil, ErrBatchImageCancelFailed } - _ = s.Repo.AppendBatchImageEvent(ctx, job.BatchID, "job_cancel_requested", map[string]any{"batch_id": job.BatchID}) + if eventErr := s.Repo.AppendBatchImageEvent(ctx, job.BatchID, "job_cancel_requested", map[string]any{"batch_id": job.BatchID}); eventErr != nil { + logger.L().Warn("batch_image.cancel_event_failed", + zap.String("batch_id", job.BatchID), + zap.Error(eventErr), + ) + } if s.Queue != nil { if err := s.Queue.Enqueue(ctx, job.BatchID); err != nil && !errors.Is(err, ErrBatchImageAlreadyQueued) { return nil, ErrBatchImageCancelFailed @@ -930,6 +1052,16 @@ func (s *BatchImagePublicService) resolvePricingSnapshot(ctx context.Context, ow } unit = resolvedUnit } + // 定价不变式:hold 比例不得低于 discount 比例,否则成功率足够高时 + // actualCost > holdAmount,结算永远失败、冻结余额无法解冻。 + // 管理端已校验新配置,此处兜底钳制存量脏数据。 + if holdMultiplier < discountMultiplier { + logger.L().Warn("batch_image.hold_multiplier_below_discount_clamped", + zap.Float64("hold_multiplier", holdMultiplier), + zap.Float64("discount_multiplier", discountMultiplier), + ) + holdMultiplier = discountMultiplier + } accountMultiplier := 1.0 if account != nil { accountMultiplier = account.BillingRateMultiplier() diff --git a/backend/internal/service/batch_image_public_test.go b/backend/internal/service/batch_image_public_test.go index d59784269f..2e5b930a7b 100644 --- a/backend/internal/service/batch_image_public_test.go +++ b/backend/internal/service/batch_image_public_test.go @@ -93,10 +93,12 @@ func TestBatchImagePublicService_Submit(t *testing.T) { require.InDelta(t, 0.5, job.GroupRateMultiplier, 1e-12) require.InDelta(t, 1.25, job.AccountRateMultiplier, 1e-12) require.InDelta(t, 0.8, job.BatchDiscountMultiplier, 1e-12) - require.InDelta(t, 0.6, job.HoldMultiplier, 1e-12) + // 配置的 hold(0.6) < discount(0.8) 属于会导致结算死锁的脏数据, + // 快照时被钳制为 discount,保证 holdAmount >= 实际成本上限。 + require.InDelta(t, 0.8, job.HoldMultiplier, 1e-12) require.InDelta(t, 0.125, job.BillableUnitPrice, 1e-12) - require.InDelta(t, 0.09375, job.HoldUnitPrice, 1e-12) - require.InDelta(t, 0.1875, *job.HoldAmount, 1e-12) + require.InDelta(t, 0.125, job.HoldUnitPrice, 1e-12) + require.InDelta(t, 0.25, *job.HoldAmount, 1e-12) }) t.Run("uses configured group 1k image price for batch image base price", func(t *testing.T) { diff --git a/backend/internal/service/batch_image_settlement.go b/backend/internal/service/batch_image_settlement.go index cbd1ca7ae6..26d2abb0fd 100644 --- a/backend/internal/service/batch_image_settlement.go +++ b/backend/internal/service/batch_image_settlement.go @@ -11,6 +11,8 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "go.uber.org/zap" ) const ( @@ -98,28 +100,40 @@ func (s *BatchImageSettlementService) Settle(ctx context.Context, batchID string if job.Status != BatchImageJobStatusSettling { return nil, ErrBatchImageSettlementInvalidStatus } - if job.SuccessCount < 0 || job.FailCount < 0 || job.ItemCount < 0 || job.SuccessCount+job.FailCount > job.ItemCount { - return nil, ErrBatchImageSettlementInvalidCounts - } - if strings.TrimSpace(batchImageDerefString(job.ManifestHash)) != "" && batchImageDerefString(job.ManifestHash) != manifestHash { - return nil, ErrBatchImageSettlementManifestConflict - } if job.APIKeyID == nil || *job.APIKeyID <= 0 { return nil, ErrBatchImageSettlementMissingAPIKeyID } if job.AccountID == nil || *job.AccountID <= 0 { return nil, ErrBatchImageSettlementMissingAccountID } + // 重试耗尽检查必须先于各类可重复失败的校验(counts/manifest/定价/超冻结), + // 否则这些错误路径会绕过耗尽出口,settling job 无限 requeue、冻结余额永不释放。 if isBatchImageSettlementRetryExhausted(job) { - return nil, s.failExhaustedSettlement(ctx, job, manifestHash, "settlement billing retry limit reached") + return nil, s.failExhaustedSettlement(ctx, job, "settlement retry limit reached: "+batchImageDerefString(job.LastErrorCode)) + } + if job.SuccessCount < 0 || job.FailCount < 0 || job.ItemCount < 0 || job.SuccessCount+job.FailCount > job.ItemCount { + if failErr := s.recordSettlementFailure(ctx, job, "SETTLEMENT_INVALID_COUNTS", + fmt.Sprintf("success=%d fail=%d item_count=%d", job.SuccessCount, job.FailCount, job.ItemCount)); failErr != nil { + return nil, failErr + } + return nil, ErrBatchImageSettlementInvalidCounts + } + if strings.TrimSpace(batchImageDerefString(job.ManifestHash)) != "" && batchImageDerefString(job.ManifestHash) != manifestHash { + if failErr := s.recordSettlementFailure(ctx, job, "SETTLEMENT_MANIFEST_CONFLICT", "manifest hash conflict"); failErr != nil { + return nil, failErr + } + return nil, ErrBatchImageSettlementManifestConflict } unitPrice, err := s.settlementUnitPrice(ctx, job) - if err != nil { - return nil, err + if err == nil && unitPrice < 0 { + err = ErrBatchImageSettlementPricingMissing } - if unitPrice < 0 { - return nil, ErrBatchImageSettlementPricingMissing + if err != nil { + if failErr := s.recordSettlementFailure(ctx, job, "SETTLEMENT_PRICING_MISSING", err.Error()); failErr != nil { + return nil, failErr + } + return nil, err } actualCost := float64(job.SuccessCount) * unitPrice result.ActualCost = actualCost @@ -129,16 +143,16 @@ func (s *BatchImageSettlementService) Settle(ctx context.Context, batchID string } if actualCost-holdAmount > batchImageCostEpsilon { msg := fmt.Sprintf("actual cost %.10f exceeds held amount %.10f", actualCost, holdAmount) - _, _ = s.Repo.SetBatchImageJobSettlementFailed(ctx, job.BatchID, "SETTLEMENT_COST_EXCEEDS_HOLD", msg) + if failErr := s.recordSettlementFailure(ctx, job, "SETTLEMENT_COST_EXCEEDS_HOLD", msg); failErr != nil { + return nil, failErr + } return nil, ErrBatchImageSettlementCostExceedsHold } if err := captureBatchImageBalanceHold(ctx, s.BillingRepo, job, actualCost, manifestHash); err != nil { msg := truncateBatchImageMessage(err.Error(), batchImageMaxErrorMessageLength) - retryCount, recordErr := s.Repo.SetBatchImageJobSettlementFailed(ctx, job.BatchID, "SETTLEMENT_BILLING_FAILED", msg) - if recordErr == nil && retryCount >= batchImageSettlementMaxRetries { - job.RetryCount = retryCount - return nil, s.failExhaustedSettlement(ctx, job, manifestHash, msg) + if failErr := s.recordSettlementFailure(ctx, job, "SETTLEMENT_BILLING_FAILED", msg); failErr != nil { + return nil, failErr } return nil, err } @@ -168,20 +182,52 @@ func (s *BatchImageSettlementService) Settle(ctx context.Context, batchID string return result, nil } +// isBatchImageSettlementRetryExhausted 判断 settling job 是否已达重试上限。 +// 必须覆盖所有 SETTLEMENT_* 失败码(而非仅 SETTLEMENT_BILLING_FAILED), +// 否则 SETTLEMENT_COST_EXCEEDS_HOLD / SETTLEMENT_INVALID_COUNTS 等错误会无限 requeue。 func isBatchImageSettlementRetryExhausted(job *BatchImageJob) bool { return job != nil && job.Status == BatchImageJobStatusSettling && job.RetryCount >= batchImageSettlementMaxRetries && - batchImageDerefString(job.LastErrorCode) == "SETTLEMENT_BILLING_FAILED" + strings.HasPrefix(batchImageDerefString(job.LastErrorCode), "SETTLEMENT_") } -func (s *BatchImageSettlementService) failExhaustedSettlement(ctx context.Context, job *BatchImageJob, manifestHash, message string) error { +// recordSettlementFailure 记录一次结算失败并递增 retry_count。 +// 重试达到上限时立即走耗尽出口(释放冻结余额并转 failed); +// 返回非 nil 时调用方应直接返回该错误。 +func (s *BatchImageSettlementService) recordSettlementFailure(ctx context.Context, job *BatchImageJob, code, message string) error { + retryCount, recordErr := s.Repo.SetBatchImageJobSettlementFailed(ctx, job.BatchID, code, truncateBatchImageMessage(message, batchImageMaxErrorMessageLength)) + if recordErr != nil { + logger.L().Warn("batch_image.settlement_failure_record_failed", + zap.String("batch_id", job.BatchID), + zap.String("code", code), + zap.Error(recordErr), + ) + return nil + } + job.RetryCount = retryCount + job.LastErrorCode = &code + if retryCount >= batchImageSettlementMaxRetries { + return s.failExhaustedSettlement(ctx, job, message) + } + return nil +} + +func (s *BatchImageSettlementService) failExhaustedSettlement(ctx context.Context, job *BatchImageJob, message string) error { if s == nil || s.Repo == nil { return ErrBatchImageSettlementBillingFailed } - if err := releaseBatchImageBalanceHold(ctx, s.BillingRepo, job, manifestHash); err != nil { + // 释放指纹必须与其余所有释放点(processor/Cancel/recovery)一致地使用 RequestHash: + // 它们共享同一 request id,payloadHash 不同会触发 ErrUsageBillingRequestConflict, + // 导致后续 Cancel/重试永远失败、terminal job 变成毒消息。 + if err := releaseBatchImageBalanceHold(ctx, s.BillingRepo, job, batchImageDerefString(job.RequestHash)); err != nil { msg := truncateBatchImageMessage(err.Error(), batchImageMaxErrorMessageLength) - _, _ = s.Repo.SetBatchImageJobSettlementFailed(ctx, job.BatchID, "SETTLEMENT_RELEASE_FAILED", msg) + if _, recordErr := s.Repo.SetBatchImageJobSettlementFailed(ctx, job.BatchID, "SETTLEMENT_RELEASE_FAILED", msg); recordErr != nil { + logger.L().Warn("batch_image.settlement_release_failure_record_failed", + zap.String("batch_id", job.BatchID), + zap.Error(recordErr), + ) + } return ErrBatchImageSettlementBillingFailed.WithCause(err) } s.invalidateAuthCache(ctx, job.UserID) diff --git a/backend/internal/service/batch_image_settlement_test.go b/backend/internal/service/batch_image_settlement_test.go index 8837a85358..a1fe3e899d 100644 --- a/backend/internal/service/batch_image_settlement_test.go +++ b/backend/internal/service/batch_image_settlement_test.go @@ -278,6 +278,73 @@ func TestBatchImageSettlementRetryExhaustedReleaseIsIdempotentAfterTransitionFai require.Len(t, billing.seen, 1) } +func TestBatchImageSettlementService_CostExceedsHoldExhaustsAndReleases(t *testing.T) { + repo := newFakeBatchImageRepository() + job := testSettlingBatchImageJob("imgbatch_over_hold_exhausted") + job.SuccessCount = 2 + job.FailCount = 0 + job.ItemCount = 2 + holdAmount := 0.5 + job.HoldAmount = &holdAmount + job.EstimatedCost = holdAmount + requestHash := "request-hash-over-hold" + job.RequestHash = &requestHash + repo.jobs[job.BatchID] = job + billing := &fakeBatchImageBillingRepo{} + svc := &BatchImageSettlementService{Repo: repo, BillingRepo: billing, Pricing: &fakeBatchImagePricingResolver{unitPrice: 0.50}} + + // 前 N-1 次:记录失败并返回错误(等待 worker 重试)。 + for i := 0; i < batchImageSettlementMaxRetries-1; i++ { + _, err := svc.Settle(context.Background(), job.BatchID) + require.ErrorIs(t, err, ErrBatchImageSettlementCostExceedsHold) + require.Equal(t, BatchImageJobStatusSettling, repo.jobs[job.BatchID].Status) + } + // 达到上限:必须走耗尽出口释放冻结并转 failed,而不是无限 requeue。 + _, err := svc.Settle(context.Background(), job.BatchID) + require.ErrorIs(t, err, ErrBatchImageSettlementBillingFailed) + require.Equal(t, BatchImageJobStatusFailed, repo.jobs[job.BatchID].Status) + require.Empty(t, billing.captures) + require.Len(t, billing.releases, 1) + require.Equal(t, BatchImageReleaseRequestID(job.BatchID), billing.releases[0].RequestID) + // 释放指纹必须与 processor/Cancel/recovery 一致地使用 RequestHash, + // 否则共享同一 request id 的后续释放会命中指纹冲突(毒消息)。 + require.Equal(t, requestHash, billing.releases[0].RequestPayloadHash) +} + +func TestBatchImageSettlementService_InvalidCountsExhaustsAndReleases(t *testing.T) { + repo := newFakeBatchImageRepository() + job := testSettlingBatchImageJob("imgbatch_bad_counts_exhausted") + job.SuccessCount = 2 + job.FailCount = 2 + job.ItemCount = 3 + requestHash := "request-hash-bad-counts" + job.RequestHash = &requestHash + repo.jobs[job.BatchID] = job + billing := &fakeBatchImageBillingRepo{} + svc := &BatchImageSettlementService{Repo: repo, BillingRepo: billing, Pricing: &fakeBatchImagePricingResolver{unitPrice: 0.25}} + + for i := 0; i < batchImageSettlementMaxRetries-1; i++ { + _, err := svc.Settle(context.Background(), job.BatchID) + require.ErrorIs(t, err, ErrBatchImageSettlementInvalidCounts) + } + _, err := svc.Settle(context.Background(), job.BatchID) + require.ErrorIs(t, err, ErrBatchImageSettlementBillingFailed) + require.Equal(t, BatchImageJobStatusFailed, repo.jobs[job.BatchID].Status) + require.Empty(t, billing.captures) + require.Len(t, billing.releases, 1) + require.Equal(t, requestHash, billing.releases[0].RequestPayloadHash) +} + +func TestReleaseBatchImageBalanceHold_TreatsFingerprintConflictAsReleased(t *testing.T) { + job := testSettlingBatchImageJob("imgbatch_release_conflict") + // 历史版本用 manifestHash 释放过一次:同一 request id 再以 RequestHash + // 释放会命中指纹冲突。资金已归还,必须视为幂等成功而非毒消息。 + billing := &fakeBatchImageBillingRepo{releaseErr: ErrUsageBillingRequestConflict} + err := releaseBatchImageBalanceHold(context.Background(), billing, job, "request-hash") + require.NoError(t, err) + require.Len(t, billing.releases, 1) +} + func TestBatchImageSettlementManifestHash(t *testing.T) { job := testSettlingBatchImageJob("imgbatch_hash") first := BuildBatchImageSettlementManifestHash(job) diff --git a/backend/internal/service/batch_image_worker.go b/backend/internal/service/batch_image_worker.go index 5027350689..f090a67a52 100644 --- a/backend/internal/service/batch_image_worker.go +++ b/backend/internal/service/batch_image_worker.go @@ -150,13 +150,23 @@ func (w *BatchImageWorker) RunOnce(ctx context.Context) error { return err } if !ok { - return nil + // 锁被其他实例持有:按冲突延迟重新入队。直接丢弃会让 job 滞留在 + // active zset,最早要等 StaleActiveAfter 才被恢复,造成分钟级停摆。 + return w.queue.RequeueAfter(ctx, reserved.BatchID, w.opts.LockConflictDelay) } defer func() { _ = lock.Release(ctx) }() + // 处理期间持续心跳:刷新 active zset 时间戳防止 stale 恢复把在处理的 + // job 重投给其他 worker,并对支持续期的锁实现延长锁 TTL。 + hbStop := make(chan struct{}) + hbDone := make(chan struct{}) + go w.runJobHeartbeat(ctx, reserved.BatchID, lock, hbStop, hbDone) + result, err := w.processor.Process(ctx, reserved.BatchID) + close(hbStop) + <-hbDone if err != nil { logger.L().Warn("batch_image.worker_process_failed", zap.String("batch_id", reserved.BatchID), @@ -174,6 +184,52 @@ func (w *BatchImageWorker) RunOnce(ctx context.Context) error { return w.queue.RequeueAfter(ctx, reserved.BatchID, delay) } +// BatchImageJobLockRefresher 是可选的锁续期能力;由具体锁实现按需提供。 +type BatchImageJobLockRefresher interface { + Refresh(ctx context.Context, ttl time.Duration) error +} + +func (w *BatchImageWorker) heartbeatInterval() time.Duration { + interval := w.opts.JobLockTTL + if w.opts.StaleActiveAfter < interval { + interval = w.opts.StaleActiveAfter + } + interval /= 3 + if interval < time.Second { + interval = time.Second + } + return interval +} + +func (w *BatchImageWorker) runJobHeartbeat(ctx context.Context, batchID string, lock BatchImageJobLock, stop <-chan struct{}, done chan<- struct{}) { + defer close(done) + ticker := time.NewTicker(w.heartbeatInterval()) + defer ticker.Stop() + for { + select { + case <-stop: + return + case <-ctx.Done(): + return + case <-ticker.C: + if err := w.queue.Heartbeat(ctx, batchID); err != nil && ctx.Err() == nil { + logger.L().Warn("batch_image.worker_heartbeat_failed", + zap.String("batch_id", batchID), + zap.Error(err), + ) + } + if refresher, ok := lock.(BatchImageJobLockRefresher); ok { + if err := refresher.Refresh(ctx, w.opts.JobLockTTL); err != nil && ctx.Err() == nil { + logger.L().Warn("batch_image.worker_lock_refresh_failed", + zap.String("batch_id", batchID), + zap.Error(err), + ) + } + } + } + } +} + func (w *BatchImageWorker) MoveDueDelayedOnce(ctx context.Context) (int, error) { if w == nil || w.queue == nil { return 0, nil diff --git a/backend/internal/service/batch_image_worker_runtime.go b/backend/internal/service/batch_image_worker_runtime.go index de3b4cb42a..f5373cb6d6 100644 --- a/backend/internal/service/batch_image_worker_runtime.go +++ b/backend/internal/service/batch_image_worker_runtime.go @@ -53,6 +53,7 @@ func ProvideBatchImageWorkerRuntime( Repo: repo, Billing: billingRepo, AuthCache: authCache, + Queue: queue, StaleAfter: NewBatchImageWorkerOptionsFromConfig(cfg).StaleActiveAfter, Limit: NewBatchImageWorkerOptionsFromConfig(cfg).RecoverLimit, } diff --git a/backend/internal/service/batch_image_worker_test.go b/backend/internal/service/batch_image_worker_test.go index 934cd9fd1c..016df7e713 100644 --- a/backend/internal/service/batch_image_worker_test.go +++ b/backend/internal/service/batch_image_worker_test.go @@ -56,15 +56,18 @@ func TestBatchImageWorker_RequeuesOnProcessorError(t *testing.T) { require.Empty(t, queue.acked) } -func TestBatchImageWorker_SkipsWhenJobLockNotAcquired(t *testing.T) { +func TestBatchImageWorker_RequeuesWhenJobLockNotAcquired(t *testing.T) { queue := newFakeBatchImageQueue("imgbatch_worker_locked") queue.lockAcquired = false processor := &fakeBatchImageProcessor{} - worker := NewBatchImageWorker(queue, processor, BatchImageWorkerOptions{}) + worker := NewBatchImageWorker(queue, processor, BatchImageWorkerOptions{LockConflictDelay: 3 * time.Second}) + // 锁冲突必须按冲突延迟重新入队;直接丢弃会让 job 滞留 active zset, + // 要等 StaleActiveAfter(默认 10 分钟)才被恢复。 require.NoError(t, worker.RunOnce(context.Background())) require.Empty(t, processor.processed) - require.Empty(t, queue.requeued) + require.Len(t, queue.requeued, 1) + require.Equal(t, 3*time.Second, queue.requeued[0].delay) require.Empty(t, queue.acked) } diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index dc54a1b1f3..4c265aed3c 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -713,6 +713,14 @@ func (s *BillingService) GetModelPricing(model string) (*ModelPricing, error) { // 1. 优先从动态价格服务获取 if s.pricingService != nil { litellmPricing := s.pricingService.GetModelPricing(model) + // 仅有图片价、无 token 价的条目(如 LiteLLM 的 imagen 类模型)不能用于 + // token 计费:直接返回会把 token 流量按 $0 计费。跳过后走 fallback, + // 无 fallback 则 fail-closed(ErrModelPricingUnavailable)。 + // 图片计费路径(getDefaultImagePrice / getImageUnitPrice)直接读 + // PricingService,不受影响。 + if litellmPricing != nil && litellmPricing.TokenPricingAbsent { + litellmPricing = nil + } if litellmPricing != nil { // 启用 5m/1h 分类计费的条件: // 1. 存在 1h 价格 diff --git a/backend/internal/service/pricing_service.go b/backend/internal/service/pricing_service.go index f3fb36d73a..2ae15df507 100644 --- a/backend/internal/service/pricing_service.go +++ b/backend/internal/service/pricing_service.go @@ -73,6 +73,11 @@ type LiteLLMModelPricing struct { SupportsPromptCaching bool `json:"supports_prompt_caching"` OutputCostPerImage float64 `json:"output_cost_per_image"` // 图片生成模型每张图片价格 OutputCostPerImageToken float64 `json:"output_cost_per_image_token"` // 图片输出 token 价格 + + // TokenPricingAbsent 表示源数据中 input/output token 价格均缺失(仅有图片价)。 + // 此类条目只可用于图片计费,token 计费必须回退到 fallback 或 fail-closed, + // 否则 token 流量会被按 $0 计费。零值(false)表示条目具备 token 价格。 + TokenPricingAbsent bool `json:"-"` } // PricingRemoteClient 远程价格数据获取接口 @@ -383,6 +388,7 @@ func (s *PricingService) parsePricingData(body []byte) (map[string]*LiteLLMModel Mode: entry.Mode, SupportsPromptCaching: entry.SupportsPromptCaching, SupportsServiceTier: entry.SupportsServiceTier, + TokenPricingAbsent: entry.InputCostPerToken == nil && entry.OutputCostPerToken == nil, } if entry.InputCostPerToken != nil { diff --git a/backend/internal/service/pricing_service_test.go b/backend/internal/service/pricing_service_test.go index 11c1b58da9..4bf8f2379e 100644 --- a/backend/internal/service/pricing_service_test.go +++ b/backend/internal/service/pricing_service_test.go @@ -54,6 +54,44 @@ func TestParsePricingData_KeepsImageOnlyPricing(t *testing.T) { require.NotNil(t, pricing) require.InDelta(t, 0.034, pricing.OutputCostPerImage, 1e-12) require.Equal(t, "image_generation", pricing.Mode) + // 仅有图片价的条目必须标记 token 价缺失,供 token 计费路径 fail-closed。 + require.True(t, pricing.TokenPricingAbsent) +} + +func TestBillingService_GetModelPricing_FailsClosedForImageOnlyEntries(t *testing.T) { + pricingSvc := &PricingService{} + data, err := pricingSvc.parsePricingData([]byte(`{ + "imagen-9.0-generate": { + "output_cost_per_image": 0.04, + "litellm_provider": "vertex_ai-image-models", + "mode": "image_generation" + }, + "gemini-image-with-token-price": { + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "output_cost_per_image": 0.034, + "litellm_provider": "vertex_ai-language-models", + "mode": "image_generation" + } + }`)) + require.NoError(t, err) + pricingSvc.pricingData = data + billingSvc := NewBillingService(&config.Config{}, pricingSvc) + + // image-only 条目不得进入 token 计费(否则 token 流量按 $0 计费), + // 必须落到 fallback / ErrModelPricingUnavailable 的 fail-closed 路径。 + _, err = billingSvc.GetModelPricing("imagen-9.0-generate") + require.ErrorIs(t, err, ErrModelPricingUnavailable) + + // 显式 0 token 价的免费条目保持历史行为:正常返回。 + pricing, err := billingSvc.GetModelPricing("gemini-image-with-token-price") + require.NoError(t, err) + require.Zero(t, pricing.InputPricePerToken) + + // 图片计费路径不受影响:仍能读到 image-only 条目的图片单价。 + raw := pricingSvc.GetModelPricing("imagen-9.0-generate") + require.NotNil(t, raw) + require.InDelta(t, 0.04, raw.OutputCostPerImage, 1e-12) } func TestPricingService_MergesFallbackOnlyModels(t *testing.T) { From b5604afd3d7ecb4b79f9b5348391a0fe11dae02a Mon Sep 17 00:00:00 2001 From: shaw Date: Tue, 7 Jul 2026 20:13:28 +0800 Subject: [PATCH 2/2] =?UTF-8?q?test(batch-image):=20=E9=9B=86=E6=88=90?= =?UTF-8?q?=E6=B5=8B=E8=AF=95=E9=80=82=E9=85=8D=20ReplaceBatchImageItemsFo?= =?UTF-8?q?rJob=20=E7=9A=84=20indexing=20=E7=8A=B6=E6=80=81=E5=AE=88?= =?UTF-8?q?=E5=8D=AB?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 守卫要求 job 处于 indexing 状态才允许重建 item 表;原测试在 created 状态下直接调用被 409 拒绝。改为先断言非 indexing 状态的冲突行为, 再经 created→submitted→indexing 合法链推进后执行原有断言。 --- .../repository/batch_image_repo_integration_test.go | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/backend/internal/repository/batch_image_repo_integration_test.go b/backend/internal/repository/batch_image_repo_integration_test.go index 174e0657e3..df6be4a954 100644 --- a/backend/internal/repository/batch_image_repo_integration_test.go +++ b/backend/internal/repository/batch_image_repo_integration_test.go @@ -208,6 +208,16 @@ func TestBatchImageRepository_ReplaceBatchImageItemsForJob(t *testing.T) { }) require.NoError(t, err) + // 非 indexing 状态不允许重建 item 表:防止锁过期后掉队的 worker + // 重写已完成/已结算 job 的条目。 + err = repo.ReplaceBatchImageItemsForJob(ctx, batchID, []service.CreateBatchImageItemParams{ + {CustomID: "old", Status: service.BatchImageItemStatusSuccess, SourceLineNumber: &lineOne, ImageCount: 1}, + }, service.BatchImageCounts{SuccessCount: 1}) + require.ErrorIs(t, err, service.ErrBatchImageIndexStateConflict) + + require.NoError(t, repo.TransitionBatchImageJobStatus(ctx, batchID, service.BatchImageJobStatusSubmitted, service.BatchImageTransitionOptions{})) + require.NoError(t, repo.TransitionBatchImageJobStatus(ctx, batchID, service.BatchImageJobStatusIndexing, service.BatchImageTransitionOptions{})) + err = repo.ReplaceBatchImageItemsForJob(ctx, batchID, []service.CreateBatchImageItemParams{ {CustomID: "old", Status: service.BatchImageItemStatusSuccess, SourceLineNumber: &lineOne, ImageCount: 1}, }, service.BatchImageCounts{SuccessCount: 1})