-
+
{{ t('admin.accounts.selectedModels', { count: allowedModels.length }) }}
{{
@@ -1562,7 +1562,7 @@
-
+
{{ t('admin.accounts.selectedModels', { count: allowedModels.length }) }}
{{ t('admin.accounts.supportsAllModels') }}
@@ -1813,7 +1813,7 @@
-
+
{{ t('admin.accounts.selectedModels', { count: allowedModels.length }) }}
{{
@@ -3361,6 +3361,17 @@ const accountCategory = ref<'oauth-based' | 'apikey' | 'bedrock' | 'service_acco
const addMethod = ref('oauth') // For oauth-based: 'oauth' or 'setup-token'
const apiKeyBaseUrl = ref('https://api.anthropic.com')
const apiKeyValue = ref('')
+
+const syncPreviewCredentials = computed(() => {
+ if (!apiKeyValue.value) return undefined
+ return {
+ platform: form.platform,
+ type: form.type,
+ base_url: apiKeyBaseUrl.value || undefined,
+ api_key: apiKeyValue.value
+ }
+})
+
const editQuotaLimit = ref(null)
const editQuotaDailyLimit = ref(null)
const editQuotaWeeklyLimit = ref(null)
diff --git a/frontend/src/components/account/ModelWhitelistSelector.vue b/frontend/src/components/account/ModelWhitelistSelector.vue
index 9a0d6af8..d4d726d0 100644
--- a/frontend/src/components/account/ModelWhitelistSelector.vue
+++ b/frontend/src/components/account/ModelWhitelistSelector.vue
@@ -133,6 +133,7 @@ import { ref, computed } from 'vue'
import { useI18n } from 'vue-i18n'
import { useAppStore } from '@/stores/app'
import { accountsAPI } from '@/api/admin/accounts'
+import type { SyncUpstreamPreviewParams } from '@/api/admin/accounts'
import ModelIcon from '@/components/common/ModelIcon.vue'
import Icon from '@/components/icons/Icon.vue'
import { allModels, getModelsByPlatform } from '@/composables/useModelWhitelist'
@@ -144,6 +145,12 @@ const props = defineProps<{
platform?: string
platforms?: string[]
accountId?: number
+ syncCredentials?: {
+ platform: string
+ type: string
+ base_url?: string
+ api_key: string
+ }
}>()
const emit = defineEmits<{
@@ -176,9 +183,14 @@ const normalizedPlatforms = computed(() => {
const upstreamSyncPlatforms = new Set(['anthropic', 'openai', 'gemini', 'antigravity'])
const canSyncUpstream = computed(() => {
- if (!props.accountId) return false
- if (normalizedPlatforms.value.length === 0) return true
- return normalizedPlatforms.value.some(platform => upstreamSyncPlatforms.has(platform.toLowerCase()))
+ if (props.accountId) {
+ if (normalizedPlatforms.value.length === 0) return true
+ return normalizedPlatforms.value.some(platform => upstreamSyncPlatforms.has(platform.toLowerCase()))
+ }
+ if (props.syncCredentials) {
+ return upstreamSyncPlatforms.has(props.syncCredentials.platform.toLowerCase())
+ }
+ return false
})
const availableOptions = computed(() => {
@@ -249,11 +261,20 @@ const fillRelated = () => {
}
const syncUpstreamModels = async () => {
- if (!props.accountId || isSyncingUpstream.value) return
+ if (isSyncingUpstream.value) return
+ if (!props.accountId && !props.syncCredentials) return
isSyncingUpstream.value = true
try {
- const result = await accountsAPI.syncUpstreamModels(props.accountId)
+ let result
+ if (props.accountId) {
+ result = await accountsAPI.syncUpstreamModels(props.accountId)
+ } else if (props.syncCredentials) {
+ result = await accountsAPI.syncUpstreamModelsPreview(props.syncCredentials as SyncUpstreamPreviewParams)
+ } else {
+ return
+ }
+
const upstreamModels = result.models.map(model => model.trim()).filter(Boolean)
if (upstreamModels.length === 0) {
appStore.showInfo(t('admin.accounts.syncUpstreamModelsEmpty'))
From b60d8bb4cca030fa0f8678b34bd408f14d7392e1 Mon Sep 17 00:00:00 2001
From: DaydreamCoding <22166516+DaydreamCoding@users.noreply.github.com>
Date: Thu, 28 May 2026 17:41:29 +0800
Subject: [PATCH 46/79] =?UTF-8?q?feat(usage):=20=E5=9C=A8=20/admin/usage?=
=?UTF-8?q?=20=E6=94=AF=E6=8C=81=E6=9F=A5=E7=9C=8B=E5=B7=B2=E5=88=A0?=
=?UTF-8?q?=E9=99=A4=E7=94=A8=E6=88=B7=E7=9A=84=E5=8E=86=E5=8F=B2=E4=BD=BF?=
=?UTF-8?q?=E7=94=A8=E6=83=85=E5=86=B5?=
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit
用户软删除后使用记录仍在,但身份(邮箱)被 ent 软删除拦截器隐藏。本次在
三条管理员只读路径定点穿透软删除过滤,并把删除状态传播到前端标记,零新表/
迁移/回填:
- 后端穿透:富化 usage 日志(loadUsers)、用户搜索(ListWithFilters +
UserListFilters.IncludeDeleted)、点击详情(GetByIDIncludeDeleted /
GetUserIncludeDeleted + getById ?include_deleted 分支)
- 状态传播:service.User / dto.User 新增 DeletedAt;SearchUsers 标记 deleted
- 前端:表格与余额弹窗展示"已删除"徽标、筛选下拉标注并排序、点击走
include_deleted;新增 i18n admin.usage.userDeletedBadge
- 安全:普通用户 usage 仅查本人(无 PII 泄漏);主用户列表与默认 getById
行为不变(已删用户仍 404);仅 admin 搜索设 IncludeDeleted
后端 build / 三态 vet / unit 全量 / 仓储集成全绿;前端 typecheck / vitest /
改动文件 eslint 全清。
Co-Authored-By: Claude Opus 4.7 (1M context)
---
.../handler/admin/admin_service_stub_test.go | 4 +
.../internal/handler/admin/usage_handler.go | 14 +-
.../admin/usage_handler_search_users_test.go | 56 ++++++
.../internal/handler/admin/user_handler.go | 7 +-
.../admin/user_handler_get_deleted_test.go | 51 ++++++
.../handler/auth_oauth_pending_flow_test.go | 4 +
backend/internal/handler/dto/mappers.go | 1 +
.../handler/dto/mappers_deleted_user_test.go | 20 +++
backend/internal/handler/dto/types.go | 1 +
backend/internal/handler/user_handler_test.go | 3 +
backend/internal/repository/api_key_repo.go | 1 +
backend/internal/repository/usage_log_repo.go | 4 +-
..._log_repo_deleted_user_integration_test.go | 65 +++++++
backend/internal/repository/user_repo.go | 28 ++-
...r_repo_include_deleted_integration_test.go | 69 +++++++
backend/internal/server/api_contract_test.go | 4 +
.../server/middleware/admin_auth_test.go | 4 +
backend/internal/service/admin_service.go | 5 +
.../service/admin_service_apikey_test.go | 11 +-
.../service/admin_service_delete_test.go | 4 +
.../admin_service_email_identity_sync_test.go | 11 +-
.../service/admin_service_get_deleted_test.go | 22 +++
.../service/auth_service_email_bind_test.go | 3 +
.../service/content_moderation_test.go | 4 +
backend/internal/service/user.go | 1 +
backend/internal/service/user_service.go | 5 +
backend/internal/service/user_service_test.go | 4 +
frontend/src/api/admin/usage.ts | 1 +
frontend/src/api/admin/users.ts | 6 +-
.../components/admin/usage/UsageFilters.vue | 5 +-
.../src/components/admin/usage/UsageTable.vue | 3 +
.../usage/__tests__/UsageFilters.spec.ts | 168 ++++++++++++++++++
.../admin/usage/__tests__/UsageTable.spec.ts | 89 ++++++++++
.../admin/user/UserBalanceHistoryModal.vue | 3 +
frontend/src/i18n/locales/en.ts | 1 +
frontend/src/i18n/locales/zh.ts | 1 +
frontend/src/types/index.ts | 1 +
frontend/src/views/admin/UsageView.vue | 2 +-
.../views/admin/__tests__/UsageView.spec.ts | 60 +++++++
39 files changed, 727 insertions(+), 19 deletions(-)
create mode 100644 backend/internal/handler/admin/usage_handler_search_users_test.go
create mode 100644 backend/internal/handler/admin/user_handler_get_deleted_test.go
create mode 100644 backend/internal/handler/dto/mappers_deleted_user_test.go
create mode 100644 backend/internal/repository/usage_log_repo_deleted_user_integration_test.go
create mode 100644 backend/internal/repository/user_repo_include_deleted_integration_test.go
create mode 100644 backend/internal/service/admin_service_get_deleted_test.go
create mode 100644 frontend/src/components/admin/usage/__tests__/UsageFilters.spec.ts
diff --git a/backend/internal/handler/admin/admin_service_stub_test.go b/backend/internal/handler/admin/admin_service_stub_test.go
index fd0ec459..819f0cdc 100644
--- a/backend/internal/handler/admin/admin_service_stub_test.go
+++ b/backend/internal/handler/admin/admin_service_stub_test.go
@@ -160,6 +160,10 @@ func (s *stubAdminService) GetUser(ctx context.Context, id int64) (*service.User
return &user, nil
}
+func (s *stubAdminService) GetUserIncludeDeleted(ctx context.Context, id int64) (*service.User, error) {
+ return s.GetUser(ctx, id)
+}
+
func (s *stubAdminService) CreateUser(ctx context.Context, input *service.CreateUserInput) (*service.User, error) {
user := service.User{ID: 100, Email: input.Email, Status: service.StatusActive}
return &user, nil
diff --git a/backend/internal/handler/admin/usage_handler.go b/backend/internal/handler/admin/usage_handler.go
index 0857a138..2ded3c1d 100644
--- a/backend/internal/handler/admin/usage_handler.go
+++ b/backend/internal/handler/admin/usage_handler.go
@@ -344,23 +344,25 @@ func (h *UsageHandler) SearchUsers(c *gin.Context) {
}
// Limit to 30 results
- users, _, err := h.adminService.ListUsers(c.Request.Context(), 1, 30, service.UserListFilters{Search: keyword}, "email", "asc")
+ users, _, err := h.adminService.ListUsers(c.Request.Context(), 1, 30, service.UserListFilters{Search: keyword, IncludeDeleted: true}, "email", "asc")
if err != nil {
response.ErrorFrom(c, err)
return
}
- // Return simplified user list (only id and email)
+ // Return simplified user list (only id, email and deleted flag)
type SimpleUser struct {
- ID int64 `json:"id"`
- Email string `json:"email"`
+ ID int64 `json:"id"`
+ Email string `json:"email"`
+ Deleted bool `json:"deleted"`
}
result := make([]SimpleUser, len(users))
for i, u := range users {
result[i] = SimpleUser{
- ID: u.ID,
- Email: u.Email,
+ ID: u.ID,
+ Email: u.Email,
+ Deleted: u.DeletedAt != nil,
}
}
diff --git a/backend/internal/handler/admin/usage_handler_search_users_test.go b/backend/internal/handler/admin/usage_handler_search_users_test.go
new file mode 100644
index 00000000..ca435012
--- /dev/null
+++ b/backend/internal/handler/admin/usage_handler_search_users_test.go
@@ -0,0 +1,56 @@
+package admin
+
+import (
+ "context"
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/service"
+ "github.com/gin-gonic/gin"
+ "github.com/stretchr/testify/require"
+)
+
+// 捕获 ListUsers 入参、返回一个已删用户的 admin service 桩。
+type searchUsersAdminStub struct {
+ service.AdminService
+ gotFilters service.UserListFilters
+}
+
+func (s *searchUsersAdminStub) ListUsers(ctx context.Context, page, pageSize int, filters service.UserListFilters, sortBy, sortOrder string) ([]service.User, int64, error) {
+ s.gotFilters = filters
+ ts := time.Date(2026, 5, 28, 0, 0, 0, 0, time.UTC)
+ return []service.User{
+ {ID: 1, Email: "active@test.com"},
+ {ID: 2, Email: "deleted@test.com", DeletedAt: &ts},
+ }, 2, nil
+}
+
+func TestAdminUsageSearchUsers_IncludesDeletedAndFlags(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ stub := &searchUsersAdminStub{}
+ handler := NewUsageHandler(nil, nil, stub, nil)
+ router := gin.New()
+ router.GET("/admin/usage/search-users", handler.SearchUsers)
+
+ req := httptest.NewRequest(http.MethodGet, "/admin/usage/search-users?q=test", nil)
+ rec := httptest.NewRecorder()
+ router.ServeHTTP(rec, req)
+
+ require.Equal(t, http.StatusOK, rec.Code)
+ require.True(t, stub.gotFilters.IncludeDeleted, "SearchUsers 必须请求 IncludeDeleted")
+
+ var resp struct {
+ Data []struct {
+ ID int64 `json:"id"`
+ Email string `json:"email"`
+ Deleted bool `json:"deleted"`
+ } `json:"data"`
+ }
+ require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp))
+ require.Len(t, resp.Data, 2)
+ require.False(t, resp.Data[0].Deleted)
+ require.True(t, resp.Data[1].Deleted, "已删用户必须标记 deleted=true")
+}
diff --git a/backend/internal/handler/admin/user_handler.go b/backend/internal/handler/admin/user_handler.go
index 6c0a02ff..ada82e90 100644
--- a/backend/internal/handler/admin/user_handler.go
+++ b/backend/internal/handler/admin/user_handler.go
@@ -195,7 +195,12 @@ func (h *UserHandler) GetByID(c *gin.Context) {
return
}
- user, err := h.adminService.GetUser(c.Request.Context(), userID)
+ var user *service.User
+ if c.Query("include_deleted") == "true" {
+ user, err = h.adminService.GetUserIncludeDeleted(c.Request.Context(), userID)
+ } else {
+ user, err = h.adminService.GetUser(c.Request.Context(), userID)
+ }
if err != nil {
response.ErrorFrom(c, err)
return
diff --git a/backend/internal/handler/admin/user_handler_get_deleted_test.go b/backend/internal/handler/admin/user_handler_get_deleted_test.go
new file mode 100644
index 00000000..1b3070cd
--- /dev/null
+++ b/backend/internal/handler/admin/user_handler_get_deleted_test.go
@@ -0,0 +1,51 @@
+package admin
+
+import (
+ "context"
+ "net/http"
+ "net/http/httptest"
+ "testing"
+
+ "github.com/Wei-Shaw/sub2api/internal/service"
+ "github.com/gin-gonic/gin"
+ "github.com/stretchr/testify/require"
+)
+
+type getByIDAdminStub struct {
+ service.AdminService
+}
+
+func (s *getByIDAdminStub) GetUser(_ context.Context, _ int64) (*service.User, error) {
+ return nil, service.ErrUserNotFound
+}
+
+func (s *getByIDAdminStub) GetUserIncludeDeleted(_ context.Context, id int64) (*service.User, error) {
+ return &service.User{ID: id, Email: "del@test.com"}, nil
+}
+
+func setupGetByIDRouter(svc service.AdminService) *gin.Engine {
+ gin.SetMode(gin.TestMode)
+ r := gin.New()
+ h := NewUserHandler(svc, nil, nil, nil)
+ r.GET("/admin/users/:id", h.GetByID)
+ return r
+}
+
+func TestAdminUserGetByID_IncludeDeleted(t *testing.T) {
+ svc := &getByIDAdminStub{AdminService: newStubAdminService()}
+ router := setupGetByIDRouter(svc)
+
+ t.Run("normal path returns 404 for deleted user", func(t *testing.T) {
+ w := httptest.NewRecorder()
+ req, _ := http.NewRequest(http.MethodGet, "/admin/users/7", nil)
+ router.ServeHTTP(w, req)
+ require.Equal(t, http.StatusNotFound, w.Code)
+ })
+
+ t.Run("include_deleted=true returns 200", func(t *testing.T) {
+ w := httptest.NewRecorder()
+ req, _ := http.NewRequest(http.MethodGet, "/admin/users/7?include_deleted=true", nil)
+ router.ServeHTTP(w, req)
+ require.Equal(t, http.StatusOK, w.Code)
+ })
+}
diff --git a/backend/internal/handler/auth_oauth_pending_flow_test.go b/backend/internal/handler/auth_oauth_pending_flow_test.go
index 70fb160a..2f8f4e58 100644
--- a/backend/internal/handler/auth_oauth_pending_flow_test.go
+++ b/backend/internal/handler/auth_oauth_pending_flow_test.go
@@ -2914,6 +2914,10 @@ func (r *oauthPendingFlowUserRepo) DisableTotp(ctx context.Context, userID int64
Exec(ctx)
}
+func (r *oauthPendingFlowUserRepo) GetByIDIncludeDeleted(ctx context.Context, id int64) (*service.User, error) {
+ return r.GetByID(ctx, id)
+}
+
func oauthPendingFlowServiceUser(entity *dbent.User) *service.User {
if entity == nil {
return nil
diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go
index 51a11ea7..86f98f15 100644
--- a/backend/internal/handler/dto/mappers.go
+++ b/backend/internal/handler/dto/mappers.go
@@ -30,6 +30,7 @@ func UserFromServiceShallow(u *service.User) *User {
BalanceNotifyExtraEmails: NotifyEmailEntriesFromService(u.BalanceNotifyExtraEmails),
TotalRecharged: u.TotalRecharged,
RPMLimit: u.RPMLimit,
+ DeletedAt: u.DeletedAt,
}
}
diff --git a/backend/internal/handler/dto/mappers_deleted_user_test.go b/backend/internal/handler/dto/mappers_deleted_user_test.go
new file mode 100644
index 00000000..8ce5388e
--- /dev/null
+++ b/backend/internal/handler/dto/mappers_deleted_user_test.go
@@ -0,0 +1,20 @@
+package dto
+
+import (
+ "testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/service"
+ "github.com/stretchr/testify/require"
+)
+
+func TestUserFromServiceShallow_MapsDeletedAt(t *testing.T) {
+ ts := time.Date(2026, 5, 28, 10, 0, 0, 0, time.UTC)
+
+ deleted := UserFromServiceShallow(&service.User{ID: 1, Email: "d@test.com", DeletedAt: &ts})
+ require.NotNil(t, deleted.DeletedAt)
+ require.Equal(t, ts, *deleted.DeletedAt)
+
+ active := UserFromServiceShallow(&service.User{ID: 2, Email: "a@test.com"})
+ require.Nil(t, active.DeletedAt, "active user must have nil DeletedAt")
+}
diff --git a/backend/internal/handler/dto/types.go b/backend/internal/handler/dto/types.go
index b1841c62..08dc6572 100644
--- a/backend/internal/handler/dto/types.go
+++ b/backend/internal/handler/dto/types.go
@@ -20,6 +20,7 @@ type User struct {
LastActiveAt *time.Time `json:"last_active_at,omitempty"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
+ DeletedAt *time.Time `json:"deleted_at,omitempty"`
// 余额不足通知
BalanceNotifyEnabled bool `json:"balance_notify_enabled"`
diff --git a/backend/internal/handler/user_handler_test.go b/backend/internal/handler/user_handler_test.go
index 41647802..2e366c23 100644
--- a/backend/internal/handler/user_handler_test.go
+++ b/backend/internal/handler/user_handler_test.go
@@ -118,6 +118,9 @@ func (s *userHandlerRepoStub) RemoveGroupFromUserAllowedGroups(context.Context,
func (s *userHandlerRepoStub) UpdateTotpSecret(context.Context, int64, *string) error { return nil }
func (s *userHandlerRepoStub) EnableTotp(context.Context, int64) error { return nil }
func (s *userHandlerRepoStub) DisableTotp(context.Context, int64) error { return nil }
+func (s *userHandlerRepoStub) GetByIDIncludeDeleted(ctx context.Context, id int64) (*service.User, error) {
+ return s.GetByID(ctx, id)
+}
func (s *userHandlerRepoStub) ListUserAuthIdentities(context.Context, int64) ([]service.UserAuthIdentityRecord, error) {
out := make([]service.UserAuthIdentityRecord, len(s.identities))
copy(out, s.identities)
diff --git a/backend/internal/repository/api_key_repo.go b/backend/internal/repository/api_key_repo.go
index bfe09283..7db35ecc 100644
--- a/backend/internal/repository/api_key_repo.go
+++ b/backend/internal/repository/api_key_repo.go
@@ -679,6 +679,7 @@ func userEntityToService(u *dbent.User) *service.User {
RPMLimit: u.RpmLimit,
CreatedAt: u.CreatedAt,
UpdatedAt: u.UpdatedAt,
+ DeletedAt: u.DeletedAt,
}
// Parse extra emails JSON (supports both old []string and new []NotifyEmailEntry format)
if u.BalanceNotifyExtraEmails != "" && u.BalanceNotifyExtraEmails != "[]" {
diff --git a/backend/internal/repository/usage_log_repo.go b/backend/internal/repository/usage_log_repo.go
index f11910a0..1835ebe5 100644
--- a/backend/internal/repository/usage_log_repo.go
+++ b/backend/internal/repository/usage_log_repo.go
@@ -17,6 +17,7 @@ import (
dbaccount "github.com/Wei-Shaw/sub2api/ent/account"
dbapikey "github.com/Wei-Shaw/sub2api/ent/apikey"
dbgroup "github.com/Wei-Shaw/sub2api/ent/group"
+ "github.com/Wei-Shaw/sub2api/ent/schema/mixins"
dbuser "github.com/Wei-Shaw/sub2api/ent/user"
dbusersub "github.com/Wei-Shaw/sub2api/ent/usersubscription"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
@@ -4121,7 +4122,8 @@ func (r *usageLogRepository) loadUsers(ctx context.Context, ids []int64) (map[in
if len(ids) == 0 {
return out, nil
}
- models, err := r.client.User.Query().Where(dbuser.IDIn(ids...)).All(ctx)
+ // 无条件穿透软删除:ids 来自调用方已按 user_id 筛选的日志行;普通用户路径强制 UserID=本人(本人必为活跃用户),不会借此解析他人已删身份;仅 admin 路径可借此显示已删用户。
+ models, err := r.client.User.Query().Where(dbuser.IDIn(ids...)).All(mixins.SkipSoftDelete(ctx))
if err != nil {
return nil, err
}
diff --git a/backend/internal/repository/usage_log_repo_deleted_user_integration_test.go b/backend/internal/repository/usage_log_repo_deleted_user_integration_test.go
new file mode 100644
index 00000000..70835b03
--- /dev/null
+++ b/backend/internal/repository/usage_log_repo_deleted_user_integration_test.go
@@ -0,0 +1,65 @@
+//go:build integration
+
+package repository
+
+import (
+ "context"
+ "testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/pkg/pagination"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats"
+ "github.com/Wei-Shaw/sub2api/internal/service"
+ "github.com/stretchr/testify/require"
+)
+
+func TestUsageLog_ListWithFilters_ResolvesSoftDeletedUser(t *testing.T) {
+ ctx := context.Background()
+ tx := testEntTx(t)
+ client := tx.Client()
+ repo := newUsageLogRepositoryWithSQL(client, tx)
+
+ // 一个活跃用户、一个将被软删的用户,各一条日志。
+ active := mustCreateUser(t, client, &service.User{Email: "active-listfilter@test.com"})
+ deleted := mustCreateUser(t, client, &service.User{Email: "deleted-listfilter@test.com"})
+ apiKey := mustCreateApiKey(t, client, &service.APIKey{UserID: deleted.ID, Key: "sk-del-1", Name: "k"})
+ apiKey2 := mustCreateApiKey(t, client, &service.APIKey{UserID: active.ID, Key: "sk-act-1", Name: "k"})
+ account := mustCreateAccount(t, client, &service.Account{Name: "acc-listfilter"})
+
+ now := time.Now().UTC()
+ for _, u := range []struct {
+ uid int64
+ kid int64
+ }{{deleted.ID, apiKey.ID}, {active.ID, apiKey2.ID}} {
+ _, err := repo.Create(ctx, &service.UsageLog{
+ UserID: u.uid, APIKeyID: u.kid, AccountID: account.ID,
+ Model: "claude-3", InputTokens: 1, OutputTokens: 1,
+ TotalCost: 0.1, ActualCost: 0.1, CreatedAt: now,
+ })
+ require.NoError(t, err)
+ }
+
+ // 软删除该用户(触发 SoftDeleteMixin Hook → UPDATE deleted_at)。
+ require.NoError(t, client.User.DeleteOneID(deleted.ID).Exec(ctx))
+
+ logs, _, err := repo.ListWithFilters(ctx, pagination.PaginationParams{Page: 1, PageSize: 50},
+ usagestats.UsageLogFilters{ExactTotal: true})
+ require.NoError(t, err)
+
+ byUser := map[int64]service.UsageLog{}
+ for _, l := range logs {
+ byUser[l.UserID] = l
+ }
+
+ // 已删用户的日志行:富化后 User 非 nil、邮箱正确、DeletedAt 非 nil。
+ delLog, ok := byUser[deleted.ID]
+ require.True(t, ok, "deleted user's usage log must still be listed")
+ require.NotNil(t, delLog.User, "deleted user identity must resolve")
+ require.Equal(t, "deleted-listfilter@test.com", delLog.User.Email)
+ require.NotNil(t, delLog.User.DeletedAt, "DeletedAt must be set for soft-deleted user")
+
+ // 活跃用户:DeletedAt 为 nil。
+ actLog := byUser[active.ID]
+ require.NotNil(t, actLog.User)
+ require.Nil(t, actLog.User.DeletedAt)
+}
diff --git a/backend/internal/repository/user_repo.go b/backend/internal/repository/user_repo.go
index 610d9a7b..fb05452d 100644
--- a/backend/internal/repository/user_repo.go
+++ b/backend/internal/repository/user_repo.go
@@ -16,6 +16,7 @@ import (
dbgroup "github.com/Wei-Shaw/sub2api/ent/group"
"github.com/Wei-Shaw/sub2api/ent/identityadoptiondecision"
"github.com/Wei-Shaw/sub2api/ent/predicate"
+ "github.com/Wei-Shaw/sub2api/ent/schema/mixins"
dbuser "github.com/Wei-Shaw/sub2api/ent/user"
"github.com/Wei-Shaw/sub2api/ent/userallowedgroup"
"github.com/Wei-Shaw/sub2api/ent/usersubscription"
@@ -133,6 +134,23 @@ func (r *userRepository) GetByID(ctx context.Context, id int64) (*service.User,
return out, nil
}
+func (r *userRepository) GetByIDIncludeDeleted(ctx context.Context, id int64) (*service.User, error) {
+ ctx = mixins.SkipSoftDelete(ctx)
+ m, err := r.client.User.Query().Where(dbuser.IDEQ(id)).Only(ctx)
+ if err != nil {
+ return nil, translatePersistenceError(err, service.ErrUserNotFound, nil)
+ }
+ out := userEntityToService(m)
+ groups, err := r.loadAllowedGroups(ctx, []int64{id})
+ if err != nil {
+ return nil, err
+ }
+ if v, ok := groups[id]; ok {
+ out.AllowedGroups = v
+ }
+ return out, nil
+}
+
func (r *userRepository) GetByEmail(ctx context.Context, email string) (*service.User, error) {
matches, err := r.client.User.Query().
Where(userEmailLookupPredicate(email)).
@@ -405,6 +423,12 @@ func (r *userRepository) List(ctx context.Context, params pagination.PaginationP
}
func (r *userRepository) ListWithFilters(ctx context.Context, params pagination.PaginationParams, filters service.UserListFilters) ([]service.User, *pagination.PaginationResult, error) {
+ // SkipSoftDelete 仅作用于 User 身份解析(下方 Count/All);订阅、分组等关联实体沿用原始 ctx,避免穿透到这些同样带软删除的实体而带出已删除行。
+ userCtx := ctx
+ if filters.IncludeDeleted {
+ userCtx = mixins.SkipSoftDelete(ctx)
+ }
+
q := r.client.User.Query()
if filters.Status != "" {
@@ -445,7 +469,7 @@ func (r *userRepository) ListWithFilters(ctx context.Context, params pagination.
q = q.Where(dbuser.IDIn(allowedUserIDs...))
}
- total, err := q.Clone().Count(ctx)
+ total, err := q.Clone().Count(userCtx)
if err != nil {
return nil, nil, err
}
@@ -457,7 +481,7 @@ func (r *userRepository) ListWithFilters(ctx context.Context, params pagination.
usersQuery = usersQuery.Order(order)
}
- users, err := usersQuery.All(ctx)
+ users, err := usersQuery.All(userCtx)
if err != nil {
return nil, nil, err
}
diff --git a/backend/internal/repository/user_repo_include_deleted_integration_test.go b/backend/internal/repository/user_repo_include_deleted_integration_test.go
new file mode 100644
index 00000000..014b24f9
--- /dev/null
+++ b/backend/internal/repository/user_repo_include_deleted_integration_test.go
@@ -0,0 +1,69 @@
+//go:build integration
+
+package repository
+
+import (
+ "context"
+ "testing"
+
+ "github.com/Wei-Shaw/sub2api/internal/pkg/pagination"
+ "github.com/Wei-Shaw/sub2api/internal/service"
+ "github.com/stretchr/testify/require"
+)
+
+func TestUserRepo_ListWithFilters_IncludeDeleted(t *testing.T) {
+ ctx := context.Background()
+ tx := testEntTx(t)
+ client := tx.Client()
+ repo := NewUserRepository(client, integrationDB)
+
+ active := mustCreateUser(t, client, &service.User{Email: "shared-keyword-active@test.com"})
+ deleted := mustCreateUser(t, client, &service.User{Email: "shared-keyword-deleted@test.com"})
+ require.NoError(t, client.User.DeleteOneID(deleted.ID).Exec(ctx))
+
+ params := pagination.PaginationParams{Page: 1, PageSize: 50, SortBy: "email", SortOrder: "asc"}
+
+ // 默认(不含已删):只返回活跃用户。
+ usersDefault, resDefault, err := repo.ListWithFilters(ctx, params,
+ service.UserListFilters{Search: "shared-keyword-"})
+ require.NoError(t, err)
+ require.Len(t, usersDefault, 1)
+ require.Equal(t, active.ID, usersDefault[0].ID)
+ require.EqualValues(t, 1, resDefault.Total)
+
+ // IncludeDeleted=true:两个都返回,且 Total 与结果集一致。
+ usersAll, resAll, err := repo.ListWithFilters(ctx, params,
+ service.UserListFilters{Search: "shared-keyword-", IncludeDeleted: true})
+ require.NoError(t, err)
+ require.Len(t, usersAll, 2)
+ require.EqualValues(t, 2, resAll.Total, "Count 必须与结果集行数一致")
+
+ var delUser *service.User
+ for i := range usersAll {
+ if usersAll[i].ID == deleted.ID {
+ delUser = &usersAll[i]
+ }
+ }
+ require.NotNil(t, delUser)
+ require.NotNil(t, delUser.DeletedAt)
+}
+
+func TestUserRepo_GetByIDIncludeDeleted(t *testing.T) {
+ ctx := context.Background()
+ tx := testEntTx(t)
+ client := tx.Client()
+ repo := NewUserRepository(client, integrationDB)
+
+ u := mustCreateUser(t, client, &service.User{Email: "getbyid-deleted@test.com"})
+ require.NoError(t, client.User.DeleteOneID(u.ID).Exec(ctx))
+
+ // 默认 GetByID:找不到(被软删过滤)。
+ _, err := repo.GetByID(ctx, u.ID)
+ require.ErrorIs(t, err, service.ErrUserNotFound)
+
+ // GetByIDIncludeDeleted:找得到,且 DeletedAt 非空。
+ got, err := repo.GetByIDIncludeDeleted(ctx, u.ID)
+ require.NoError(t, err)
+ require.Equal(t, "getbyid-deleted@test.com", got.Email)
+ require.NotNil(t, got.DeletedAt)
+}
diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go
index 9eea0924..6bb87995 100644
--- a/backend/internal/server/api_contract_test.go
+++ b/backend/internal/server/api_contract_test.go
@@ -1492,6 +1492,10 @@ func (r *stubUserRepo) DisableTotp(ctx context.Context, userID int64) error {
return errors.New("not implemented")
}
+func (r *stubUserRepo) GetByIDIncludeDeleted(ctx context.Context, id int64) (*service.User, error) {
+ panic("unexpected GetByIDIncludeDeleted call")
+}
+
type stubApiKeyCache struct{}
func (stubApiKeyCache) GetCreateAttemptCount(ctx context.Context, userID int64) (int, error) {
diff --git a/backend/internal/server/middleware/admin_auth_test.go b/backend/internal/server/middleware/admin_auth_test.go
index 303d0db8..3110c6c1 100644
--- a/backend/internal/server/middleware/admin_auth_test.go
+++ b/backend/internal/server/middleware/admin_auth_test.go
@@ -236,3 +236,7 @@ func (s *stubUserRepo) EnableTotp(ctx context.Context, userID int64) error {
func (s *stubUserRepo) DisableTotp(ctx context.Context, userID int64) error {
panic("unexpected DisableTotp call")
}
+
+func (s *stubUserRepo) GetByIDIncludeDeleted(ctx context.Context, id int64) (*service.User, error) {
+ panic("unexpected GetByIDIncludeDeleted call")
+}
diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go
index d46b636f..81d1f022 100644
--- a/backend/internal/service/admin_service.go
+++ b/backend/internal/service/admin_service.go
@@ -33,6 +33,7 @@ type AdminService interface {
// User management
ListUsers(ctx context.Context, page, pageSize int, filters UserListFilters, sortBy, sortOrder string) ([]User, int64, error)
GetUser(ctx context.Context, id int64) (*User, error)
+ GetUserIncludeDeleted(ctx context.Context, id int64) (*User, error)
CreateUser(ctx context.Context, input *CreateUserInput) (*User, error)
UpdateUser(ctx context.Context, id int64, input *UpdateUserInput) (*User, error)
DeleteUser(ctx context.Context, id int64) error
@@ -674,6 +675,10 @@ func (s *adminServiceImpl) GetUser(ctx context.Context, id int64) (*User, error)
return user, nil
}
+func (s *adminServiceImpl) GetUserIncludeDeleted(ctx context.Context, id int64) (*User, error) {
+ return s.userRepo.GetByIDIncludeDeleted(ctx, id)
+}
+
func (s *adminServiceImpl) CreateUser(ctx context.Context, input *CreateUserInput) (*User, error) {
user := &User{
Email: input.Email,
diff --git a/backend/internal/service/admin_service_apikey_test.go b/backend/internal/service/admin_service_apikey_test.go
index 3b3dbc21..f26fadb8 100644
--- a/backend/internal/service/admin_service_apikey_test.go
+++ b/backend/internal/service/admin_service_apikey_test.go
@@ -69,8 +69,12 @@ func (s *userRepoStubForGroupUpdate) UpdateConcurrency(context.Context, int64, i
panic("unexpected")
}
-func (s *userRepoStubForGroupUpdate) BatchSetConcurrency(context.Context, []int64, int) (int, error) { return 0, nil }
-func (s *userRepoStubForGroupUpdate) BatchAddConcurrency(context.Context, []int64, int) (int, error) { return 0, nil }
+func (s *userRepoStubForGroupUpdate) BatchSetConcurrency(context.Context, []int64, int) (int, error) {
+ return 0, nil
+}
+func (s *userRepoStubForGroupUpdate) BatchAddConcurrency(context.Context, []int64, int) (int, error) {
+ return 0, nil
+}
func (s *userRepoStubForGroupUpdate) ExistsByEmail(context.Context, string) (bool, error) {
panic("unexpected")
}
@@ -82,6 +86,9 @@ func (s *userRepoStubForGroupUpdate) UpdateTotpSecret(context.Context, int64, *s
}
func (s *userRepoStubForGroupUpdate) EnableTotp(context.Context, int64) error { panic("unexpected") }
func (s *userRepoStubForGroupUpdate) DisableTotp(context.Context, int64) error { panic("unexpected") }
+func (s *userRepoStubForGroupUpdate) GetByIDIncludeDeleted(ctx context.Context, id int64) (*User, error) {
+ panic("unexpected GetByIDIncludeDeleted call")
+}
func (s *userRepoStubForGroupUpdate) ListUserAuthIdentities(context.Context, int64) ([]UserAuthIdentityRecord, error) {
panic("unexpected")
}
diff --git a/backend/internal/service/admin_service_delete_test.go b/backend/internal/service/admin_service_delete_test.go
index 2aae73a9..150c4f53 100644
--- a/backend/internal/service/admin_service_delete_test.go
+++ b/backend/internal/service/admin_service_delete_test.go
@@ -173,6 +173,10 @@ func (s *userRepoStub) DisableTotp(ctx context.Context, userID int64) error {
panic("unexpected DisableTotp call")
}
+func (s *userRepoStub) GetByIDIncludeDeleted(ctx context.Context, id int64) (*User, error) {
+ return s.GetByID(ctx, id)
+}
+
type groupRepoStub struct {
affectedUserIDs []int64
deleteErr error
diff --git a/backend/internal/service/admin_service_email_identity_sync_test.go b/backend/internal/service/admin_service_email_identity_sync_test.go
index c791b747..c3737f5a 100644
--- a/backend/internal/service/admin_service_email_identity_sync_test.go
+++ b/backend/internal/service/admin_service_email_identity_sync_test.go
@@ -113,8 +113,12 @@ func (s *emailSyncRepoStub) RemoveGroupFromAllowedGroups(context.Context, int64)
return 0, nil
}
-func (s *emailSyncRepoStub) BatchSetConcurrency(context.Context, []int64, int) (int, error) { return 0, nil }
-func (s *emailSyncRepoStub) BatchAddConcurrency(context.Context, []int64, int) (int, error) { return 0, nil }
+func (s *emailSyncRepoStub) BatchSetConcurrency(context.Context, []int64, int) (int, error) {
+ return 0, nil
+}
+func (s *emailSyncRepoStub) BatchAddConcurrency(context.Context, []int64, int) (int, error) {
+ return 0, nil
+}
func (s *emailSyncRepoStub) AddGroupToAllowedGroups(context.Context, int64, int64) error { return nil }
@@ -133,6 +137,9 @@ func (s *emailSyncRepoStub) UpdateTotpSecret(context.Context, int64, *string) er
func (s *emailSyncRepoStub) EnableTotp(context.Context, int64) error { return nil }
func (s *emailSyncRepoStub) DisableTotp(context.Context, int64) error { return nil }
+func (s *emailSyncRepoStub) GetByIDIncludeDeleted(ctx context.Context, id int64) (*User, error) {
+ return s.GetByID(ctx, id)
+}
func (s *emailSyncRepoStub) EnsureEmailAuthIdentity(_ context.Context, userID int64, email string) error {
s.ensureCalls = append(s.ensureCalls, ensureEmailCall{userID: userID, email: email})
diff --git a/backend/internal/service/admin_service_get_deleted_test.go b/backend/internal/service/admin_service_get_deleted_test.go
new file mode 100644
index 00000000..6ad17f60
--- /dev/null
+++ b/backend/internal/service/admin_service_get_deleted_test.go
@@ -0,0 +1,22 @@
+//go:build unit
+
+package service
+
+import (
+ "context"
+ "testing"
+ "time"
+
+ "github.com/stretchr/testify/require"
+)
+
+func TestAdminService_GetUserIncludeDeleted(t *testing.T) {
+ ts := time.Date(2026, 5, 28, 0, 0, 0, 0, time.UTC)
+ repo := &userRepoStub{user: &User{ID: 7, Email: "del@test.com", DeletedAt: &ts}}
+ svc := &adminServiceImpl{userRepo: repo}
+
+ got, err := svc.GetUserIncludeDeleted(context.Background(), 7)
+ require.NoError(t, err)
+ require.Equal(t, int64(7), got.ID)
+ require.NotNil(t, got.DeletedAt)
+}
diff --git a/backend/internal/service/auth_service_email_bind_test.go b/backend/internal/service/auth_service_email_bind_test.go
index 87867395..28bb0a3b 100644
--- a/backend/internal/service/auth_service_email_bind_test.go
+++ b/backend/internal/service/auth_service_email_bind_test.go
@@ -850,6 +850,9 @@ func (s *emailBindUserRepoStub) UnbindUserAuthProvider(context.Context, int64, s
func (s *emailBindUserRepoStub) UpdateTotpSecret(context.Context, int64, *string) error { return nil }
func (s *emailBindUserRepoStub) EnableTotp(context.Context, int64) error { return nil }
func (s *emailBindUserRepoStub) DisableTotp(context.Context, int64) error { return nil }
+func (s *emailBindUserRepoStub) GetByIDIncludeDeleted(ctx context.Context, id int64) (*service.User, error) {
+ return s.GetByID(ctx, id)
+}
func cloneEmailBindUser(user *service.User) *service.User {
if user == nil {
diff --git a/backend/internal/service/content_moderation_test.go b/backend/internal/service/content_moderation_test.go
index 1fb72f36..6c6fef44 100644
--- a/backend/internal/service/content_moderation_test.go
+++ b/backend/internal/service/content_moderation_test.go
@@ -277,6 +277,10 @@ func (r *contentModerationTestUserRepo) DisableTotp(ctx context.Context, userID
panic("unexpected DisableTotp call")
}
+func (r *contentModerationTestUserRepo) GetByIDIncludeDeleted(ctx context.Context, id int64) (*User, error) {
+ return r.GetByID(ctx, id)
+}
+
type contentModerationTestAuthCacheInvalidator struct {
userIDs []int64
}
diff --git a/backend/internal/service/user.go b/backend/internal/service/user.go
index f9833611..edb944ee 100644
--- a/backend/internal/service/user.go
+++ b/backend/internal/service/user.go
@@ -32,6 +32,7 @@ type User struct {
LastUsedAt *time.Time
CreatedAt time.Time
UpdatedAt time.Time
+ DeletedAt *time.Time // 非 nil 表示用户已软删除
// GroupRates 用户专属分组倍率配置
// map[groupID]rateMultiplier
diff --git a/backend/internal/service/user_service.go b/backend/internal/service/user_service.go
index 36bcf1c8..f801e2c4 100644
--- a/backend/internal/service/user_service.go
+++ b/backend/internal/service/user_service.go
@@ -74,11 +74,16 @@ type UserListFilters struct {
// For large datasets this can be expensive; admin list pages should enable it on demand.
// nil means not specified (default: load subscriptions for backward compatibility).
IncludeSubscriptions *bool
+ // IncludeDeleted 为 true 时绕过软删除过滤,返回含已删除(deleted_at 非空)的用户。
+ // 仅供 /admin/usage 的 SearchUsers 端点使用,其他列表调用方不要设置。
+ IncludeDeleted bool
}
type UserRepository interface {
Create(ctx context.Context, user *User) error
GetByID(ctx context.Context, id int64) (*User, error)
+ // GetByIDIncludeDeleted 绕过软删除过滤按 ID 取用户(含已删)。仅供管理员审计/usage 点击使用。
+ GetByIDIncludeDeleted(ctx context.Context, id int64) (*User, error)
GetByEmail(ctx context.Context, email string) (*User, error)
GetFirstAdmin(ctx context.Context) (*User, error)
Update(ctx context.Context, user *User) error
diff --git a/backend/internal/service/user_service_test.go b/backend/internal/service/user_service_test.go
index 1a18e70a..417140ad 100644
--- a/backend/internal/service/user_service_test.go
+++ b/backend/internal/service/user_service_test.go
@@ -236,6 +236,10 @@ func (m *mockUserRepo) UnbindUserAuthProvider(_ context.Context, _ int64, provid
return nil
}
+func (m *mockUserRepo) GetByIDIncludeDeleted(ctx context.Context, id int64) (*User, error) {
+ return m.GetByID(ctx, id)
+}
+
func (m *mockUserRepo) WithUserProfileIdentityTx(ctx context.Context, fn func(txCtx context.Context) error) error {
m.txCalls++
txState := &mockUserRepoTxState{
diff --git a/frontend/src/api/admin/usage.ts b/frontend/src/api/admin/usage.ts
index 7ad00742..c3a38c0a 100644
--- a/frontend/src/api/admin/usage.ts
+++ b/frontend/src/api/admin/usage.ts
@@ -27,6 +27,7 @@ export interface AdminUsageStatsResponse {
export interface SimpleUser {
id: number
email: string
+ deleted: boolean
}
export interface SimpleApiKey {
diff --git a/frontend/src/api/admin/users.ts b/frontend/src/api/admin/users.ts
index bfe5e3ba..b84eb1e3 100644
--- a/frontend/src/api/admin/users.ts
+++ b/frontend/src/api/admin/users.ts
@@ -100,10 +100,12 @@ export async function list(
/**
* Get user by ID
* @param id - User ID
+ * @param includeDeleted - Whether to include soft-deleted users
* @returns User details
*/
-export async function getById(id: number): Promise {
- const { data } = await apiClient.get(`/admin/users/${id}`)
+export async function getById(id: number, includeDeleted = false): Promise {
+ const url = includeDeleted ? `/admin/users/${id}?include_deleted=true` : `/admin/users/${id}`
+ const { data } = await apiClient.get(url)
return data
}
diff --git a/frontend/src/components/admin/usage/UsageFilters.vue b/frontend/src/components/admin/usage/UsageFilters.vue
index 66c2b4fa..40f2c9f8 100644
--- a/frontend/src/components/admin/usage/UsageFilters.vue
+++ b/frontend/src/components/admin/usage/UsageFilters.vue
@@ -35,7 +35,7 @@
@click="selectUser(u)"
class="w-full px-4 py-2 text-left hover:bg-gray-100 dark:hover:bg-gray-700"
>
- {{ u.email }}
+ {{ u.email }}({{ t('admin.usage.userDeletedBadge') }})
#{{ u.id }}
@@ -255,7 +255,8 @@ const debounceUserSearch = () => {
return
}
try {
- userResults.value = await adminAPI.usage.searchUsers(userKeyword.value)
+ const results = await adminAPI.usage.searchUsers(userKeyword.value)
+ userResults.value = results.sort((a, b) => Number(a.deleted) - Number(b.deleted))
} catch {
userResults.value = []
}
diff --git a/frontend/src/components/admin/usage/UsageTable.vue b/frontend/src/components/admin/usage/UsageTable.vue
index 65ac1548..ba17ff35 100644
--- a/frontend/src/components/admin/usage/UsageTable.vue
+++ b/frontend/src/components/admin/usage/UsageTable.vue
@@ -21,6 +21,9 @@
{{ row.user.email }}
-
+
+ {{ t('admin.usage.userDeletedBadge') }}
+
#{{ row.user_id }}
diff --git a/frontend/src/components/admin/usage/__tests__/UsageFilters.spec.ts b/frontend/src/components/admin/usage/__tests__/UsageFilters.spec.ts
new file mode 100644
index 00000000..6f1780d1
--- /dev/null
+++ b/frontend/src/components/admin/usage/__tests__/UsageFilters.spec.ts
@@ -0,0 +1,168 @@
+import { describe, expect, it, vi, beforeEach, afterEach } from 'vitest'
+import { mount, flushPromises } from '@vue/test-utils'
+
+import UsageFilters from '../UsageFilters.vue'
+
+// --- i18n messages (only what UsageFilters needs) ---
+const messages: Record = {
+ 'admin.usage.userDeletedBadge': 'deleted',
+ 'admin.usage.userFilter': 'User',
+ 'admin.usage.searchUserPlaceholder': 'Search user...',
+ 'usage.apiKeyFilter': 'API Key',
+ 'admin.usage.searchApiKeyPlaceholder': 'Search API key...',
+ 'usage.model': 'Model',
+ 'admin.usage.allModels': 'All Models',
+ 'admin.usage.account': 'Account',
+ 'admin.usage.searchAccountPlaceholder': 'Search account...',
+ 'usage.type': 'Type',
+ 'admin.usage.allTypes': 'All Types',
+ 'usage.ws': 'WS',
+ 'usage.stream': 'Stream',
+ 'usage.sync': 'Sync',
+ 'admin.usage.billingType': 'Billing Type',
+ 'admin.usage.allBillingTypes': 'All Billing Types',
+ 'admin.usage.billingTypeBalance': 'Balance',
+ 'admin.usage.billingTypeSubscription': 'Subscription',
+ 'admin.usage.billingMode': 'Billing Mode',
+ 'admin.usage.allBillingModes': 'All Billing Modes',
+ 'admin.usage.billingModeToken': 'Token',
+ 'admin.usage.billingModePerRequest': 'Per Request',
+ 'admin.usage.billingModeImage': 'Image',
+ 'admin.usage.group': 'Group',
+ 'admin.usage.allGroups': 'All Groups',
+ 'common.refresh': 'Refresh',
+ 'common.reset': 'Reset',
+ 'admin.usage.cleanup.button': 'Cleanup',
+ 'usage.exportExcel': 'Export',
+}
+
+// Mock vue-i18n
+vi.mock('vue-i18n', async () => {
+ const actual = await vi.importActual('vue-i18n')
+ return {
+ ...actual,
+ useI18n: () => ({
+ t: (key: string) => messages[key] ?? key,
+ }),
+ }
+})
+
+// Mock the admin API module — we control searchUsers return value per test
+const mockSearchUsers = vi.fn()
+const mockSearchApiKeys = vi.fn().mockResolvedValue([])
+
+vi.mock('@/api/admin', () => ({
+ adminAPI: {
+ usage: {
+ searchUsers: (...args: any[]) => mockSearchUsers(...args),
+ searchApiKeys: (...args: any[]) => mockSearchApiKeys(...args),
+ },
+ groups: {
+ list: vi.fn().mockResolvedValue({ items: [] }),
+ },
+ dashboard: {
+ getModelStats: vi.fn().mockResolvedValue({ models: [] }),
+ },
+ accounts: {
+ list: vi.fn().mockResolvedValue({ items: [] }),
+ },
+ },
+}))
+
+// Default props helper
+const defaultFilters = () => ({
+ user_id: undefined,
+ api_key_id: undefined,
+ account_id: undefined,
+ model: null,
+ request_type: null,
+ billing_type: null,
+ billing_mode: null,
+ group_id: null,
+ start_date: '',
+ end_date: '',
+})
+
+function mountFilters(filters = defaultFilters()) {
+ return mount(UsageFilters, {
+ props: {
+ modelValue: filters,
+ exporting: false,
+ startDate: '2026-05-01',
+ endDate: '2026-05-28',
+ showActions: false,
+ },
+ global: {
+ stubs: {
+ Select: true,
+ Teleport: true,
+ },
+ },
+ })
+}
+
+describe('UsageFilters — user search dropdown', () => {
+ beforeEach(() => {
+ vi.useFakeTimers()
+ mockSearchUsers.mockReset()
+ mockSearchApiKeys.mockResolvedValue([])
+ })
+
+ afterEach(() => {
+ vi.useRealTimers()
+ })
+
+ it('(a) labels deleted users with the i18n badge and (b) sorts active users before deleted ones, (c) selection sets user_id', async () => {
+ // Arrange: mock returns deleted FIRST (proves sorting re-orders to active-first)
+ mockSearchUsers.mockResolvedValue([
+ { id: 2, email: 'gone@test.com', deleted: true },
+ { id: 1, email: 'active@test.com', deleted: false },
+ ])
+
+ const wrapper = mountFilters()
+
+ // Trigger focus (sets showUserDropdown = true) then input (fires debounceUserSearch)
+ const input = wrapper.find('input[type="text"]')
+ await input.trigger('focus')
+ await input.setValue('test')
+ await input.trigger('input')
+
+ // Advance debounce timer (300ms) then flush the resolved promise
+ vi.advanceTimersByTime(300)
+ await flushPromises()
+
+ // --- (b) Sort: active user should appear BEFORE deleted user ---
+ // Check the underlying component state via rendered DOM order
+ const buttons = wrapper.findAll('.usage-filter-dropdown button[type="button"]')
+ const emailTexts = buttons.map((b) => b.text())
+
+ // active@test.com should be listed first
+ const activeIdx = emailTexts.findIndex((t) => t.includes('active@test.com'))
+ const deletedIdx = emailTexts.findIndex((t) => t.includes('gone@test.com'))
+ expect(activeIdx).toBeGreaterThanOrEqual(0)
+ expect(deletedIdx).toBeGreaterThanOrEqual(0)
+ expect(activeIdx).toBeLessThan(deletedIdx)
+
+ // --- (a) Label: deleted user's button shows the badge text ---
+ const deletedButton = buttons[deletedIdx]
+ expect(deletedButton.text()).toContain('deleted')
+
+ // active user's button does NOT show the badge text
+ const activeButton = buttons[activeIdx]
+ expect(activeButton.text()).not.toContain('deleted')
+
+ // --- (c) Selection: clicking active user button sets filters.user_id ---
+ await activeButton.trigger('click')
+ await flushPromises()
+
+ // The component emits 'update:modelValue' or modifies filters.user_id via toRef
+ // selectUser sets filters.value.user_id = u.id and emits 'change'
+ const changeEmits = wrapper.emitted('change')
+ expect(changeEmits).toBeTruthy()
+ expect(changeEmits!.length).toBeGreaterThan(0)
+
+ // Also confirm user_id was set by checking the emitted change came through
+ // (the component uses toRef so modelValue is mutated in place and 'change' is emitted)
+ expect(wrapper.props('modelValue').user_id).toBe(1)
+ })
+})
diff --git a/frontend/src/components/admin/usage/__tests__/UsageTable.spec.ts b/frontend/src/components/admin/usage/__tests__/UsageTable.spec.ts
index ece0dbda..7856d19a 100644
--- a/frontend/src/components/admin/usage/__tests__/UsageTable.spec.ts
+++ b/frontend/src/components/admin/usage/__tests__/UsageTable.spec.ts
@@ -5,6 +5,7 @@ import { nextTick } from 'vue'
import UsageTable from '../UsageTable.vue'
const messages: Record = {
+ 'admin.usage.userDeletedBadge': 'Deleted',
'usage.costDetails': 'Cost Breakdown',
'admin.usage.inputCost': 'Input Cost',
'admin.usage.outputCost': 'Output Cost',
@@ -321,3 +322,91 @@ describe('admin UsageTable tooltip', () => {
expect(text).not.toContain('(2K)')
})
})
+
+// A DataTable stub that also renders cell-user, so the deleted badge can be asserted.
+const DataTableStubWithUser = {
+ props: ['data'],
+ template: `
+
+ `,
+}
+
+describe('admin UsageTable deleted-user badge', () => {
+ it('renders deleted badge for a soft-deleted user row', () => {
+ const row = {
+ request_id: 'req-deleted-user-1',
+ model: 'claude-3',
+ user_id: 2,
+ user: { id: 2, email: 'd@test.com', deleted_at: '2026-05-28T00:00:00Z' },
+ actual_cost: 0,
+ total_cost: 0,
+ input_cost: 0,
+ output_cost: 0,
+ rate_multiplier: 1,
+ input_tokens: 1,
+ output_tokens: 1,
+ }
+
+ const wrapper = mount(UsageTable, {
+ props: {
+ data: [row],
+ loading: false,
+ columns: [{ key: 'user', label: 'User' }],
+ },
+ global: {
+ stubs: {
+ DataTable: DataTableStubWithUser,
+ EmptyState: true,
+ Icon: true,
+ Teleport: true,
+ },
+ },
+ })
+
+ expect(wrapper.text()).toContain('Deleted')
+ expect(wrapper.text()).toContain('d@test.com')
+ })
+
+ it('does NOT render deleted badge for an active user row', () => {
+ const row = {
+ request_id: 'req-active-user-1',
+ model: 'claude-3',
+ user_id: 3,
+ user: { id: 3, email: 'active@test.com', deleted_at: null },
+ actual_cost: 0,
+ total_cost: 0,
+ input_cost: 0,
+ output_cost: 0,
+ rate_multiplier: 1,
+ input_tokens: 1,
+ output_tokens: 1,
+ }
+
+ const wrapper = mount(UsageTable, {
+ props: {
+ data: [row],
+ loading: false,
+ columns: [{ key: 'user', label: 'User' }],
+ },
+ global: {
+ stubs: {
+ DataTable: DataTableStubWithUser,
+ EmptyState: true,
+ Icon: true,
+ Teleport: true,
+ },
+ },
+ })
+
+ expect(wrapper.text()).not.toContain('Deleted')
+ expect(wrapper.text()).toContain('active@test.com')
+ })
+})
diff --git a/frontend/src/components/admin/user/UserBalanceHistoryModal.vue b/frontend/src/components/admin/user/UserBalanceHistoryModal.vue
index 6d48ed77..4bce43ad 100644
--- a/frontend/src/components/admin/user/UserBalanceHistoryModal.vue
+++ b/frontend/src/components/admin/user/UserBalanceHistoryModal.vue
@@ -13,6 +13,9 @@
{{ user.email }}
+
+ {{ t('admin.usage.userDeletedBadge') }}
+
{
const handleUserClick = async (userId: number) => {
try {
- const user = await adminAPI.users.getById(userId)
+ const user = await adminAPI.users.getById(userId, true)
balanceHistoryUser.value = user
showBalanceHistoryModal.value = true
} catch {
diff --git a/frontend/src/views/admin/__tests__/UsageView.spec.ts b/frontend/src/views/admin/__tests__/UsageView.spec.ts
index 1a5d285a..9cdb2a3c 100644
--- a/frontend/src/views/admin/__tests__/UsageView.spec.ts
+++ b/frontend/src/views/admin/__tests__/UsageView.spec.ts
@@ -84,6 +84,10 @@ vi.mock('vue-router', () => ({
const AppLayoutStub = { template: '
' }
const UsageFiltersStub = { template: '
' }
+const UsageTableStub = {
+ emits: ['userClick'],
+ template: '',
+}
const ModelDistributionChartStub = {
props: ['metric'],
emits: ['update:metric'],
@@ -194,3 +198,59 @@ describe('admin UsageView distribution metric toggles', () => {
expect(getSnapshotV2).toHaveBeenCalledTimes(1)
})
})
+
+describe('admin UsageView handleUserClick', () => {
+ beforeEach(() => {
+ vi.useFakeTimers()
+ list.mockReset()
+ getStats.mockReset()
+ getSnapshotV2.mockReset()
+ getById.mockReset()
+
+ list.mockResolvedValue({ items: [], total: 0, pages: 0 })
+ getStats.mockResolvedValue({
+ total_requests: 0, total_input_tokens: 0, total_output_tokens: 0,
+ total_cache_tokens: 0, total_tokens: 0, total_cost: 0, total_actual_cost: 0, average_duration_ms: 0,
+ })
+ getSnapshotV2.mockResolvedValue({ trend: [], models: [], groups: [] })
+ })
+
+ afterEach(() => {
+ vi.useRealTimers()
+ })
+
+ it('opens user via include_deleted when clicking a usage row user', async () => {
+ getById.mockResolvedValue({ id: 2, email: 'd@test.com', deleted_at: '2026-05-28T00:00:00Z' })
+
+ const wrapper = mount(UsageView, {
+ global: {
+ stubs: {
+ AppLayout: AppLayoutStub,
+ UsageStatsCards: true,
+ UsageFilters: UsageFiltersStub,
+ UsageTable: UsageTableStub,
+ UsageExportProgress: true,
+ UsageCleanupDialog: true,
+ UserBalanceHistoryModal: true,
+ AuditLogModal: true,
+ Pagination: true,
+ Select: true,
+ DateRangePicker: true,
+ Icon: true,
+ TokenUsageTrend: true,
+ ModelDistributionChart: true,
+ GroupDistributionChart: true,
+ EndpointDistributionChart: true,
+ },
+ },
+ })
+
+ vi.advanceTimersByTime(120)
+ await flushPromises()
+
+ await wrapper.find('[data-test="usage-table"] .user-click').trigger('click')
+ await flushPromises()
+
+ expect(getById).toHaveBeenCalledWith(2, true)
+ })
+})
From bf24b611399fd43d34a8d1801deac30a7b41fe4e Mon Sep 17 00:00:00 2001
From: DaydreamCoding <22166516+DaydreamCoding@users.noreply.github.com>
Date: Sat, 30 May 2026 18:02:13 +0800
Subject: [PATCH 47/79] =?UTF-8?q?perf(usage):=20=E4=BC=98=E5=8C=96=20/admi?=
=?UTF-8?q?n/usage=20=E6=89=93=E5=BC=80=E9=80=9F=E5=BA=A6=E4=B8=8E?=
=?UTF-8?q?=E5=88=B7=E6=96=B0=E5=93=8D=E5=BA=94?=
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit
根因:页面 mount 并发 6 个请求,其中 5 个在原始 usage_logs 上 live 聚合,
且有 1 个重复 getModelStats。优化(不引入预聚合):
前端
- UsageView 经 :model-options 下传 model 列表,移除 UsageFilters 重复的
getModelStats(mount 请求 6→5,少一次 usage_logs 全表 GROUP BY model)
- 刷新/换筛选保留旧模型数据(invalidateModelStatsCache 只失效标记不清空数据),
图表不再闪空,刷新期间页面保持可交互
后端
- GetStatsWithFilters 4 条聚合 errgroup 并行(仅 *sql.DB 连接池路径,ent.Tx
顺序回退以保事务内不并发),endpoint 明细 best-effort;抑制取消级联噪声日志
- /admin/usage/stats 复用 dashboard 的 newSnapshotCache 30s 处理器层缓存,按
filters+窗口为 key;前端手动刷新带 nocache=1 强制回源(刷新=最新)
注:自合并提交 4c8396c 迁移而来,仅取"列表打开速度与刷新响应"部分;原提交的
"审计查看弹窗大 body 渲染"改动(AuditLogModal / audit-log-format / 审计 i18n)
依赖尚未迁移的审计功能(d8389ade),本次已排除。
Co-Authored-By: Claude Opus 4.8 (1M context)
---
.../internal/handler/admin/usage_handler.go | 22 ++++-
.../handler/admin/usage_query_cache.go | 62 ++++++++++++
.../handler/admin/usage_query_cache_test.go | 28 ++++++
backend/internal/repository/usage_log_repo.go | 98 +++++++++++++------
.../usage_log_repo_stats_integration_test.go | 51 ++++++++++
frontend/src/api/admin/usage.ts | 1 +
.../components/admin/usage/UsageFilters.vue | 27 ++---
.../usage/__tests__/UsageFilters.spec.ts | 45 +++++++--
frontend/src/views/admin/UsageView.vue | 28 ++++--
.../views/admin/__tests__/UsageView.spec.ts | 37 ++++++-
10 files changed, 325 insertions(+), 74 deletions(-)
create mode 100644 backend/internal/handler/admin/usage_query_cache.go
create mode 100644 backend/internal/handler/admin/usage_query_cache_test.go
create mode 100644 backend/internal/repository/usage_log_repo_stats_integration_test.go
diff --git a/backend/internal/handler/admin/usage_handler.go b/backend/internal/handler/admin/usage_handler.go
index 2ded3c1d..11a4aeb8 100644
--- a/backend/internal/handler/admin/usage_handler.go
+++ b/backend/internal/handler/admin/usage_handler.go
@@ -325,10 +325,24 @@ func (h *UsageHandler) Stats(c *gin.Context) {
EndTime: &endTime,
}
- stats, err := h.usageService.GetStatsWithFilters(c.Request.Context(), filters)
- if err != nil {
- response.ErrorFrom(c, err)
- return
+ var stats *usagestats.UsageStats
+ // nocache: 绕过缓存直接回源,刷新者本人拿最新;不回写缓存(管理台"我刷新我自己拿最新"语义,非全局失效)。
+ if parseBoolQueryWithDefault(c.Query("nocache"), false) {
+ s, err := h.usageService.GetStatsWithFilters(c.Request.Context(), filters)
+ if err != nil {
+ response.ErrorFrom(c, err)
+ return
+ }
+ stats = s
+ c.Header("X-Usage-Stats-Cache", "bypass")
+ } else {
+ s, hit, err := h.getStatsCached(c.Request.Context(), filters)
+ if err != nil {
+ response.ErrorFrom(c, err)
+ return
+ }
+ stats = s
+ c.Header("X-Usage-Stats-Cache", cacheStatusValue(hit))
}
response.Success(c, stats)
diff --git a/backend/internal/handler/admin/usage_query_cache.go b/backend/internal/handler/admin/usage_query_cache.go
new file mode 100644
index 00000000..b288a95b
--- /dev/null
+++ b/backend/internal/handler/admin/usage_query_cache.go
@@ -0,0 +1,62 @@
+package admin
+
+import (
+ "context"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats"
+)
+
+// 与 dashboard 查询缓存同款:30s TTL 进程内缓存,仅服务 /admin/usage/stats 读路径。
+var usageStatsCache = newSnapshotCache(30 * time.Second)
+
+type usageStatsCacheKeyData struct {
+ StartTime string `json:"start_time"`
+ EndTime string `json:"end_time"`
+ UserID int64 `json:"user_id"`
+ APIKeyID int64 `json:"api_key_id"`
+ AccountID int64 `json:"account_id"`
+ GroupID int64 `json:"group_id"`
+ Model string `json:"model"`
+ BillingMode string `json:"billing_mode"`
+ RequestType *int16 `json:"request_type"`
+ Stream *bool `json:"stream"`
+ BillingType *int8 `json:"billing_type"`
+}
+
+func usageStatsCacheKey(filters usagestats.UsageLogFilters) string {
+ start := ""
+ if filters.StartTime != nil {
+ start = filters.StartTime.UTC().Format(time.RFC3339)
+ }
+ end := ""
+ if filters.EndTime != nil {
+ end = filters.EndTime.UTC().Format(time.RFC3339)
+ }
+ return mustMarshalDashboardCacheKey(usageStatsCacheKeyData{
+ StartTime: start,
+ EndTime: end,
+ UserID: filters.UserID,
+ APIKeyID: filters.APIKeyID,
+ AccountID: filters.AccountID,
+ GroupID: filters.GroupID,
+ Model: filters.Model,
+ BillingMode: filters.BillingMode,
+ RequestType: filters.RequestType,
+ Stream: filters.Stream,
+ BillingType: filters.BillingType,
+ })
+}
+
+// getStatsCached 命中则返回缓存,未命中则回源 usageService 并写缓存。
+func (h *UsageHandler) getStatsCached(ctx context.Context, filters usagestats.UsageLogFilters) (*usagestats.UsageStats, bool, error) {
+ key := usageStatsCacheKey(filters)
+ entry, hit, err := usageStatsCache.GetOrLoad(key, func() (any, error) {
+ return h.usageService.GetStatsWithFilters(ctx, filters)
+ })
+ if err != nil {
+ return nil, hit, err
+ }
+ stats, err := snapshotPayloadAs[*usagestats.UsageStats](entry.Payload)
+ return stats, hit, err
+}
diff --git a/backend/internal/handler/admin/usage_query_cache_test.go b/backend/internal/handler/admin/usage_query_cache_test.go
new file mode 100644
index 00000000..857e507a
--- /dev/null
+++ b/backend/internal/handler/admin/usage_query_cache_test.go
@@ -0,0 +1,28 @@
+package admin
+
+import (
+ "testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats"
+ "github.com/stretchr/testify/require"
+)
+
+func TestUsageStatsCacheKey_StableAndDistinct(t *testing.T) {
+ start := time.Date(2026, 5, 29, 0, 0, 0, 0, time.UTC)
+ end := time.Date(2026, 5, 31, 0, 0, 0, 0, time.UTC)
+ base := usagestats.UsageLogFilters{StartTime: &start, EndTime: &end, Model: "claude-3"}
+
+ k1 := usageStatsCacheKey(base)
+ k2 := usageStatsCacheKey(base)
+ require.NotEmpty(t, k1)
+ require.Equal(t, k1, k2, "same filters must produce same key")
+
+ other := base
+ other.Model = "gpt-4o"
+ require.NotEqual(t, k1, usageStatsCacheKey(other), "different model must change key")
+
+ withUser := base
+ withUser.UserID = 7
+ require.NotEqual(t, k1, usageStatsCacheKey(withUser), "different user must change key")
+}
diff --git a/backend/internal/repository/usage_log_repo.go b/backend/internal/repository/usage_log_repo.go
index 1835ebe5..b0992dae 100644
--- a/backend/internal/repository/usage_log_repo.go
+++ b/backend/internal/repository/usage_log_repo.go
@@ -27,6 +27,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/lib/pq"
gocache "github.com/patrickmn/go-cache"
+ "golang.org/x/sync/errgroup"
)
const usageLogSelectColumns = "id, user_id, api_key_id, account_id, request_id, model, requested_model, upstream_model, group_id, subscription_id, input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens, cache_creation_5m_tokens, cache_creation_1h_tokens, image_output_tokens, image_output_cost, input_cost, output_cost, cache_creation_cost, cache_read_cost, total_cost, actual_cost, rate_multiplier, account_rate_multiplier, billing_type, request_type, stream, openai_ws_mode, duration_ms, first_token_ms, user_agent, ip_address, image_count, image_size, image_input_size, image_output_size, image_size_source, image_size_breakdown, service_tier, reasoning_effort, inbound_endpoint, upstream_endpoint, cache_ttl_overridden, channel_id, model_mapping_chain, billing_tier, billing_mode, account_stats_cost, created_at"
@@ -3538,24 +3539,6 @@ func (r *usageLogRepository) GetStatsWithFilters(ctx context.Context, filters Us
stats := &UsageStats{}
var totalAccountCost float64
- if err := scanSingleRow(
- ctx,
- r.sql,
- query,
- args,
- &stats.TotalRequests,
- &stats.TotalInputTokens,
- &stats.TotalOutputTokens,
- &stats.TotalCacheTokens,
- &stats.TotalCost,
- &stats.TotalActualCost,
- &totalAccountCost,
- &stats.AverageDurationMs,
- ); err != nil {
- return nil, err
- }
- stats.TotalAccountCost = &totalAccountCost
- stats.TotalTokens = stats.TotalInputTokens + stats.TotalOutputTokens + stats.TotalCacheTokens
start := time.Unix(0, 0).UTC()
if filters.StartTime != nil {
@@ -3566,21 +3549,76 @@ func (r *usageLogRepository) GetStatsWithFilters(ctx context.Context, filters Us
end = *filters.EndTime
}
- endpoints, endpointErr := r.GetEndpointStatsWithFilters(ctx, start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType)
- if endpointErr != nil {
- logger.LegacyPrintf("repository.usage_log", "GetEndpointStatsWithFilters failed in GetStatsWithFilters: %v", endpointErr)
- endpoints = []EndpointStat{}
+ var endpoints, upstreamEndpoints, endpointPaths []EndpointStat
+
+ // 汇总查询:失败即致命。
+ runSummary := func(c context.Context) error {
+ return scanSingleRow(
+ c, r.sql, query, args,
+ &stats.TotalRequests,
+ &stats.TotalInputTokens,
+ &stats.TotalOutputTokens,
+ &stats.TotalCacheTokens,
+ &stats.TotalCost,
+ &stats.TotalActualCost,
+ &totalAccountCost,
+ &stats.AverageDurationMs,
+ )
}
- upstreamEndpoints, upstreamEndpointErr := r.GetUpstreamEndpointStatsWithFilters(ctx, start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType)
- if upstreamEndpointErr != nil {
- logger.LegacyPrintf("repository.usage_log", "GetUpstreamEndpointStatsWithFilters failed in GetStatsWithFilters: %v", upstreamEndpointErr)
- upstreamEndpoints = []EndpointStat{}
+ // endpoint 明细:best-effort(失败 log + 返空),不致命。
+ runEndpoints := func(c context.Context) {
+ res, err := r.GetEndpointStatsWithFilters(c, start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType)
+ if err != nil {
+ if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) {
+ logger.LegacyPrintf("repository.usage_log", "GetEndpointStatsWithFilters failed in GetStatsWithFilters: %v", err)
+ }
+ res = []EndpointStat{}
+ }
+ endpoints = res
}
- endpointPaths, endpointPathErr := r.getEndpointPathStatsWithFilters(ctx, start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType)
- if endpointPathErr != nil {
- logger.LegacyPrintf("repository.usage_log", "getEndpointPathStatsWithFilters failed in GetStatsWithFilters: %v", endpointPathErr)
- endpointPaths = []EndpointStat{}
+ runUpstream := func(c context.Context) {
+ res, err := r.GetUpstreamEndpointStatsWithFilters(c, start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType)
+ if err != nil {
+ if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) {
+ logger.LegacyPrintf("repository.usage_log", "GetUpstreamEndpointStatsWithFilters failed in GetStatsWithFilters: %v", err)
+ }
+ res = []EndpointStat{}
+ }
+ upstreamEndpoints = res
}
+ runPaths := func(c context.Context) {
+ res, err := r.getEndpointPathStatsWithFilters(c, start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType)
+ if err != nil {
+ if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) {
+ logger.LegacyPrintf("repository.usage_log", "getEndpointPathStatsWithFilters failed in GetStatsWithFilters: %v", err)
+ }
+ res = []EndpointStat{}
+ }
+ endpointPaths = res
+ }
+
+ if r.db != nil {
+ // 生产路径:r.sql 是 *sql.DB 连接池,可并发。4 条查询并行,延迟取最大值。
+ g, gctx := errgroup.WithContext(ctx)
+ g.Go(func() error { return runSummary(gctx) })
+ g.Go(func() error { runEndpoints(gctx); return nil })
+ g.Go(func() error { runUpstream(gctx); return nil })
+ g.Go(func() error { runPaths(gctx); return nil })
+ if err := g.Wait(); err != nil {
+ return nil, err
+ }
+ } else {
+ // 事务路径(ent.Tx 不能并发查询):顺序执行,行为与重构前一致。
+ if err := runSummary(ctx); err != nil {
+ return nil, err
+ }
+ runEndpoints(ctx)
+ runUpstream(ctx)
+ runPaths(ctx)
+ }
+
+ stats.TotalAccountCost = &totalAccountCost
+ stats.TotalTokens = stats.TotalInputTokens + stats.TotalOutputTokens + stats.TotalCacheTokens
stats.Endpoints = endpoints
stats.UpstreamEndpoints = upstreamEndpoints
stats.EndpointPaths = endpointPaths
diff --git a/backend/internal/repository/usage_log_repo_stats_integration_test.go b/backend/internal/repository/usage_log_repo_stats_integration_test.go
new file mode 100644
index 00000000..09ac2aee
--- /dev/null
+++ b/backend/internal/repository/usage_log_repo_stats_integration_test.go
@@ -0,0 +1,51 @@
+//go:build integration
+
+package repository
+
+import (
+ "context"
+ "testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats"
+ "github.com/Wei-Shaw/sub2api/internal/service"
+ "github.com/stretchr/testify/require"
+)
+
+func TestUsageLog_GetStatsWithFilters_AggregatesAndEndpoints(t *testing.T) {
+ ctx := context.Background()
+ tx := testEntTx(t)
+ client := tx.Client()
+ repo := newUsageLogRepositoryWithSQL(client, tx)
+
+ user := mustCreateUser(t, client, &service.User{Email: "stats@test.com"})
+ apiKey := mustCreateApiKey(t, client, &service.APIKey{UserID: user.ID, Key: "sk-stats-1", Name: "k"})
+ account := mustCreateAccount(t, client, &service.Account{Name: "acc-stats"})
+
+ now := time.Now().UTC()
+ inboundEndpoint := "/v1/messages"
+ upstreamEndpoint := "/v1/responses"
+ for i := 0; i < 3; i++ {
+ _, err := repo.Create(ctx, &service.UsageLog{
+ UserID: user.ID, APIKeyID: apiKey.ID, AccountID: account.ID,
+ Model: "claude-3", InputTokens: 2, OutputTokens: 3,
+ TotalCost: 0.5, ActualCost: 0.4, CreatedAt: now,
+ InboundEndpoint: &inboundEndpoint, UpstreamEndpoint: &upstreamEndpoint,
+ })
+ require.NoError(t, err)
+ }
+
+ start := now.Add(-1 * time.Hour)
+ end := now.Add(1 * time.Hour)
+ // 按本测试创建的 user 维度过滤:集成库为共享实例,其它用 testEntClient 的兄弟测试会留下
+ // 已提交的 usage_log 行(含零 token 的失败请求),不限定 user 会把它们计入 TotalRequests。
+ stats, err := repo.GetStatsWithFilters(ctx, usagestats.UsageLogFilters{UserID: user.ID, StartTime: &start, EndTime: &end})
+ require.NoError(t, err)
+ require.Equal(t, int64(3), stats.TotalRequests)
+ require.Equal(t, int64(6), stats.TotalInputTokens)
+ require.Equal(t, int64(9), stats.TotalOutputTokens)
+ require.InDelta(t, 1.2, stats.TotalActualCost, 1e-9)
+ require.NotEmpty(t, stats.Endpoints)
+ require.NotEmpty(t, stats.UpstreamEndpoints)
+ require.NotEmpty(t, stats.EndpointPaths)
+}
diff --git a/frontend/src/api/admin/usage.ts b/frontend/src/api/admin/usage.ts
index c3a38c0a..d933ac63 100644
--- a/frontend/src/api/admin/usage.ts
+++ b/frontend/src/api/admin/usage.ts
@@ -121,6 +121,7 @@ export async function getStats(params: {
start_date?: string
end_date?: string
timezone?: string
+ nocache?: number
}): Promise {
const { data } = await apiClient.get('/admin/usage/stats', {
params
diff --git a/frontend/src/components/admin/usage/UsageFilters.vue b/frontend/src/components/admin/usage/UsageFilters.vue
index 40f2c9f8..a800f190 100644
--- a/frontend/src/components/admin/usage/UsageFilters.vue
+++ b/frontend/src/components/admin/usage/UsageFilters.vue
@@ -168,7 +168,7 @@
diff --git a/frontend/src/views/admin/__tests__/UsageView.spec.ts b/frontend/src/views/admin/__tests__/UsageView.spec.ts
index 9cdb2a3c..8c644a75 100644
--- a/frontend/src/views/admin/__tests__/UsageView.spec.ts
+++ b/frontend/src/views/admin/__tests__/UsageView.spec.ts
@@ -3,7 +3,7 @@ import { flushPromises, mount } from '@vue/test-utils'
import UsageView from '../UsageView.vue'
-const { list, getStats, getSnapshotV2, getById } = vi.hoisted(() => {
+const { list, getStats, getSnapshotV2, getById, getModelStats } = vi.hoisted(() => {
vi.stubGlobal('localStorage', {
getItem: vi.fn(() => null),
setItem: vi.fn(),
@@ -15,6 +15,7 @@ const { list, getStats, getSnapshotV2, getById } = vi.hoisted(() => {
getStats: vi.fn(),
getSnapshotV2: vi.fn(),
getById: vi.fn(),
+ getModelStats: vi.fn(),
}
})
@@ -40,6 +41,7 @@ vi.mock('@/api/admin', () => ({
},
dashboard: {
getSnapshotV2,
+ getModelStats,
},
users: {
getById,
@@ -116,6 +118,7 @@ describe('admin UsageView distribution metric toggles', () => {
getStats.mockReset()
getSnapshotV2.mockReset()
getById.mockReset()
+ getModelStats.mockReset()
list.mockResolvedValue({
items: [],
@@ -137,12 +140,44 @@ describe('admin UsageView distribution metric toggles', () => {
models: [],
groups: [],
})
+ getModelStats.mockResolvedValue({ models: [] })
})
afterEach(() => {
vi.useRealTimers()
})
+ it('keeps previous model stats visible during refresh until new data arrives', async () => {
+ // 首次加载返回 A
+ getModelStats.mockResolvedValueOnce({ models: [{ model: 'A', total_tokens: 10 }] })
+
+ const wrapper = mount(UsageView, {
+ global: { stubs: {
+ AppLayout: AppLayoutStub, UsageStatsCards: true, UsageFilters: UsageFiltersStub,
+ UsageTable: true, UsageExportProgress: true, UsageCleanupDialog: true,
+ UserBalanceHistoryModal: true, AuditLogModal: true, Pagination: true, Select: true,
+ DateRangePicker: true, Icon: true, TokenUsageTrend: true,
+ ModelDistributionChart: ModelDistributionChartStub, GroupDistributionChart: GroupDistributionChartStub,
+ EndpointDistributionChart: true,
+ } },
+ })
+ vi.advanceTimersByTime(120)
+ await flushPromises()
+ expect((wrapper.vm as any).requestedModelStats).toEqual([{ model: 'A', total_tokens: 10 }])
+
+ // 刷新:让第二次 getModelStats 处于 pending,断言旧数据 A 仍在(不被清空成 [])
+ let resolveSecond: (v: any) => void = () => {}
+ getModelStats.mockReturnValueOnce(new Promise((res) => { resolveSecond = res }))
+ ;(wrapper.vm as any).refreshData()
+ await flushPromises()
+ expect((wrapper.vm as any).requestedModelStats).toEqual([{ model: 'A', total_tokens: 10 }])
+
+ // 新数据到达后替换为 B
+ resolveSecond({ models: [{ model: 'B', total_tokens: 20 }] })
+ await flushPromises()
+ expect((wrapper.vm as any).requestedModelStats).toEqual([{ model: 'B', total_tokens: 20 }])
+ })
+
it('keeps model and group metric toggles independent without refetching chart data', async () => {
const wrapper = mount(UsageView, {
global: {
From d8cbf9ab5cb9c5ec4d3511e63d642384a6c60e62 Mon Sep 17 00:00:00 2001
From: name <136912576+is7Qin@users.noreply.github.com>
Date: Fri, 29 May 2026 21:05:47 +0800
Subject: [PATCH 48/79] refactor(gateway): introduce request body refs
---
backend/internal/handler/gateway_handler.go | 12 +++--
.../gateway_handler_chat_completions.go | 5 +-
.../handler/gateway_handler_responses.go | 5 +-
.../internal/handler/gemini_v1beta_handler.go | 2 +-
...teway_anthropic_apikey_passthrough_test.go | 16 +++----
backend/internal/service/gateway_request.go | 46 +++++++++++++++----
.../internal/service/gateway_request_test.go | 44 +++++++++---------
backend/internal/service/gateway_service.go | 12 ++---
.../service/gateway_service_benchmark_test.go | 2 +-
.../service/gateway_websearch_emulation.go | 2 +-
.../service/generate_session_hash_test.go | 10 ++--
11 files changed, 95 insertions(+), 61 deletions(-)
diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go
index eb5c4a42..79bed8b9 100644
--- a/backend/internal/handler/gateway_handler.go
+++ b/backend/internal/handler/gateway_handler.go
@@ -154,7 +154,8 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
setOpsRequestContext(c, "", false)
- parsedReq, err := service.ParseGatewayRequest(body, domain.PlatformAnthropic)
+ bodyRef := service.NewRequestBodyRef(body)
+ parsedReq, err := service.ParseGatewayRequest(bodyRef, domain.PlatformAnthropic)
if err != nil {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
return
@@ -746,11 +747,11 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
// 应用渠道模型映射到请求
if channelMapping.Mapped {
parsedReq.Model = channelMapping.MappedModel
- parsedReq.Body = h.gatewayService.ReplaceModelInBody(parsedReq.Body, channelMapping.MappedModel)
+ parsedReq.Body.Replace(h.gatewayService.ReplaceModelInBody(parsedReq.Body.Bytes(), channelMapping.MappedModel))
}
// Bedrock CC 兼容:渠道模型映射后,清理 Anthropic API 专有字段、注入 Bedrock 必需字段
- parsedReq.Body = h.gatewayService.ApplyBedrockCCCompat(c.Request.Context(), parsedReq.Body, parsedReq.Model, account, apiKey.GroupID)
- body = parsedReq.Body
+ parsedReq.Body.Replace(h.gatewayService.ApplyBedrockCCCompat(c.Request.Context(), parsedReq.Body.Bytes(), parsedReq.Model, account, apiKey.GroupID))
+ body = parsedReq.Body.Bytes()
// 转发请求 - 根据账号平台分流
c.Set("parsed_request", parsedReq)
@@ -1683,7 +1684,8 @@ func (h *GatewayHandler) CountTokens(c *gin.Context) {
setOpsRequestContext(c, "", false)
- parsedReq, err := service.ParseGatewayRequest(body, domain.PlatformAnthropic)
+ bodyRef := service.NewRequestBodyRef(body)
+ parsedReq, err := service.ParseGatewayRequest(bodyRef, domain.PlatformAnthropic)
if err != nil {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
return
diff --git a/backend/internal/handler/gateway_handler_chat_completions.go b/backend/internal/handler/gateway_handler_chat_completions.go
index daf6e6ea..719700aa 100644
--- a/backend/internal/handler/gateway_handler_chat_completions.go
+++ b/backend/internal/handler/gateway_handler_chat_completions.go
@@ -151,9 +151,10 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) {
}
// Parse request for session hash
- parsedReq, _ := service.ParseGatewayRequest(body, "chat_completions")
+ bodyRef := service.NewRequestBodyRef(body)
+ parsedReq, _ := service.ParseGatewayRequest(bodyRef, "chat_completions")
if parsedReq == nil {
- parsedReq = &service.ParsedRequest{Model: reqModel, Stream: reqStream, Body: body}
+ parsedReq = &service.ParsedRequest{Model: reqModel, Stream: reqStream, Body: bodyRef}
}
parsedReq.SessionContext = &service.SessionContext{
ClientIP: ip.GetClientIP(c),
diff --git a/backend/internal/handler/gateway_handler_responses.go b/backend/internal/handler/gateway_handler_responses.go
index f57b9989..49f80d19 100644
--- a/backend/internal/handler/gateway_handler_responses.go
+++ b/backend/internal/handler/gateway_handler_responses.go
@@ -156,9 +156,10 @@ func (h *GatewayHandler) Responses(c *gin.Context) {
}
// Parse request for session hash
- parsedReq, _ := service.ParseGatewayRequest(body, "responses")
+ bodyRef := service.NewRequestBodyRef(body)
+ parsedReq, _ := service.ParseGatewayRequest(bodyRef, "responses")
if parsedReq == nil {
- parsedReq = &service.ParsedRequest{Model: reqModel, Stream: reqStream, Body: body}
+ parsedReq = &service.ParsedRequest{Model: reqModel, Stream: reqStream, Body: bodyRef}
}
parsedReq.SessionContext = &service.SessionContext{
ClientIP: ip.GetClientIP(c),
diff --git a/backend/internal/handler/gemini_v1beta_handler.go b/backend/internal/handler/gemini_v1beta_handler.go
index 0b33ca3e..5d8e6fa8 100644
--- a/backend/internal/handler/gemini_v1beta_handler.go
+++ b/backend/internal/handler/gemini_v1beta_handler.go
@@ -262,7 +262,7 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
sessionHash := extractGeminiCLISessionHash(c, body)
if sessionHash == "" {
// Fallback: 使用通用的会话哈希生成逻辑(适用于其他客户端)
- parsedReq, _ := service.ParseGatewayRequest(body, domain.PlatformGemini)
+ parsedReq, _ := service.ParseGatewayRequest(service.NewRequestBodyRef(body), domain.PlatformGemini)
if parsedReq != nil {
parsedReq.SessionContext = &service.SessionContext{
ClientIP: ip.GetClientIP(c),
diff --git a/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go b/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go
index 9062c517..a67a3dc2 100644
--- a/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go
+++ b/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go
@@ -112,7 +112,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_ForwardStreamPreservesBodyAnd
body := []byte(`{"model":"claude-3-7-sonnet-20250219","stream":true,"system":[{"type":"text","text":"x-anthropic-billing-header keep"}],"messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}]}`)
parsed := &ParsedRequest{
- Body: body,
+ Body: NewRequestBodyRef(body),
Model: "claude-3-7-sonnet-20250219",
Stream: true,
}
@@ -202,7 +202,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_ForwardCountTokensPreservesBo
body := []byte(`{"model":"claude-3-5-sonnet-latest","messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}],"thinking":{"type":"enabled"}}`)
parsed := &ParsedRequest{
- Body: body,
+ Body: NewRequestBodyRef(body),
Model: "claude-3-5-sonnet-latest",
}
@@ -344,7 +344,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_ModelMappingEdgeCases(t *test
body := []byte(`{"model":"` + tt.model + `","messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}]}`)
parsed := &ParsedRequest{
- Body: body,
+ Body: NewRequestBodyRef(body),
Model: tt.model,
}
@@ -429,7 +429,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_ModelMappingPreservesOtherFie
// 包含复杂字段的请求体:system、thinking、messages
body := []byte(`{"model":"claude-sonnet-4-20250514","system":[{"type":"text","text":"You are a helpful assistant."}],"messages":[{"role":"user","content":[{"type":"text","text":"hello world"}]}],"thinking":{"type":"enabled","budget_tokens":5000},"max_tokens":1024}`)
parsed := &ParsedRequest{
- Body: body,
+ Body: NewRequestBodyRef(body),
Model: "claude-sonnet-4-20250514",
}
@@ -485,7 +485,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_CountTokensFiltersGenerationF
body := []byte(`{"model":"claude-sonnet-4-20250514","system":[{"type":"text","text":"sys"}],"messages":[{"role":"user","content":"hello"}],"tools":[{"name":"tool","input_schema":{"type":"object"}}],"temperature":0.7,"top_p":0.9,"top_k":40,"stream":true,"stop_sequences":["END"],"max_tokens":1024,"thinking":{"type":"enabled","budget_tokens":5000}}`)
parsed := &ParsedRequest{
- Body: body,
+ Body: NewRequestBodyRef(body),
Model: "claude-sonnet-4-20250514",
}
@@ -547,7 +547,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_EmptyModelSkipsMapping(t *tes
body := []byte(`{"messages":[{"role":"user","content":"hello"}]}`)
parsed := &ParsedRequest{
- Body: body,
+ Body: NewRequestBodyRef(body),
Model: "", // 空模型
}
@@ -636,7 +636,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_CountTokens404PassthroughNotE
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", nil)
body := []byte(`{"model":"claude-sonnet-4-5-20250929","messages":[{"role":"user","content":"hi"}]}`)
- parsed := &ParsedRequest{Body: body, Model: "claude-sonnet-4-5-20250929"}
+ parsed := &ParsedRequest{Body: NewRequestBodyRef(body), Model: "claude-sonnet-4-5-20250929"}
upstream := &anthropicHTTPUpstreamRecorder{
resp: &http.Response{
@@ -767,7 +767,7 @@ func TestGatewayService_AnthropicOAuth_ForwardPreservesBillingHeaderSystemBlock(
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
- parsed, err := ParseGatewayRequest([]byte(tt.body), PlatformAnthropic)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(tt.body)), PlatformAnthropic)
require.NoError(t, err)
upstream := &anthropicHTTPUpstreamRecorder{
diff --git a/backend/internal/service/gateway_request.go b/backend/internal/service/gateway_request.go
index 91f7601c..819bb0a8 100644
--- a/backend/internal/service/gateway_request.go
+++ b/backend/internal/service/gateway_request.go
@@ -51,6 +51,35 @@ type SessionContext struct {
APIKeyID int64
}
+type RequestBodyRef struct {
+ data []byte
+}
+
+func NewRequestBodyRef(data []byte) *RequestBodyRef {
+ return &RequestBodyRef{data: data}
+}
+
+func (b *RequestBodyRef) Bytes() []byte {
+ if b == nil {
+ return nil
+ }
+ return b.data
+}
+
+func (b *RequestBodyRef) Len() int {
+ if b == nil {
+ return 0
+ }
+ return len(b.data)
+}
+
+func (b *RequestBodyRef) Replace(data []byte) {
+ if b == nil {
+ return
+ }
+ b.data = data
+}
+
// ParsedRequest 保存网关请求的预解析结果
//
// 性能优化说明:
@@ -64,7 +93,7 @@ type SessionContext struct {
// 2. 将解析结果 ParsedRequest 传递给 Service 层
// 3. 避免重复 json.Unmarshal,减少 CPU 和内存开销
type ParsedRequest struct {
- Body []byte // 原始请求体(保留用于转发)
+ Body *RequestBodyRef // 原始请求体引用(保留用于转发)
Model string // 请求的模型名称
Stream bool // 是否为流式请求
MetadataUserID string // metadata.user_id(用于会话亲和)
@@ -130,17 +159,18 @@ func normalizeSessionUserAgentFallback(raw string) string {
// ParseGatewayRequest 解析网关请求体并返回结构化结果。
// protocol 指定请求协议格式(domain.PlatformAnthropic / domain.PlatformGemini),
// 不同协议使用不同的 system/messages 字段名。
-func ParseGatewayRequest(body []byte, protocol string) (*ParsedRequest, error) {
+func ParseGatewayRequest(body *RequestBodyRef, protocol string) (*ParsedRequest, error) {
+ bodyBytes := body.Bytes()
// 保持与旧实现一致:请求体必须是合法 JSON。
// 注意:gjson.GetBytes 对非法 JSON 不会报错,因此需要显式校验。
- if !gjson.ValidBytes(body) {
+ if !gjson.ValidBytes(bodyBytes) {
return nil, fmt.Errorf("invalid json")
}
// 性能:
// - gjson.GetBytes 会把匹配的 Raw/Str 安全复制成 string(对于巨大 messages 会产生额外拷贝)。
// - 这里将 body 通过 unsafe 零拷贝视为 string,仅在本函数内使用,且 body 不会被修改。
- jsonStr := *(*string)(unsafe.Pointer(&body))
+ jsonStr := *(*string)(unsafe.Pointer(&bodyBytes))
parsed := &ParsedRequest{
Body: body,
@@ -197,7 +227,7 @@ func ParseGatewayRequest(body []byte, protocol string) (*ParsedRequest, error) {
// Gemini 原生格式: systemInstruction.parts / contents
if sysParts := gjson.Get(jsonStr, "systemInstruction.parts"); sysParts.Exists() && sysParts.IsArray() {
var parts []any
- if err := json.Unmarshal(sliceRawFromBody(body, sysParts), &parts); err != nil {
+ if err := json.Unmarshal(sliceRawFromBody(bodyBytes, sysParts), &parts); err != nil {
return nil, err
}
parsed.System = parts
@@ -205,7 +235,7 @@ func ParseGatewayRequest(body []byte, protocol string) (*ParsedRequest, error) {
if contents := gjson.Get(jsonStr, "contents"); contents.Exists() && contents.IsArray() {
var msgs []any
- if err := json.Unmarshal(sliceRawFromBody(body, contents), &msgs); err != nil {
+ if err := json.Unmarshal(sliceRawFromBody(bodyBytes, contents), &msgs); err != nil {
return nil, err
}
parsed.Messages = msgs
@@ -224,7 +254,7 @@ func ParseGatewayRequest(body []byte, protocol string) (*ParsedRequest, error) {
parsed.System = sys.String()
default:
var system any
- if err := json.Unmarshal(sliceRawFromBody(body, sys), &system); err != nil {
+ if err := json.Unmarshal(sliceRawFromBody(bodyBytes, sys), &system); err != nil {
return nil, err
}
parsed.System = system
@@ -233,7 +263,7 @@ func ParseGatewayRequest(body []byte, protocol string) (*ParsedRequest, error) {
if msgs := gjson.Get(jsonStr, "messages"); msgs.Exists() && msgs.IsArray() {
var messages []any
- if err := json.Unmarshal(sliceRawFromBody(body, msgs), &messages); err != nil {
+ if err := json.Unmarshal(sliceRawFromBody(bodyBytes, msgs), &messages); err != nil {
return nil, err
}
parsed.Messages = messages
diff --git a/backend/internal/service/gateway_request_test.go b/backend/internal/service/gateway_request_test.go
index 045dc66c..d415b871 100644
--- a/backend/internal/service/gateway_request_test.go
+++ b/backend/internal/service/gateway_request_test.go
@@ -14,7 +14,7 @@ import (
func TestParseGatewayRequest(t *testing.T) {
body := []byte(`{"model":"claude-3-7-sonnet","stream":true,"metadata":{"user_id":"session_123e4567-e89b-12d3-a456-426614174000"},"system":[{"type":"text","text":"hello","cache_control":{"type":"ephemeral"}}],"messages":[{"content":"hi"}]}`)
- parsed, err := ParseGatewayRequest(body, "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.NoError(t, err)
require.Equal(t, "claude-3-7-sonnet", parsed.Model)
require.True(t, parsed.Stream)
@@ -27,7 +27,7 @@ func TestParseGatewayRequest(t *testing.T) {
func TestParseGatewayRequest_ThinkingEnabled(t *testing.T) {
body := []byte(`{"model":"claude-sonnet-4-5","thinking":{"type":"enabled"},"messages":[{"content":"hi"}]}`)
- parsed, err := ParseGatewayRequest(body, "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.NoError(t, err)
require.Equal(t, "claude-sonnet-4-5", parsed.Model)
require.True(t, parsed.ThinkingEnabled)
@@ -35,7 +35,7 @@ func TestParseGatewayRequest_ThinkingEnabled(t *testing.T) {
func TestParseGatewayRequest_ThinkingAdaptiveEnabled(t *testing.T) {
body := []byte(`{"model":"claude-sonnet-4-5","thinking":{"type":"adaptive"},"messages":[{"content":"hi"}]}`)
- parsed, err := ParseGatewayRequest(body, "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.NoError(t, err)
require.Equal(t, "claude-sonnet-4-5", parsed.Model)
require.True(t, parsed.ThinkingEnabled)
@@ -43,21 +43,21 @@ func TestParseGatewayRequest_ThinkingAdaptiveEnabled(t *testing.T) {
func TestParseGatewayRequest_MaxTokens(t *testing.T) {
body := []byte(`{"model":"claude-haiku-4-5","max_tokens":1}`)
- parsed, err := ParseGatewayRequest(body, "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.NoError(t, err)
require.Equal(t, 1, parsed.MaxTokens)
}
func TestParseGatewayRequest_MaxTokensNonIntegralIgnored(t *testing.T) {
body := []byte(`{"model":"claude-haiku-4-5","max_tokens":1.5}`)
- parsed, err := ParseGatewayRequest(body, "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.NoError(t, err)
require.Equal(t, 0, parsed.MaxTokens)
}
func TestParseGatewayRequest_SystemNull(t *testing.T) {
body := []byte(`{"model":"claude-3","system":null}`)
- parsed, err := ParseGatewayRequest(body, "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.NoError(t, err)
// 显式传入 system:null 也应视为“字段已存在”,避免默认 system 被注入。
require.True(t, parsed.HasSystem)
@@ -66,13 +66,13 @@ func TestParseGatewayRequest_SystemNull(t *testing.T) {
func TestParseGatewayRequest_InvalidModelType(t *testing.T) {
body := []byte(`{"model":123}`)
- _, err := ParseGatewayRequest(body, "")
+ _, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.Error(t, err)
}
func TestParseGatewayRequest_InvalidStreamType(t *testing.T) {
body := []byte(`{"stream":"true"}`)
- _, err := ParseGatewayRequest(body, "")
+ _, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.Error(t, err)
}
@@ -86,7 +86,7 @@ func TestParseGatewayRequest_GeminiContents(t *testing.T) {
{"role": "user", "parts": [{"text": "How are you?"}]}
]
}`)
- parsed, err := ParseGatewayRequest(body, domain.PlatformGemini)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err)
require.Len(t, parsed.Messages, 3, "should parse contents as Messages")
require.False(t, parsed.HasSystem, "Gemini format should not set HasSystem")
@@ -102,7 +102,7 @@ func TestParseGatewayRequest_GeminiSystemInstruction(t *testing.T) {
{"role": "user", "parts": [{"text": "Hello"}]}
]
}`)
- parsed, err := ParseGatewayRequest(body, domain.PlatformGemini)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err)
require.NotNil(t, parsed.System, "should parse systemInstruction.parts as System")
parts, ok := parsed.System.([]any)
@@ -119,7 +119,7 @@ func TestParseGatewayRequest_GeminiWithModel(t *testing.T) {
"model": "gemini-2.5-pro",
"contents": [{"role": "user", "parts": [{"text": "test"}]}]
}`)
- parsed, err := ParseGatewayRequest(body, domain.PlatformGemini)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err)
require.Equal(t, "gemini-2.5-pro", parsed.Model)
require.Len(t, parsed.Messages, 1)
@@ -132,7 +132,7 @@ func TestParseGatewayRequest_GeminiIgnoresAnthropicFields(t *testing.T) {
"messages": [{"role": "user", "content": "ignored"}],
"contents": [{"role": "user", "parts": [{"text": "real content"}]}]
}`)
- parsed, err := ParseGatewayRequest(body, domain.PlatformGemini)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err)
require.False(t, parsed.HasSystem, "Gemini protocol should not parse Anthropic system field")
require.Nil(t, parsed.System, "no systemInstruction = nil System")
@@ -141,14 +141,14 @@ func TestParseGatewayRequest_GeminiIgnoresAnthropicFields(t *testing.T) {
func TestParseGatewayRequest_GeminiEmptyContents(t *testing.T) {
body := []byte(`{"contents": []}`)
- parsed, err := ParseGatewayRequest(body, domain.PlatformGemini)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err)
require.Empty(t, parsed.Messages)
}
func TestParseGatewayRequest_GeminiNoContents(t *testing.T) {
body := []byte(`{"model": "gemini-2.5-flash"}`)
- parsed, err := ParseGatewayRequest(body, domain.PlatformGemini)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err)
require.Nil(t, parsed.Messages)
require.Equal(t, "gemini-2.5-flash", parsed.Model)
@@ -162,7 +162,7 @@ func TestParseGatewayRequest_AnthropicIgnoresGeminiFields(t *testing.T) {
"contents": [{"role": "user", "parts": [{"text": "ignored"}]}],
"systemInstruction": {"parts": [{"text": "ignored"}]}
}`)
- parsed, err := ParseGatewayRequest(body, domain.PlatformAnthropic)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformAnthropic)
require.NoError(t, err)
require.True(t, parsed.HasSystem)
require.Equal(t, "real system", parsed.System)
@@ -897,7 +897,7 @@ func TestParseGatewayRequest_TypeValidation(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
- _, err := ParseGatewayRequest([]byte(tt.body), "")
+ _, err := ParseGatewayRequest(NewRequestBodyRef([]byte(tt.body)), "")
if tt.wantErr {
require.Error(t, err)
if tt.errSubstr != "" {
@@ -959,7 +959,7 @@ func TestParseGatewayRequest_OptionalFieldsMissing(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
- parsed, err := ParseGatewayRequest([]byte(tt.body), "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(tt.body)), "")
require.NoError(t, err)
require.Equal(t, tt.wantModel, parsed.Model)
@@ -1023,7 +1023,7 @@ func TestParseGatewayRequest_MaxTokensBoundary(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
- parsed, err := ParseGatewayRequest([]byte(tt.body), "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(tt.body)), "")
if tt.wantErr {
require.Error(t, err)
return
@@ -1040,7 +1040,7 @@ func TestParseGatewayRequest_MaxTokensBoundary(t *testing.T) {
// 核心路径:先 Unmarshal 到 map[string]any,再逐字段提取。
func parseGatewayRequestOld(body []byte, protocol string) (*ParsedRequest, error) {
parsed := &ParsedRequest{
- Body: body,
+ Body: NewRequestBodyRef(body),
}
var req map[string]any
@@ -1151,7 +1151,7 @@ func BenchmarkParseGatewayRequest_New_Small(b *testing.B) {
b.SetBytes(int64(len(data)))
b.ResetTimer()
for i := 0; i < b.N; i++ {
- _, _ = ParseGatewayRequest(data, "")
+ _, _ = ParseGatewayRequest(NewRequestBodyRef(data), "")
}
}
@@ -1203,7 +1203,7 @@ func TestParseGatewayRequest_OutputEffort(t *testing.T) {
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
- parsed, err := ParseGatewayRequest([]byte(tt.body), "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(tt.body)), "")
require.NoError(t, err)
require.Equal(t, tt.wantEffort, parsed.OutputEffort)
})
@@ -1245,6 +1245,6 @@ func BenchmarkParseGatewayRequest_New_Large(b *testing.B) {
b.SetBytes(int64(len(data)))
b.ResetTimer()
for i := 0; i < b.N; i++ {
- _, _ = ParseGatewayRequest(data, "")
+ _, _ = ParseGatewayRequest(NewRequestBodyRef(data), "")
}
}
diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go
index f807f3ec..8f55bf13 100644
--- a/backend/internal/service/gateway_service.go
+++ b/backend/internal/service/gateway_service.go
@@ -4400,12 +4400,12 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
}
// Web Search 模拟:纯 web_search 请求时,直接调用搜索 API 构造响应
- if account != nil && s.shouldEmulateWebSearch(ctx, account, parsed.GroupID, parsed.Body) {
+ if account != nil && s.shouldEmulateWebSearch(ctx, account, parsed.GroupID, parsed.Body.Bytes()) {
return s.handleWebSearchEmulation(ctx, c, account, parsed)
}
if account != nil && account.IsAnthropicAPIKeyPassthroughEnabled() {
- passthroughBody := parsed.Body
+ passthroughBody := parsed.Body.Bytes()
passthroughModel := parsed.Model
if passthroughModel != "" {
if mappedModel := account.GetMappedModel(passthroughModel); mappedModel != passthroughModel {
@@ -4441,7 +4441,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
c.Set(betaPolicyFilterSetKey, filterSet)
}
- body := parsed.Body
+ body := parsed.Body.Bytes()
reqModel := parsed.Model
reqStream := parsed.Stream
originalModel := reqModel
@@ -5735,7 +5735,7 @@ func (s *GatewayService) forwardBedrock(
) (*ForwardResult, error) {
reqModel := parsed.Model
reqStream := parsed.Stream
- body := parsed.Body
+ body := parsed.Body.Bytes()
region := bedrockRuntimeRegion(account)
mappedModel, ok := ResolveBedrockModelID(account, reqModel)
@@ -9172,7 +9172,7 @@ func (s *GatewayService) ForwardCountTokens(ctx context.Context, c *gin.Context,
}
if account != nil && account.IsAnthropicAPIKeyPassthroughEnabled() {
- passthroughBody := parsed.Body
+ passthroughBody := parsed.Body.Bytes()
if reqModel := parsed.Model; reqModel != "" {
if mappedModel := account.GetMappedModel(reqModel); mappedModel != reqModel {
passthroughBody = s.replaceModelInBody(passthroughBody, mappedModel)
@@ -9188,7 +9188,7 @@ func (s *GatewayService) ForwardCountTokens(ctx context.Context, c *gin.Context,
return nil
}
- body := parsed.Body
+ body := parsed.Body.Bytes()
reqModel := parsed.Model
// Pre-filter: strip empty text blocks to prevent upstream 400.
diff --git a/backend/internal/service/gateway_service_benchmark_test.go b/backend/internal/service/gateway_service_benchmark_test.go
index c9c4d3dd..5637680b 100644
--- a/backend/internal/service/gateway_service_benchmark_test.go
+++ b/backend/internal/service/gateway_service_benchmark_test.go
@@ -14,7 +14,7 @@ func BenchmarkGenerateSessionHash_Metadata(b *testing.B) {
b.ReportAllocs()
for i := 0; i < b.N; i++ {
- parsed, err := ParseGatewayRequest(body, "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
if err != nil {
b.Fatalf("解析请求失败: %v", err)
}
diff --git a/backend/internal/service/gateway_websearch_emulation.go b/backend/internal/service/gateway_websearch_emulation.go
index a42b5585..2f9c8e0c 100644
--- a/backend/internal/service/gateway_websearch_emulation.go
+++ b/backend/internal/service/gateway_websearch_emulation.go
@@ -150,7 +150,7 @@ func (s *GatewayService) handleWebSearchEmulation(
parsed.OnUpstreamAccepted()
}
- query := extractSearchQueryFromBody(parsed.Body)
+ query := extractSearchQueryFromBody(parsed.Body.Bytes())
if query == "" {
return nil, fmt.Errorf("web search emulation: no query found in messages")
}
diff --git a/backend/internal/service/generate_session_hash_test.go b/backend/internal/service/generate_session_hash_test.go
index 39679c3d..8f3258b7 100644
--- a/backend/internal/service/generate_session_hash_test.go
+++ b/backend/internal/service/generate_session_hash_test.go
@@ -1198,7 +1198,7 @@ func TestGenerateSessionHash_GeminiMultiTurnHashNotSticky(t *testing.T) {
hashes := make([]string, 3)
for i, body := range [][]byte{round1Body, round2Body, round3Body} {
- parsed, err := ParseGatewayRequest(body, "gemini")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "gemini")
require.NoError(t, err)
parsed.SessionContext = ctx
hashes[i] = svc.GenerateSessionHash(parsed)
@@ -1211,7 +1211,7 @@ func TestGenerateSessionHash_GeminiMultiTurnHashNotSticky(t *testing.T) {
require.NotEqual(t, hashes[0], hashes[2], "round 1 vs 3 hash should differ")
// 同一轮重试应产生相同 hash
- parsed1Again, err := ParseGatewayRequest(round2Body, "gemini")
+ parsed1Again, err := ParseGatewayRequest(NewRequestBodyRef(round2Body), "gemini")
require.NoError(t, err)
parsed1Again.SessionContext = ctx
h2Again := svc.GenerateSessionHash(parsed1Again)
@@ -1234,7 +1234,7 @@ func TestGenerateSessionHash_GeminiEndToEnd(t *testing.T) {
]
}`)
- parsed, err := ParseGatewayRequest(body, "gemini")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "gemini")
require.NoError(t, err)
parsed.SessionContext = &SessionContext{
ClientIP: "10.0.0.1",
@@ -1246,7 +1246,7 @@ func TestGenerateSessionHash_GeminiEndToEnd(t *testing.T) {
require.NotEmpty(t, h, "end-to-end Gemini flow should produce a hash")
// 同一请求再次解析应产生相同 hash
- parsed2, err := ParseGatewayRequest(body, "gemini")
+ parsed2, err := ParseGatewayRequest(NewRequestBodyRef(body), "gemini")
require.NoError(t, err)
parsed2.SessionContext = &SessionContext{
ClientIP: "10.0.0.1",
@@ -1258,7 +1258,7 @@ func TestGenerateSessionHash_GeminiEndToEnd(t *testing.T) {
require.Equal(t, h, h2, "same request should produce same hash")
// 不同用户发送相同请求应产生不同 hash
- parsed3, err := ParseGatewayRequest(body, "gemini")
+ parsed3, err := ParseGatewayRequest(NewRequestBodyRef(body), "gemini")
require.NoError(t, err)
parsed3.SessionContext = &SessionContext{
ClientIP: "10.0.0.2",
From b1c4be4ac81519df4c562f9563e7be8fe94ada2a Mon Sep 17 00:00:00 2001
From: name <136912576+is7Qin@users.noreply.github.com>
Date: Sat, 30 May 2026 00:37:34 +0800
Subject: [PATCH 49/79] refactor(gateway): remove parsed request object graphs
Keep large gateway payloads as raw body ranges and bind OpenAI parsed-body caches to the body bytes so failover and mapping do not reuse stale mutable state.
---
backend/internal/handler/gateway_handler.go | 4 +-
backend/internal/handler/gateway_helper.go | 14 +-
.../handler/gateway_helper_hotpath_test.go | 11 +-
.../handler/openai_gateway_handler.go | 2 +-
backend/internal/service/anthropic_session.go | 45 +-
.../service/anthropic_session_test.go | 148 +--
.../service/gateway_oauth_metadata_test.go | 2 -
backend/internal/service/gateway_request.go | 202 +++-
.../internal/service/gateway_request_test.go | 68 +-
backend/internal/service/gateway_service.go | 199 ++--
.../service/gateway_service_benchmark_test.go | 24 +-
.../service/generate_session_hash_test.go | 1009 +++--------------
.../service/openai_gateway_service.go | 69 +-
.../openai_gateway_service_hotpath_test.go | 20 +-
.../service/openai_oauth_passthrough_test.go | 52 +
.../service/user_msg_queue_service.go | 58 +-
16 files changed, 699 insertions(+), 1228 deletions(-)
diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go
index 79bed8b9..10d67ba3 100644
--- a/backend/internal/handler/gateway_handler.go
+++ b/backend/internal/handler/gateway_handler.go
@@ -747,10 +747,10 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
// 应用渠道模型映射到请求
if channelMapping.Mapped {
parsedReq.Model = channelMapping.MappedModel
- parsedReq.Body.Replace(h.gatewayService.ReplaceModelInBody(parsedReq.Body.Bytes(), channelMapping.MappedModel))
+ parsedReq.ReplaceBody(h.gatewayService.ReplaceModelInBody(parsedReq.Body.Bytes(), channelMapping.MappedModel))
}
// Bedrock CC 兼容:渠道模型映射后,清理 Anthropic API 专有字段、注入 Bedrock 必需字段
- parsedReq.Body.Replace(h.gatewayService.ApplyBedrockCCCompat(c.Request.Context(), parsedReq.Body.Bytes(), parsedReq.Model, account, apiKey.GroupID))
+ parsedReq.ReplaceBody(h.gatewayService.ApplyBedrockCCCompat(c.Request.Context(), parsedReq.Body.Bytes(), parsedReq.Model, account, apiKey.GroupID))
body = parsedReq.Body.Bytes()
// 转发请求 - 根据账号平台分流
diff --git a/backend/internal/handler/gateway_helper.go b/backend/internal/handler/gateway_helper.go
index e4897502..52362dd1 100644
--- a/backend/internal/handler/gateway_helper.go
+++ b/backend/internal/handler/gateway_helper.go
@@ -74,8 +74,12 @@ func claudeCodeBodyMapFromParsedRequest(parsedReq *service.ParsedRequest) map[st
bodyMap := map[string]any{
"model": parsedReq.Model,
}
- if parsedReq.System != nil || parsedReq.HasSystem {
- bodyMap["system"] = parsedReq.System
+ if parsedReq.HasSystem {
+ if system, ok := parsedReq.SystemValue(); ok {
+ bodyMap["system"] = system
+ } else {
+ bodyMap["system"] = nil
+ }
}
if parsedReq.MetadataUserID != "" {
bodyMap["metadata"] = map[string]any{"user_id": parsedReq.MetadataUserID}
@@ -87,10 +91,8 @@ func claudeCodeBodyMapFromContextCache(c *gin.Context) map[string]any {
if c == nil {
return nil
}
- if cached, ok := c.Get(service.OpenAIParsedRequestBodyKey); ok {
- if bodyMap, ok := cached.(map[string]any); ok {
- return bodyMap
- }
+ if bodyMap := service.CachedOpenAIParsedRequestBody(c); bodyMap != nil {
+ return bodyMap
}
if cached, ok := c.Get(claudeCodeParsedRequestContextKey); ok {
switch v := cached.(type) {
diff --git a/backend/internal/handler/gateway_helper_hotpath_test.go b/backend/internal/handler/gateway_helper_hotpath_test.go
index d57c396c..d973bedb 100644
--- a/backend/internal/handler/gateway_helper_hotpath_test.go
+++ b/backend/internal/handler/gateway_helper_hotpath_test.go
@@ -185,13 +185,8 @@ func TestSetClaudeCodeClientContext_ReuseParsedRequestAndContextCache(t *testing
c.Request.Header.Set("anthropic-beta", "message-batches-2024-09-24")
c.Request.Header.Set("anthropic-version", "2023-06-01")
- parsedReq := &service.ParsedRequest{
- Model: "claude-3-5-sonnet-20241022",
- System: []any{
- map[string]any{"text": "You are Claude Code, Anthropic's official CLI for Claude."},
- },
- MetadataUserID: "user_aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa_account__session_aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa",
- }
+ parsedReq, err := service.ParseGatewayRequest(service.NewRequestBodyRef(validClaudeCodeBodyJSON()), "")
+ require.NoError(t, err)
// body 非法 JSON,如果函数复用 parsedReq 成功则仍应判定为 Claude Code。
SetClaudeCodeClientContext(c, []byte(`{invalid`), parsedReq)
@@ -204,7 +199,7 @@ func TestSetClaudeCodeClientContext_ReuseParsedRequestAndContextCache(t *testing
c.Request.Header.Set("X-App", "claude-code")
c.Request.Header.Set("anthropic-beta", "message-batches-2024-09-24")
c.Request.Header.Set("anthropic-version", "2023-06-01")
- c.Set(service.OpenAIParsedRequestBodyKey, map[string]any{
+ service.CacheOpenAIParsedRequestBody(c, []byte(`{invalid`), map[string]any{
"model": "claude-3-5-sonnet-20241022",
"system": []any{
map[string]any{"text": "You are Claude Code, Anthropic's official CLI for Claude."},
diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go
index 620c6861..a131cbd9 100644
--- a/backend/internal/handler/openai_gateway_handler.go
+++ b/backend/internal/handler/openai_gateway_handler.go
@@ -954,7 +954,7 @@ func (h *OpenAIGatewayHandler) validateFunctionCallOutputRequest(c *gin.Context,
return true
}
- c.Set(service.OpenAIParsedRequestBodyKey, reqBody)
+ service.CacheOpenAIParsedRequestBody(c, body, reqBody)
validation := service.ValidateFunctionCallOutputContext(reqBody)
if !validation.HasFunctionCallOutput {
return true
diff --git a/backend/internal/service/anthropic_session.go b/backend/internal/service/anthropic_session.go
index 26544c68..bca8cc7f 100644
--- a/backend/internal/service/anthropic_session.go
+++ b/backend/internal/service/anthropic_session.go
@@ -4,6 +4,8 @@ import (
"encoding/json"
"strings"
"time"
+
+ "github.com/tidwall/gjson"
)
// Anthropic 会话 Fallback 相关常量
@@ -30,30 +32,39 @@ func BuildAnthropicDigestChain(parsed *ParsedRequest) string {
var parts []string
- // 1. system prompt
- if parsed.System != nil {
- systemData, _ := json.Marshal(parsed.System)
- if len(systemData) > 0 && string(systemData) != "null" {
- parts = append(parts, "s:"+shortHash(systemData))
- }
+ if systemRaw := parsed.SystemRaw(); len(systemRaw) > 0 && string(systemRaw) != "null" {
+ parts = append(parts, "s:"+shortHash(canonicalAnthropicDigestJSON(systemRaw)))
}
- // 2. messages
- for _, msg := range parsed.Messages {
- msgMap, ok := msg.(map[string]any)
- if !ok {
- continue
- }
- role, _ := msgMap["role"].(string)
- prefix := rolePrefix(role)
- content := msgMap["content"]
- contentData, _ := json.Marshal(content)
- parts = append(parts, prefix+":"+shortHash(contentData))
+ messages := parsed.MessagesRaw()
+ if len(messages) > 0 {
+ gjson.ParseBytes(messages).ForEach(func(_, msg gjson.Result) bool {
+ prefix := rolePrefix(msg.Get("role").String())
+ content := msg.Get("content")
+ parts = append(parts, prefix+":"+shortHash(canonicalAnthropicDigestJSON([]byte(content.Raw))))
+ return true
+ })
}
return strings.Join(parts, "-")
}
+// canonicalAnthropicDigestJSON 保持 digest 对 JSON key 顺序和空白不敏感。
+func canonicalAnthropicDigestJSON(raw []byte) []byte {
+ if len(raw) == 0 {
+ return raw
+ }
+ var value any
+ if err := json.Unmarshal(raw, &value); err != nil {
+ return raw
+ }
+ canonical, err := json.Marshal(value)
+ if err != nil {
+ return raw
+ }
+ return canonical
+}
+
// rolePrefix 将 Anthropic 的 role 映射为单字符前缀
func rolePrefix(role string) string {
switch role {
diff --git a/backend/internal/service/anthropic_session_test.go b/backend/internal/service/anthropic_session_test.go
index 10406643..4d88fc8b 100644
--- a/backend/internal/service/anthropic_session_test.go
+++ b/backend/internal/service/anthropic_session_test.go
@@ -1,3 +1,5 @@
+//go:build unit
+
package service
import (
@@ -5,6 +7,15 @@ import (
"testing"
)
+func mustParseAnthropicDigestRequest(t *testing.T, body string) *ParsedRequest {
+ t.Helper()
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(body)), "")
+ if err != nil {
+ t.Fatalf("ParseGatewayRequest failed: %v", err)
+ }
+ return parsed
+}
+
func TestBuildAnthropicDigestChain_NilRequest(t *testing.T) {
result := BuildAnthropicDigestChain(nil)
if result != "" {
@@ -13,9 +24,7 @@ func TestBuildAnthropicDigestChain_NilRequest(t *testing.T) {
}
func TestBuildAnthropicDigestChain_EmptyMessages(t *testing.T) {
- parsed := &ParsedRequest{
- Messages: []any{},
- }
+ parsed := mustParseAnthropicDigestRequest(t, `{"messages":[]}`)
result := BuildAnthropicDigestChain(parsed)
if result != "" {
t.Errorf("expected empty string for empty messages, got: %s", result)
@@ -23,11 +32,7 @@ func TestBuildAnthropicDigestChain_EmptyMessages(t *testing.T) {
}
func TestBuildAnthropicDigestChain_SingleUserMessage(t *testing.T) {
- parsed := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ parsed := mustParseAnthropicDigestRequest(t, `{"messages":[{"role":"user","content":"hello"}]}`)
result := BuildAnthropicDigestChain(parsed)
parts := splitChain(result)
if len(parts) != 1 {
@@ -39,12 +44,7 @@ func TestBuildAnthropicDigestChain_SingleUserMessage(t *testing.T) {
}
func TestBuildAnthropicDigestChain_UserAndAssistant(t *testing.T) {
- parsed := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "hi there"},
- },
- }
+ parsed := mustParseAnthropicDigestRequest(t, `{"messages":[{"role":"user","content":"hello"},{"role":"assistant","content":"hi there"}]}`)
result := BuildAnthropicDigestChain(parsed)
parts := splitChain(result)
if len(parts) != 2 {
@@ -59,12 +59,7 @@ func TestBuildAnthropicDigestChain_UserAndAssistant(t *testing.T) {
}
func TestBuildAnthropicDigestChain_WithSystemString(t *testing.T) {
- parsed := &ParsedRequest{
- System: "You are a helpful assistant",
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ parsed := mustParseAnthropicDigestRequest(t, `{"system":"You are a helpful assistant","messages":[{"role":"user","content":"hello"}]}`)
result := BuildAnthropicDigestChain(parsed)
parts := splitChain(result)
if len(parts) != 2 {
@@ -79,14 +74,7 @@ func TestBuildAnthropicDigestChain_WithSystemString(t *testing.T) {
}
func TestBuildAnthropicDigestChain_WithSystemContentBlocks(t *testing.T) {
- parsed := &ParsedRequest{
- System: []any{
- map[string]any{"type": "text", "text": "You are a helpful assistant"},
- },
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ parsed := mustParseAnthropicDigestRequest(t, `{"system":[{"type":"text","text":"You are a helpful assistant"}],"messages":[{"role":"user","content":"hello"}]}`)
result := BuildAnthropicDigestChain(parsed)
parts := splitChain(result)
if len(parts) != 2 {
@@ -100,74 +88,33 @@ func TestBuildAnthropicDigestChain_WithSystemContentBlocks(t *testing.T) {
func TestBuildAnthropicDigestChain_ConversationPrefixRelationship(t *testing.T) {
// 核心测试:验证对话增长时链的前缀关系
// 上一轮的完整链一定是下一轮链的前缀
- system := "You are a helpful assistant"
-
- // 第 1 轮: system + user
- round1 := &ParsedRequest{
- System: system,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ round1 := mustParseAnthropicDigestRequest(t, `{"system":"You are a helpful assistant","messages":[{"role":"user","content":"hello"}]}`)
chain1 := BuildAnthropicDigestChain(round1)
- // 第 2 轮: system + user + assistant + user
- round2 := &ParsedRequest{
- System: system,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "hi there"},
- map[string]any{"role": "user", "content": "how are you?"},
- },
- }
+ round2 := mustParseAnthropicDigestRequest(t, `{"system":"You are a helpful assistant","messages":[{"role":"user","content":"hello"},{"role":"assistant","content":"hi there"},{"role":"user","content":"how are you?"}]}`)
chain2 := BuildAnthropicDigestChain(round2)
- // 第 3 轮: system + user + assistant + user + assistant + user
- round3 := &ParsedRequest{
- System: system,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "hi there"},
- map[string]any{"role": "user", "content": "how are you?"},
- map[string]any{"role": "assistant", "content": "I'm doing well"},
- map[string]any{"role": "user", "content": "great"},
- },
- }
+ round3 := mustParseAnthropicDigestRequest(t, `{"system":"You are a helpful assistant","messages":[{"role":"user","content":"hello"},{"role":"assistant","content":"hi there"},{"role":"user","content":"how are you?"},{"role":"assistant","content":"I'm doing well"},{"role":"user","content":"great"}]}`)
chain3 := BuildAnthropicDigestChain(round3)
t.Logf("Chain1: %s", chain1)
t.Logf("Chain2: %s", chain2)
t.Logf("Chain3: %s", chain3)
- // chain1 是 chain2 的前缀
if !strings.HasPrefix(chain2, chain1) {
t.Errorf("chain1 should be prefix of chain2:\n chain1: %s\n chain2: %s", chain1, chain2)
}
-
- // chain2 是 chain3 的前缀
if !strings.HasPrefix(chain3, chain2) {
t.Errorf("chain2 should be prefix of chain3:\n chain2: %s\n chain3: %s", chain2, chain3)
}
-
- // chain1 也是 chain3 的前缀(传递性)
if !strings.HasPrefix(chain3, chain1) {
t.Errorf("chain1 should be prefix of chain3:\n chain1: %s\n chain3: %s", chain1, chain3)
}
}
func TestBuildAnthropicDigestChain_DifferentSystemProducesDifferentChain(t *testing.T) {
- parsed1 := &ParsedRequest{
- System: "System A",
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
- parsed2 := &ParsedRequest{
- System: "System B",
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ parsed1 := mustParseAnthropicDigestRequest(t, `{"system":"System A","messages":[{"role":"user","content":"hello"}]}`)
+ parsed2 := mustParseAnthropicDigestRequest(t, `{"system":"System B","messages":[{"role":"user","content":"hello"}]}`)
chain1 := BuildAnthropicDigestChain(parsed1)
chain2 := BuildAnthropicDigestChain(parsed2)
@@ -176,7 +123,6 @@ func TestBuildAnthropicDigestChain_DifferentSystemProducesDifferentChain(t *test
t.Error("Different system prompts should produce different chains")
}
- // 但 user 部分的 hash 应该相同
parts1 := splitChain(chain1)
parts2 := splitChain(chain2)
if parts1[1] != parts2[1] {
@@ -185,20 +131,8 @@ func TestBuildAnthropicDigestChain_DifferentSystemProducesDifferentChain(t *test
}
func TestBuildAnthropicDigestChain_DifferentContentProducesDifferentChain(t *testing.T) {
- parsed1 := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "ORIGINAL reply"},
- map[string]any{"role": "user", "content": "next"},
- },
- }
- parsed2 := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "TAMPERED reply"},
- map[string]any{"role": "user", "content": "next"},
- },
- }
+ parsed1 := mustParseAnthropicDigestRequest(t, `{"messages":[{"role":"user","content":"hello"},{"role":"assistant","content":"ORIGINAL reply"},{"role":"user","content":"next"}]}`)
+ parsed2 := mustParseAnthropicDigestRequest(t, `{"messages":[{"role":"user","content":"hello"},{"role":"assistant","content":"TAMPERED reply"},{"role":"user","content":"next"}]}`)
chain1 := BuildAnthropicDigestChain(parsed1)
chain2 := BuildAnthropicDigestChain(parsed2)
@@ -209,24 +143,16 @@ func TestBuildAnthropicDigestChain_DifferentContentProducesDifferentChain(t *tes
parts1 := splitChain(chain1)
parts2 := splitChain(chain2)
- // 第一个 user message hash 应该相同
if parts1[0] != parts2[0] {
t.Error("First user message hash should be the same")
}
- // assistant reply hash 应该不同
if parts1[1] == parts2[1] {
t.Error("Assistant reply hash should differ")
}
}
func TestBuildAnthropicDigestChain_Deterministic(t *testing.T) {
- parsed := &ParsedRequest{
- System: "test system",
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "hi"},
- },
- }
+ parsed := mustParseAnthropicDigestRequest(t, `{"system":"test system","messages":[{"role":"user","content":"hello"},{"role":"assistant","content":"hi"}]}`)
chain1 := BuildAnthropicDigestChain(parsed)
chain2 := BuildAnthropicDigestChain(parsed)
@@ -236,6 +162,18 @@ func TestBuildAnthropicDigestChain_Deterministic(t *testing.T) {
}
}
+func TestBuildAnthropicDigestChain_CanonicalJSON(t *testing.T) {
+ parsed1 := mustParseAnthropicDigestRequest(t, `{"system":[{"type":"text","text":"system"}],"messages":[{"role":"user","content":{"type":"text","text":"hello"}}]}`)
+ parsed2 := mustParseAnthropicDigestRequest(t, `{"system":[{"text":"system","type":"text"}],"messages":[{"role":"user","content":{"text":"hello","type":"text"}}]}`)
+
+ chain1 := BuildAnthropicDigestChain(parsed1)
+ chain2 := BuildAnthropicDigestChain(parsed2)
+
+ if chain1 != chain2 {
+ t.Errorf("semantically equivalent JSON should produce same chain: %s vs %s", chain1, chain2)
+ }
+}
+
func TestGenerateAnthropicDigestSessionKey(t *testing.T) {
tests := []struct {
name string
@@ -278,7 +216,6 @@ func TestGenerateAnthropicDigestSessionKey(t *testing.T) {
})
}
- // 验证不同 uuid 产生不同 sessionKey
t.Run("different uuid different key", func(t *testing.T) {
hash := "sameprefix123456"
result1 := GenerateAnthropicDigestSessionKey(hash, "uuid0001-session-a")
@@ -297,18 +234,7 @@ func TestAnthropicSessionTTL(t *testing.T) {
}
func TestBuildAnthropicDigestChain_ContentBlocks(t *testing.T) {
- // 测试 content 为 content blocks 数组的情况
- parsed := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "content": []any{
- map[string]any{"type": "text", "text": "describe this image"},
- map[string]any{"type": "image", "source": map[string]any{"type": "base64"}},
- },
- },
- },
- }
+ parsed := mustParseAnthropicDigestRequest(t, `{"messages":[{"role":"user","content":[{"type":"text","text":"describe this image"},{"type":"image","source":{"type":"base64"}}]}]}`)
result := BuildAnthropicDigestChain(parsed)
parts := splitChain(result)
if len(parts) != 1 {
diff --git a/backend/internal/service/gateway_oauth_metadata_test.go b/backend/internal/service/gateway_oauth_metadata_test.go
index ed6f1887..b172dc6e 100644
--- a/backend/internal/service/gateway_oauth_metadata_test.go
+++ b/backend/internal/service/gateway_oauth_metadata_test.go
@@ -14,8 +14,6 @@ func TestBuildOAuthMetadataUserID_FallbackWithoutAccountUUID(t *testing.T) {
Model: "claude-sonnet-4-5",
Stream: true,
MetadataUserID: "",
- System: nil,
- Messages: nil,
}
account := &Account{
diff --git a/backend/internal/service/gateway_request.go b/backend/internal/service/gateway_request.go
index 819bb0a8..65718c27 100644
--- a/backend/internal/service/gateway_request.go
+++ b/backend/internal/service/gateway_request.go
@@ -51,6 +51,12 @@ type SessionContext struct {
APIKeyID int64
}
+type jsonRange struct {
+ start int // 原始请求体中的起始偏移(闭区间)
+ end int // 原始请求体中的结束偏移(开区间)
+ kind gjson.Type // JSON 值类型,用于调用方做轻量分支
+}
+
type RequestBodyRef struct {
data []byte
}
@@ -80,6 +86,76 @@ func (b *RequestBodyRef) Replace(data []byte) {
b.data = data
}
+func missingJSONRange() jsonRange {
+ return jsonRange{start: -1, end: -1}
+}
+
+func rangeFromResult(r gjson.Result) jsonRange {
+ if r.Raw == "" || r.Index <= 0 {
+ return missingJSONRange()
+ }
+ end := r.Index + len(r.Raw)
+ if end < r.Index {
+ return missingJSONRange()
+ }
+ return jsonRange{start: r.Index, end: end, kind: r.Type}
+}
+
+func (r jsonRange) exists() bool {
+ return r.start >= 0 && r.end >= r.start
+}
+
+func clearGatewayRequestRanges(parsed *ParsedRequest) {
+ if parsed == nil {
+ return
+ }
+ parsed.HasSystem = false
+ parsed.systemRange = missingJSONRange()
+ parsed.messagesRange = missingJSONRange()
+}
+
+func setGatewayRequestRanges(parsed *ParsedRequest, protocol string, jsonStr string) {
+ if parsed == nil {
+ return
+ }
+ switch protocol {
+ case domain.PlatformGemini:
+ if sysParts := gjson.Get(jsonStr, "systemInstruction.parts"); sysParts.Exists() && sysParts.IsArray() {
+ parsed.systemRange = rangeFromResult(sysParts)
+ }
+ if contents := gjson.Get(jsonStr, "contents"); contents.Exists() && contents.IsArray() {
+ parsed.messagesRange = rangeFromResult(contents)
+ }
+ default:
+ if sys := gjson.Get(jsonStr, "system"); sys.Exists() {
+ parsed.HasSystem = true
+ parsed.systemRange = rangeFromResult(sys)
+ }
+ if msgs := gjson.Get(jsonStr, "messages"); msgs.Exists() && msgs.IsArray() {
+ parsed.messagesRange = rangeFromResult(msgs)
+ }
+ }
+}
+
+func refreshGatewayRequestRanges(parsed *ParsedRequest, protocol string) error {
+ if parsed == nil {
+ return fmt.Errorf("empty request body")
+ }
+ clearGatewayRequestRanges(parsed)
+ if parsed.Body == nil {
+ return fmt.Errorf("empty request body")
+ }
+
+ bodyBytes := parsed.Body.Bytes()
+ if !gjson.ValidBytes(bodyBytes) {
+ return fmt.Errorf("invalid json")
+ }
+
+ jsonStr := *(*string)(unsafe.Pointer(&bodyBytes))
+ setGatewayRequestRanges(parsed, protocol, jsonStr)
+ return nil
+}
+
// ParsedRequest 保存网关请求的预解析结果
//
// 性能优化说明:
@@ -93,18 +169,20 @@ func (b *RequestBodyRef) Replace(data []byte) {
// 2. 将解析结果 ParsedRequest 传递给 Service 层
// 3. 避免重复 json.Unmarshal,减少 CPU 和内存开销
type ParsedRequest struct {
- Body *RequestBodyRef // 原始请求体引用(保留用于转发)
+ Body *RequestBodyRef // 原始请求体引用(保留用于转发);替换内容请走 ReplaceBody
Model string // 请求的模型名称
Stream bool // 是否为流式请求
MetadataUserID string // metadata.user_id(用于会话亲和)
- System any // system 字段内容
- Messages []any // messages 数组
HasSystem bool // 是否包含 system 字段(包含 null 也视为显式传入)
ThinkingEnabled bool // 是否开启 thinking(部分平台会影响最终模型名)
OutputEffort string // output_config.effort(Claude API 的推理强度控制)
MaxTokens int // max_tokens 值(用于探测请求拦截)
SessionContext *SessionContext // 可选:请求上下文区分因子(nil 时行为不变)
+ protocol string // 当前 Body 的协议格式,用于 Body 替换后刷新 raw range
+ systemRange jsonRange // system/systemInstruction.parts 的 raw JSON 范围,绑定 Body 当前内容
+ messagesRange jsonRange // messages/contents 的 raw JSON 范围,绑定 Body 当前内容
+
// GroupID 请求所属分组 ID(来自 API Key)
GroupID *int64
@@ -173,7 +251,10 @@ func ParseGatewayRequest(body *RequestBodyRef, protocol string) (*ParsedRequest,
jsonStr := *(*string)(unsafe.Pointer(&bodyBytes))
parsed := &ParsedRequest{
- Body: body,
+ Body: body,
+ protocol: protocol,
+ systemRange: missingJSONRange(),
+ messagesRange: missingJSONRange(),
}
// --- gjson 提取简单字段(避免完整 Unmarshal) ---
@@ -219,60 +300,73 @@ func ParseGatewayRequest(body *RequestBodyRef, protocol string) (*ParsedRequest,
}
// --- system/messages 提取 ---
- // 避免把整个 body Unmarshal 到 map(会产生大量 map/接口分配)。
- // 使用 gjson 抽取目标字段的 Raw,再对该子树进行 Unmarshal。
-
- switch protocol {
- case domain.PlatformGemini:
- // Gemini 原生格式: systemInstruction.parts / contents
- if sysParts := gjson.Get(jsonStr, "systemInstruction.parts"); sysParts.Exists() && sysParts.IsArray() {
- var parts []any
- if err := json.Unmarshal(sliceRawFromBody(bodyBytes, sysParts), &parts); err != nil {
- return nil, err
- }
- parsed.System = parts
- }
-
- if contents := gjson.Get(jsonStr, "contents"); contents.Exists() && contents.IsArray() {
- var msgs []any
- if err := json.Unmarshal(sliceRawFromBody(bodyBytes, contents), &msgs); err != nil {
- return nil, err
- }
- parsed.Messages = msgs
- }
- default:
- // Anthropic / OpenAI 格式: system / messages
- // system 字段只要存在就视为显式提供(即使为 null),
- // 以避免客户端传 null 时被默认 system 误注入。
- if sys := gjson.Get(jsonStr, "system"); sys.Exists() {
- parsed.HasSystem = true
- switch sys.Type {
- case gjson.Null:
- parsed.System = nil
- case gjson.String:
- // 与 encoding/json 的 Unmarshal 行为一致:返回解码后的字符串。
- parsed.System = sys.String()
- default:
- var system any
- if err := json.Unmarshal(sliceRawFromBody(bodyBytes, sys), &system); err != nil {
- return nil, err
- }
- parsed.System = system
- }
- }
-
- if msgs := gjson.Get(jsonStr, "messages"); msgs.Exists() && msgs.IsArray() {
- var messages []any
- if err := json.Unmarshal(sliceRawFromBody(bodyBytes, msgs), &messages); err != nil {
- return nil, err
- }
- parsed.Messages = messages
- }
- }
+ // 只保存大字段 raw range,不默认反序列化成 []any/map[string]any 对象图。
+ setGatewayRequestRanges(parsed, protocol, jsonStr)
return parsed, nil
}
+func (p *ParsedRequest) raw(r jsonRange) []byte {
+ if p == nil || p.Body == nil || !r.exists() {
+ return nil
+ }
+ body := p.Body.Bytes()
+ if r.end > len(body) {
+ return nil
+ }
+ return body[r.start:r.end]
+}
+
+func (p *ParsedRequest) SystemRaw() []byte {
+ return p.raw(p.systemRange)
+}
+
+func (p *ParsedRequest) MessagesRaw() []byte {
+ return p.raw(p.messagesRange)
+}
+
+func (p *ParsedRequest) DecodeSystem(dst any) error {
+ raw := p.SystemRaw()
+ if len(raw) == 0 {
+ return nil
+ }
+ return json.Unmarshal(raw, dst)
+}
+
+func (p *ParsedRequest) DecodeMessages(dst any) error {
+ raw := p.MessagesRaw()
+ if len(raw) == 0 {
+ return nil
+ }
+ return json.Unmarshal(raw, dst)
+}
+
+func (p *ParsedRequest) SystemValue() (any, bool) {
+ raw := p.SystemRaw()
+ if len(raw) == 0 {
+ return nil, false
+ }
+ var system any
+ if err := json.Unmarshal(raw, &system); err != nil {
+ return nil, false
+ }
+ return system, true
+}
+
+func (p *ParsedRequest) ReplaceBody(data []byte) {
+ if p == nil {
+ return
+ }
+ if p.Body == nil {
+ p.Body = NewRequestBodyRef(data)
+ } else {
+ p.Body.Replace(data)
+ }
+ if err := refreshGatewayRequestRanges(p, p.protocol); err != nil {
+ clearGatewayRequestRanges(p)
+ }
+}
+
// sliceRawFromBody 返回 Result.Raw 对应的原始字节切片。
// 优先使用 Result.Index 直接从 body 切片,避免对大字段(如 messages)产生额外拷贝。
// 当 Index 不可用时,退化为复制(理论上极少发生)。
diff --git a/backend/internal/service/gateway_request_test.go b/backend/internal/service/gateway_request_test.go
index d415b871..288c031c 100644
--- a/backend/internal/service/gateway_request_test.go
+++ b/backend/internal/service/gateway_request_test.go
@@ -10,6 +10,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/domain"
"github.com/stretchr/testify/require"
+ "github.com/tidwall/gjson"
)
func TestParseGatewayRequest(t *testing.T) {
@@ -20,8 +21,8 @@ func TestParseGatewayRequest(t *testing.T) {
require.True(t, parsed.Stream)
require.Equal(t, "session_123e4567-e89b-12d3-a456-426614174000", parsed.MetadataUserID)
require.True(t, parsed.HasSystem)
- require.NotNil(t, parsed.System)
- require.Len(t, parsed.Messages, 1)
+ require.NotEmpty(t, parsed.SystemRaw())
+ require.NotEmpty(t, parsed.MessagesRaw())
require.False(t, parsed.ThinkingEnabled)
}
@@ -61,7 +62,7 @@ func TestParseGatewayRequest_SystemNull(t *testing.T) {
require.NoError(t, err)
// 显式传入 system:null 也应视为“字段已存在”,避免默认 system 被注入。
require.True(t, parsed.HasSystem)
- require.Nil(t, parsed.System)
+ require.Equal(t, []byte("null"), parsed.SystemRaw())
}
func TestParseGatewayRequest_InvalidModelType(t *testing.T) {
@@ -88,9 +89,9 @@ func TestParseGatewayRequest_GeminiContents(t *testing.T) {
}`)
parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err)
- require.Len(t, parsed.Messages, 3, "should parse contents as Messages")
+ require.Len(t, gjson.ParseBytes(parsed.MessagesRaw()).Array(), 3, "should parse contents as Messages")
require.False(t, parsed.HasSystem, "Gemini format should not set HasSystem")
- require.Nil(t, parsed.System, "no systemInstruction means nil System")
+ require.Nil(t, parsed.SystemRaw(), "no systemInstruction means nil System")
}
func TestParseGatewayRequest_GeminiSystemInstruction(t *testing.T) {
@@ -104,14 +105,11 @@ func TestParseGatewayRequest_GeminiSystemInstruction(t *testing.T) {
}`)
parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err)
- require.NotNil(t, parsed.System, "should parse systemInstruction.parts as System")
- parts, ok := parsed.System.([]any)
- require.True(t, ok)
- require.Len(t, parts, 1)
- partMap, ok := parts[0].(map[string]any)
- require.True(t, ok)
- require.Equal(t, "You are a helpful assistant.", partMap["text"])
- require.Len(t, parsed.Messages, 1)
+ system := gjson.ParseBytes(parsed.SystemRaw())
+ require.True(t, system.IsArray(), "should parse systemInstruction.parts as System")
+ require.Len(t, system.Array(), 1)
+ require.Equal(t, "You are a helpful assistant.", system.Get("0.text").String())
+ require.Len(t, gjson.ParseBytes(parsed.MessagesRaw()).Array(), 1)
}
func TestParseGatewayRequest_GeminiWithModel(t *testing.T) {
@@ -122,7 +120,7 @@ func TestParseGatewayRequest_GeminiWithModel(t *testing.T) {
parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err)
require.Equal(t, "gemini-2.5-pro", parsed.Model)
- require.Len(t, parsed.Messages, 1)
+ require.Len(t, gjson.ParseBytes(parsed.MessagesRaw()).Array(), 1)
}
func TestParseGatewayRequest_GeminiIgnoresAnthropicFields(t *testing.T) {
@@ -135,22 +133,22 @@ func TestParseGatewayRequest_GeminiIgnoresAnthropicFields(t *testing.T) {
parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err)
require.False(t, parsed.HasSystem, "Gemini protocol should not parse Anthropic system field")
- require.Nil(t, parsed.System, "no systemInstruction = nil System")
- require.Len(t, parsed.Messages, 1, "should use contents, not messages")
+ require.Nil(t, parsed.SystemRaw(), "no systemInstruction = nil System")
+ require.Len(t, gjson.ParseBytes(parsed.MessagesRaw()).Array(), 1, "should use contents, not messages")
}
func TestParseGatewayRequest_GeminiEmptyContents(t *testing.T) {
body := []byte(`{"contents": []}`)
parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err)
- require.Empty(t, parsed.Messages)
+ require.Empty(t, gjson.ParseBytes(parsed.MessagesRaw()).Array())
}
func TestParseGatewayRequest_GeminiNoContents(t *testing.T) {
body := []byte(`{"model": "gemini-2.5-flash"}`)
parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err)
- require.Nil(t, parsed.Messages)
+ require.Nil(t, parsed.MessagesRaw())
require.Equal(t, "gemini-2.5-flash", parsed.Model)
}
@@ -165,11 +163,10 @@ func TestParseGatewayRequest_AnthropicIgnoresGeminiFields(t *testing.T) {
parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformAnthropic)
require.NoError(t, err)
require.True(t, parsed.HasSystem)
- require.Equal(t, "real system", parsed.System)
- require.Len(t, parsed.Messages, 1)
- msg, ok := parsed.Messages[0].(map[string]any)
- require.True(t, ok)
- require.Equal(t, "real content", msg["content"])
+ require.Equal(t, "real system", gjson.ParseBytes(parsed.SystemRaw()).String())
+ messages := gjson.ParseBytes(parsed.MessagesRaw()).Array()
+ require.Len(t, messages, 1)
+ require.Equal(t, "real content", messages[0].Get("content").String())
}
func TestFilterThinkingBlocks(t *testing.T) {
@@ -970,10 +967,10 @@ func TestParseGatewayRequest_OptionalFieldsMissing(t *testing.T) {
require.Equal(t, tt.wantMaxTokens, parsed.MaxTokens)
if tt.wantMessagesNil {
- require.Nil(t, parsed.Messages)
+ require.Nil(t, parsed.MessagesRaw())
}
if tt.wantMessagesLen > 0 {
- require.Len(t, parsed.Messages, tt.wantMessagesLen)
+ require.Len(t, gjson.ParseBytes(parsed.MessagesRaw()).Array(), tt.wantMessagesLen)
}
})
}
@@ -1087,25 +1084,8 @@ func parseGatewayRequestOld(body []byte, protocol string) (*ParsedRequest, error
}
}
- // system / messages(按协议分支)
- switch protocol {
- case domain.PlatformGemini:
- if sysInst, ok := req["systemInstruction"].(map[string]any); ok {
- if parts, ok := sysInst["parts"].([]any); ok {
- parsed.System = parts
- }
- }
- if contents, ok := req["contents"].([]any); ok {
- parsed.Messages = contents
- }
- default:
- if system, ok := req["system"]; ok {
- parsed.HasSystem = true
- parsed.System = system
- }
- if messages, ok := req["messages"].([]any); ok {
- parsed.Messages = messages
- }
+ if err := refreshGatewayRequestRanges(parsed, protocol); err != nil {
+ return nil, err
}
return parsed, nil
diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go
index 8f55bf13..438592e3 100644
--- a/backend/internal/service/gateway_service.go
+++ b/backend/internal/service/gateway_service.go
@@ -748,31 +748,10 @@ func (s *GatewayService) GenerateSessionHash(parsed *ParsedRequest) string {
_, _ = combined.WriteString(strconv.FormatInt(parsed.SessionContext.APIKeyID, 10))
_, _ = combined.WriteString("|")
}
- if parsed.System != nil {
- systemText := s.extractTextFromSystem(parsed.System)
- if systemText != "" {
- _, _ = combined.WriteString(systemText)
- }
- }
- for _, msg := range parsed.Messages {
- if m, ok := msg.(map[string]any); ok {
- if content, exists := m["content"]; exists {
- // Anthropic: messages[].content
- if msgText := s.extractTextFromContent(content); msgText != "" {
- _, _ = combined.WriteString(msgText)
- }
- } else if parts, ok := m["parts"].([]any); ok {
- // Gemini: contents[].parts[].text
- for _, part := range parts {
- if partMap, ok := part.(map[string]any); ok {
- if text, ok := partMap["text"].(string); ok {
- _, _ = combined.WriteString(text)
- }
- }
- }
- }
- }
+ if systemText := extractTextFromSystemRaw(parsed.SystemRaw()); systemText != "" {
+ _, _ = combined.WriteString(systemText)
}
+ appendMessageTextsFromRaw(&combined, parsed.MessagesRaw())
if combined.Len() > 0 {
hash := s.hashContent(combined.String())
slog.Info("sticky.hash_source",
@@ -847,82 +826,127 @@ func (s *GatewayService) extractCacheableContent(parsed *ParsedRequest) string {
return ""
}
- var builder strings.Builder
-
- // 检查 system 中的 cacheable 内容
- if system, ok := parsed.System.([]any); ok {
- for _, part := range system {
- if partMap, ok := part.(map[string]any); ok {
- if cc, ok := partMap["cache_control"].(map[string]any); ok {
- if cc["type"] == "ephemeral" {
- if text, ok := partMap["text"].(string); ok {
- _, _ = builder.WriteString(text)
- }
- }
- }
- }
- }
+ systemText := extractCacheableTextFromSystemRaw(parsed.SystemRaw())
+ if messageText := extractCacheableTextFromMessagesRaw(parsed.MessagesRaw()); messageText != "" {
+ return messageText
}
- systemText := builder.String()
-
- // 检查 messages 中的 cacheable 内容
- for _, msg := range parsed.Messages {
- if msgMap, ok := msg.(map[string]any); ok {
- if msgContent, ok := msgMap["content"].([]any); ok {
- for _, part := range msgContent {
- if partMap, ok := part.(map[string]any); ok {
- if cc, ok := partMap["cache_control"].(map[string]any); ok {
- if cc["type"] == "ephemeral" {
- return s.extractTextFromContent(msgMap["content"])
- }
- }
- }
- }
- }
- }
- }
-
return systemText
}
-func (s *GatewayService) extractTextFromSystem(system any) string {
- switch v := system.(type) {
- case string:
- return v
- case []any:
- var texts []string
- for _, part := range v {
- if partMap, ok := part.(map[string]any); ok {
- if text, ok := partMap["text"].(string); ok {
- texts = append(texts, text)
- }
- }
+func extractTextFromSystemRaw(raw []byte) string {
+ system := gjson.ParseBytes(raw)
+ switch system.Type {
+ case gjson.String:
+ return system.String()
+ case gjson.JSON:
+ if !system.IsArray() {
+ return ""
}
- return strings.Join(texts, "")
+ var builder strings.Builder
+ system.ForEach(func(_, part gjson.Result) bool {
+ if text := part.Get("text").String(); text != "" {
+ _, _ = builder.WriteString(text)
+ }
+ return true
+ })
+ return builder.String()
}
return ""
}
-func (s *GatewayService) extractTextFromContent(content any) string {
- switch v := content.(type) {
- case string:
- return v
- case []any:
- var texts []string
- for _, part := range v {
- if partMap, ok := part.(map[string]any); ok {
- if partMap["type"] == "text" {
- if text, ok := partMap["text"].(string); ok {
- texts = append(texts, text)
- }
+func extractTextFromContentRaw(content gjson.Result) string {
+ switch content.Type {
+ case gjson.String:
+ return content.String()
+ case gjson.JSON:
+ if !content.IsArray() {
+ return ""
+ }
+ var builder strings.Builder
+ content.ForEach(func(_, part gjson.Result) bool {
+ if part.Get("type").String() == "text" {
+ if text := part.Get("text").String(); text != "" {
+ _, _ = builder.WriteString(text)
}
}
- }
- return strings.Join(texts, "")
+ return true
+ })
+ return builder.String()
}
return ""
}
+func appendMessageTextsFromRaw(builder *strings.Builder, raw []byte) {
+ if builder == nil || len(raw) == 0 {
+ return
+ }
+ messages := gjson.ParseBytes(raw)
+ if !messages.IsArray() {
+ return
+ }
+ messages.ForEach(func(_, msg gjson.Result) bool {
+ if content := msg.Get("content"); content.Exists() {
+ _, _ = builder.WriteString(extractTextFromContentRaw(content))
+ return true
+ }
+ parts := msg.Get("parts")
+ if parts.IsArray() {
+ parts.ForEach(func(_, part gjson.Result) bool {
+ if text := part.Get("text").String(); text != "" {
+ _, _ = builder.WriteString(text)
+ }
+ return true
+ })
+ }
+ return true
+ })
+}
+
+func extractCacheableTextFromSystemRaw(raw []byte) string {
+ system := gjson.ParseBytes(raw)
+ if !system.IsArray() {
+ return ""
+ }
+ var builder strings.Builder
+ system.ForEach(func(_, part gjson.Result) bool {
+ if part.Get("cache_control.type").String() == "ephemeral" {
+ if text := part.Get("text").String(); text != "" {
+ _, _ = builder.WriteString(text)
+ }
+ }
+ return true
+ })
+ return builder.String()
+}
+
+func extractCacheableTextFromMessagesRaw(raw []byte) string {
+ messages := gjson.ParseBytes(raw)
+ if !messages.IsArray() {
+ return ""
+ }
+ var text string
+ messages.ForEach(func(_, msg gjson.Result) bool {
+ content := msg.Get("content")
+ if !content.IsArray() {
+ return true
+ }
+ found := false
+ content.ForEach(func(_, part gjson.Result) bool {
+ if part.Get("cache_control.type").String() == "ephemeral" {
+ found = true
+ return false
+ }
+ return true
+ })
+ if found {
+ text = extractTextFromContentRaw(content)
+ return false
+ }
+ return true
+ })
+ return text
+}
+
func (s *GatewayService) hashContent(content string) string {
h := xxhash.Sum64String(content)
return strconv.FormatUint(h, 36)
@@ -1284,7 +1308,7 @@ func (s *GatewayService) applyClaudeCodeOAuthMimicryToBody(
systemRewritten := false
if !strings.Contains(strings.ToLower(model), "haiku") {
- body = rewriteSystemForNonClaudeCode(body, systemRaw)
+ body = rewriteSystemForNonClaudeCode(body, normalizeSystemParam(systemRaw))
systemRewritten = true
}
@@ -4474,7 +4498,8 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
// Parrot 的 transform_request 从不检查客户端 system 内容,直接覆盖。
systemRewritten := false
if !strings.Contains(strings.ToLower(reqModel), "haiku") {
- body = rewriteSystemForNonClaudeCode(body, parsed.System)
+ systemRaw, _ := parsed.SystemValue()
+ body = rewriteSystemForNonClaudeCode(body, systemRaw)
systemRewritten = true
}
diff --git a/backend/internal/service/gateway_service_benchmark_test.go b/backend/internal/service/gateway_service_benchmark_test.go
index 5637680b..8b30cb24 100644
--- a/backend/internal/service/gateway_service_benchmark_test.go
+++ b/backend/internal/service/gateway_service_benchmark_test.go
@@ -2,6 +2,7 @@ package service
import (
"strconv"
+ "strings"
"testing"
)
@@ -34,17 +35,20 @@ func BenchmarkExtractCacheableContent_System(b *testing.B) {
}
func buildSystemCacheableRequest(parts int) *ParsedRequest {
- systemParts := make([]any, 0, parts)
+ var builder strings.Builder
+ builder.WriteString(`{"system":[`)
for i := 0; i < parts; i++ {
- systemParts = append(systemParts, map[string]any{
- "text": "system_part_" + strconv.Itoa(i),
- "cache_control": map[string]any{
- "type": "ephemeral",
- },
- })
+ if i > 0 {
+ builder.WriteByte(',')
+ }
+ builder.WriteString(`{"text":"system_part_`)
+ builder.WriteString(strconv.Itoa(i))
+ builder.WriteString(`","cache_control":{"type":"ephemeral"}}`)
}
- return &ParsedRequest{
- System: systemParts,
- HasSystem: true,
+ builder.WriteString(`]}`)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(builder.String())), "")
+ if err != nil {
+ panic(err)
}
+ return parsed
}
diff --git a/backend/internal/service/generate_session_hash_test.go b/backend/internal/service/generate_session_hash_test.go
index 8f3258b7..5ed3f0ae 100644
--- a/backend/internal/service/generate_session_hash_test.go
+++ b/backend/internal/service/generate_session_hash_test.go
@@ -3,12 +3,67 @@
package service
import (
+ "encoding/json"
"testing"
+ "github.com/Wei-Shaw/sub2api/internal/domain"
"github.com/stretchr/testify/require"
)
-// ============ 基础优先级测试 ============
+func mustParseSessionHashRequest(t *testing.T, body string, ctx *SessionContext) *ParsedRequest {
+ t.Helper()
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(body)), domain.PlatformAnthropic)
+ require.NoError(t, err)
+ parsed.SessionContext = ctx
+ return parsed
+}
+
+func mustParseGeminiSessionHashRequest(t *testing.T, body string, ctx *SessionContext) *ParsedRequest {
+ t.Helper()
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(body)), domain.PlatformGemini)
+ require.NoError(t, err)
+ parsed.SessionContext = ctx
+ return parsed
+}
+
+func anthropicSessionBody(system any, messages []any, metadataUserID string) string {
+ body := map[string]any{}
+ if system != nil {
+ body["system"] = system
+ }
+ if messages != nil {
+ body["messages"] = messages
+ }
+ if metadataUserID != "" {
+ body["metadata"] = map[string]any{"user_id": metadataUserID}
+ }
+ data, _ := json.Marshal(body)
+ return string(data)
+}
+
+func geminiSessionBody(systemParts []any, contents []any) string {
+ body := map[string]any{}
+ if systemParts != nil {
+ body["systemInstruction"] = map[string]any{"parts": systemParts}
+ }
+ if contents != nil {
+ body["contents"] = contents
+ }
+ data, _ := json.Marshal(body)
+ return string(data)
+}
+
+func msg(role string, content any) map[string]any {
+ return map[string]any{"role": role, "content": content}
+}
+
+func geminiMsg(role string, texts ...string) map[string]any {
+ parts := make([]any, 0, len(texts))
+ for _, text := range texts {
+ parts = append(parts, map[string]any{"text": text})
+ }
+ return map[string]any{"role": role, "parts": parts}
+}
func TestGenerateSessionHash_NilParsedRequest(t *testing.T) {
svc := &GatewayService{}
@@ -22,37 +77,17 @@ func TestGenerateSessionHash_EmptyRequest(t *testing.T) {
func TestGenerateSessionHash_MetadataHasHighestPriority(t *testing.T) {
svc := &GatewayService{}
-
- parsed := &ParsedRequest{
- MetadataUserID: "user_a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2_account__session_123e4567-e89b-12d3-a456-426614174000",
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ metadata := "user_a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2_account__session_123e4567-e89b-12d3-a456-426614174000"
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "hello")}, metadata), nil)
hash := svc.GenerateSessionHash(parsed)
require.Equal(t, "123e4567-e89b-12d3-a456-426614174000", hash, "metadata session_id should have highest priority")
}
-// ============ System + Messages 基础测试 ============
-
func TestGenerateSessionHash_SystemPlusMessages(t *testing.T) {
svc := &GatewayService{}
-
- withSystem := &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
- withoutSystem := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ withSystem := mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "hello")}, ""), nil)
+ withoutSystem := mustParseSessionHashRequest(t, anthropicSessionBody(nil, []any{msg("user", "hello")}, ""), nil)
h1 := svc.GenerateSessionHash(withSystem)
h2 := svc.GenerateSessionHash(withoutSystem)
@@ -63,32 +98,16 @@ func TestGenerateSessionHash_SystemPlusMessages(t *testing.T) {
func TestGenerateSessionHash_SystemOnlyProducesHash(t *testing.T) {
svc := &GatewayService{}
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", nil, ""), nil)
- parsed := &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- }
hash := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, hash, "system prompt alone should produce a hash as part of full digest")
}
func TestGenerateSessionHash_DifferentSystemsSameMessages(t *testing.T) {
svc := &GatewayService{}
-
- parsed1 := &ParsedRequest{
- System: "You are assistant A.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
- parsed2 := &ParsedRequest{
- System: "You are assistant B.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ parsed1 := mustParseSessionHashRequest(t, anthropicSessionBody("You are assistant A.", []any{msg("user", "hello")}, ""), nil)
+ parsed2 := mustParseSessionHashRequest(t, anthropicSessionBody("You are assistant B.", []any{msg("user", "hello")}, ""), nil)
h1 := svc.GenerateSessionHash(parsed1)
h2 := svc.GenerateSessionHash(parsed2)
@@ -97,16 +116,8 @@ func TestGenerateSessionHash_DifferentSystemsSameMessages(t *testing.T) {
func TestGenerateSessionHash_SameSystemSameMessages(t *testing.T) {
svc := &GatewayService{}
-
mk := func() *ParsedRequest {
- return &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "hi"},
- },
- }
+ return mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "hello"), msg("assistant", "hi")}, ""), nil)
}
h1 := svc.GenerateSessionHash(mk())
@@ -116,53 +127,19 @@ func TestGenerateSessionHash_SameSystemSameMessages(t *testing.T) {
func TestGenerateSessionHash_DifferentMessagesProduceDifferentHash(t *testing.T) {
svc := &GatewayService{}
-
- parsed1 := &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "help me with Go"},
- },
- }
- parsed2 := &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "help me with Python"},
- },
- }
+ parsed1 := mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "help me with Go")}, ""), nil)
+ parsed2 := mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "help me with Python")}, ""), nil)
h1 := svc.GenerateSessionHash(parsed1)
h2 := svc.GenerateSessionHash(parsed2)
require.NotEqual(t, h1, h2, "same system but different messages should produce different hashes")
}
-// ============ SessionContext 核心测试 ============
-
func TestGenerateSessionHash_DifferentSessionContextProducesDifferentHash(t *testing.T) {
svc := &GatewayService{}
-
- // 相同消息 + 不同 SessionContext → 不同 hash(解决碰撞问题的核心场景)
- parsed1 := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: &SessionContext{
- ClientIP: "192.168.1.1",
- UserAgent: "Mozilla/5.0",
- APIKeyID: 100,
- },
- }
- parsed2 := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: &SessionContext{
- ClientIP: "10.0.0.1",
- UserAgent: "curl/7.0",
- APIKeyID: 200,
- },
- }
+ body := anthropicSessionBody(nil, []any{msg("user", "hello")}, "")
+ parsed1 := mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "192.168.1.1", UserAgent: "Mozilla/5.0", APIKeyID: 100})
+ parsed2 := mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "10.0.0.1", UserAgent: "curl/7.0", APIKeyID: 200})
h1 := svc.GenerateSessionHash(parsed1)
h2 := svc.GenerateSessionHash(parsed2)
@@ -173,19 +150,9 @@ func TestGenerateSessionHash_DifferentSessionContextProducesDifferentHash(t *tes
func TestGenerateSessionHash_SameSessionContextProducesSameHash(t *testing.T) {
svc := &GatewayService{}
-
- mk := func() *ParsedRequest {
- return &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: &SessionContext{
- ClientIP: "192.168.1.1",
- UserAgent: "Mozilla/5.0",
- APIKeyID: 100,
- },
- }
- }
+ ctx := &SessionContext{ClientIP: "192.168.1.1", UserAgent: "Mozilla/5.0", APIKeyID: 100}
+ body := anthropicSessionBody(nil, []any{msg("user", "hello")}, "")
+ mk := func() *ParsedRequest { return mustParseSessionHashRequest(t, body, ctx) }
h1 := svc.GenerateSessionHash(mk())
h2 := svc.GenerateSessionHash(mk())
@@ -194,35 +161,17 @@ func TestGenerateSessionHash_SameSessionContextProducesSameHash(t *testing.T) {
func TestGenerateSessionHash_MetadataOverridesSessionContext(t *testing.T) {
svc := &GatewayService{}
-
- parsed := &ParsedRequest{
- MetadataUserID: "user_a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2_account__session_123e4567-e89b-12d3-a456-426614174000",
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: &SessionContext{
- ClientIP: "192.168.1.1",
- UserAgent: "Mozilla/5.0",
- APIKeyID: 100,
- },
- }
+ metadata := "user_a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2_account__session_123e4567-e89b-12d3-a456-426614174000"
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody(nil, []any{msg("user", "hello")}, metadata), &SessionContext{ClientIP: "192.168.1.1", UserAgent: "Mozilla/5.0", APIKeyID: 100})
hash := svc.GenerateSessionHash(parsed)
- require.Equal(t, "123e4567-e89b-12d3-a456-426614174000", hash,
- "metadata session_id should take priority over SessionContext")
+ require.Equal(t, "123e4567-e89b-12d3-a456-426614174000", hash, "metadata session_id should take priority over SessionContext")
}
func TestGenerateSessionHash_MetadataJSON_HasHighestPriority(t *testing.T) {
svc := &GatewayService{}
-
- parsed := &ParsedRequest{
- MetadataUserID: `{"device_id":"a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2","account_uuid":"","session_id":"c72554f2-1234-5678-abcd-123456789abc"}`,
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ metadata := `{"device_id":"a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2","account_uuid":"","session_id":"c72554f2-1234-5678-abcd-123456789abc"}`
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "hello")}, metadata), nil)
hash := svc.GenerateSessionHash(parsed)
require.Equal(t, "c72554f2-1234-5678-abcd-123456789abc", hash, "JSON format metadata session_id should have highest priority")
@@ -230,69 +179,25 @@ func TestGenerateSessionHash_MetadataJSON_HasHighestPriority(t *testing.T) {
func TestGenerateSessionHash_NilSessionContextBackwardCompatible(t *testing.T) {
svc := &GatewayService{}
-
- withCtx := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: nil,
- }
- withoutCtx := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ body := anthropicSessionBody(nil, []any{msg("user", "hello")}, "")
+ withCtx := mustParseSessionHashRequest(t, body, nil)
+ withoutCtx := mustParseSessionHashRequest(t, body, nil)
h1 := svc.GenerateSessionHash(withCtx)
h2 := svc.GenerateSessionHash(withoutCtx)
require.Equal(t, h1, h2, "nil SessionContext should produce same hash as no SessionContext")
}
-// ============ 多轮连续会话测试 ============
-
func TestGenerateSessionHash_ContinuousConversation_HashChangesWithMessages(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1}
-
- // 模拟连续会话:每增加一轮对话,hash 应该不同(内容累积变化)
- round1 := &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: ctx,
- }
-
- round2 := &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "Hi there!"},
- map[string]any{"role": "user", "content": "How are you?"},
- },
- SessionContext: ctx,
- }
-
- round3 := &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "Hi there!"},
- map[string]any{"role": "user", "content": "How are you?"},
- map[string]any{"role": "assistant", "content": "I'm doing well!"},
- map[string]any{"role": "user", "content": "Tell me a joke"},
- },
- SessionContext: ctx,
- }
+ round1 := mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "hello")}, ""), ctx)
+ round2 := mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "hello"), msg("assistant", "Hi there!"), msg("user", "How are you?")}, ""), ctx)
+ round3 := mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "hello"), msg("assistant", "Hi there!"), msg("user", "How are you?"), msg("assistant", "I'm doing well!"), msg("user", "Tell me a joke")}, ""), ctx)
h1 := svc.GenerateSessionHash(round1)
h2 := svc.GenerateSessionHash(round2)
h3 := svc.GenerateSessionHash(round3)
-
require.NotEmpty(t, h1)
require.NotEmpty(t, h2)
require.NotEmpty(t, h3)
@@ -303,62 +208,20 @@ func TestGenerateSessionHash_ContinuousConversation_HashChangesWithMessages(t *t
func TestGenerateSessionHash_ContinuousConversation_SameRoundSameHash(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1}
-
- // 同一轮对话重复请求(如重试)应产生相同 hash
- mk := func() *ParsedRequest {
- return &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "Hi there!"},
- map[string]any{"role": "user", "content": "How are you?"},
- },
- SessionContext: ctx,
- }
- }
+ body := anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "hello"), msg("assistant", "Hi there!"), msg("user", "How are you?")}, "")
+ mk := func() *ParsedRequest { return mustParseSessionHashRequest(t, body, ctx) }
h1 := svc.GenerateSessionHash(mk())
h2 := svc.GenerateSessionHash(mk())
require.Equal(t, h1, h2, "same conversation state should produce identical hash on retry")
}
-// ============ 消息回退测试 ============
-
func TestGenerateSessionHash_MessageRollback(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1}
-
- // 模拟消息回退:用户删掉最后一轮再重发
- original := &ParsedRequest{
- System: "System prompt",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "msg1"},
- map[string]any{"role": "assistant", "content": "reply1"},
- map[string]any{"role": "user", "content": "msg2"},
- map[string]any{"role": "assistant", "content": "reply2"},
- map[string]any{"role": "user", "content": "msg3"},
- },
- SessionContext: ctx,
- }
-
- // 回退到 msg2 后,用新的 msg3 替代
- rollback := &ParsedRequest{
- System: "System prompt",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "msg1"},
- map[string]any{"role": "assistant", "content": "reply1"},
- map[string]any{"role": "user", "content": "msg2"},
- map[string]any{"role": "assistant", "content": "reply2"},
- map[string]any{"role": "user", "content": "different msg3"},
- },
- SessionContext: ctx,
- }
+ original := mustParseSessionHashRequest(t, anthropicSessionBody("System prompt", []any{msg("user", "msg1"), msg("assistant", "reply1"), msg("user", "msg2"), msg("assistant", "reply2"), msg("user", "msg3")}, ""), ctx)
+ rollback := mustParseSessionHashRequest(t, anthropicSessionBody("System prompt", []any{msg("user", "msg1"), msg("assistant", "reply1"), msg("user", "msg2"), msg("assistant", "reply2"), msg("user", "different msg3")}, ""), ctx)
hOrig := svc.GenerateSessionHash(original)
hRollback := svc.GenerateSessionHash(rollback)
@@ -367,58 +230,19 @@ func TestGenerateSessionHash_MessageRollback(t *testing.T) {
func TestGenerateSessionHash_MessageRollbackSameContent(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1}
-
- // 回退后重新发送相同内容 → 相同 hash(合理的粘性恢复)
- mk := func() *ParsedRequest {
- return &ParsedRequest{
- System: "System prompt",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "msg1"},
- map[string]any{"role": "assistant", "content": "reply1"},
- map[string]any{"role": "user", "content": "msg2"},
- },
- SessionContext: ctx,
- }
- }
+ body := anthropicSessionBody("System prompt", []any{msg("user", "msg1"), msg("assistant", "reply1"), msg("user", "msg2")}, "")
+ mk := func() *ParsedRequest { return mustParseSessionHashRequest(t, body, ctx) }
h1 := svc.GenerateSessionHash(mk())
h2 := svc.GenerateSessionHash(mk())
require.Equal(t, h1, h2, "rollback and resend same content should produce same hash")
}
-// ============ 相同 System、不同用户消息 ============
-
func TestGenerateSessionHash_SameSystemDifferentUsers(t *testing.T) {
svc := &GatewayService{}
-
- // 两个不同用户使用相同 system prompt 但发送不同消息
- user1 := &ParsedRequest{
- System: "You are a code reviewer.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "Review this Go code"},
- },
- SessionContext: &SessionContext{
- ClientIP: "1.1.1.1",
- UserAgent: "vscode",
- APIKeyID: 1,
- },
- }
- user2 := &ParsedRequest{
- System: "You are a code reviewer.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "Review this Python code"},
- },
- SessionContext: &SessionContext{
- ClientIP: "2.2.2.2",
- UserAgent: "vscode",
- APIKeyID: 2,
- },
- }
+ user1 := mustParseSessionHashRequest(t, anthropicSessionBody("You are a code reviewer.", []any{msg("user", "Review this Go code")}, ""), &SessionContext{ClientIP: "1.1.1.1", UserAgent: "vscode", APIKeyID: 1})
+ user2 := mustParseSessionHashRequest(t, anthropicSessionBody("You are a code reviewer.", []any{msg("user", "Review this Python code")}, ""), &SessionContext{ClientIP: "2.2.2.2", UserAgent: "vscode", APIKeyID: 2})
h1 := svc.GenerateSessionHash(user1)
h2 := svc.GenerateSessionHash(user2)
@@ -427,55 +251,20 @@ func TestGenerateSessionHash_SameSystemDifferentUsers(t *testing.T) {
func TestGenerateSessionHash_SameSystemSameMessageDifferentContext(t *testing.T) {
svc := &GatewayService{}
-
- // 这是修复的核心场景:两个不同用户发送完全相同的 system + messages(如 "hello")
- // 有了 SessionContext 后应该产生不同 hash
- user1 := &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: &SessionContext{
- ClientIP: "1.1.1.1",
- UserAgent: "Mozilla/5.0",
- APIKeyID: 10,
- },
- }
- user2 := &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: &SessionContext{
- ClientIP: "2.2.2.2",
- UserAgent: "Mozilla/5.0",
- APIKeyID: 20,
- },
- }
+ body := anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "hello")}, "")
+ user1 := mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "1.1.1.1", UserAgent: "Mozilla/5.0", APIKeyID: 10})
+ user2 := mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "2.2.2.2", UserAgent: "Mozilla/5.0", APIKeyID: 20})
h1 := svc.GenerateSessionHash(user1)
h2 := svc.GenerateSessionHash(user2)
require.NotEqual(t, h1, h2, "CRITICAL: same system+messages but different users should get different hashes")
}
-// ============ SessionContext 各字段独立影响测试 ============
-
func TestGenerateSessionHash_SessionContext_IPDifference(t *testing.T) {
svc := &GatewayService{}
-
+ body := anthropicSessionBody(nil, []any{msg("user", "test")}, "")
base := func(ip string) *ParsedRequest {
- return &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "test"},
- },
- SessionContext: &SessionContext{
- ClientIP: ip,
- UserAgent: "same-ua",
- APIKeyID: 1,
- },
- }
+ return mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: ip, UserAgent: "same-ua", APIKeyID: 1})
}
h1 := svc.GenerateSessionHash(base("1.1.1.1"))
@@ -485,18 +274,9 @@ func TestGenerateSessionHash_SessionContext_IPDifference(t *testing.T) {
func TestGenerateSessionHash_SessionContext_UADifference(t *testing.T) {
svc := &GatewayService{}
-
+ body := anthropicSessionBody(nil, []any{msg("user", "test")}, "")
base := func(ua string) *ParsedRequest {
- return &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "test"},
- },
- SessionContext: &SessionContext{
- ClientIP: "1.1.1.1",
- UserAgent: ua,
- APIKeyID: 1,
- },
- }
+ return mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "1.1.1.1", UserAgent: ua, APIKeyID: 1})
}
h1 := svc.GenerateSessionHash(base("Mozilla/5.0"))
@@ -506,18 +286,9 @@ func TestGenerateSessionHash_SessionContext_UADifference(t *testing.T) {
func TestGenerateSessionHash_SessionContext_UAVersionNoiseIgnored(t *testing.T) {
svc := &GatewayService{}
-
+ body := anthropicSessionBody(nil, []any{msg("user", "test")}, "")
base := func(ua string) *ParsedRequest {
- return &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "test"},
- },
- SessionContext: &SessionContext{
- ClientIP: "1.1.1.1",
- UserAgent: ua,
- APIKeyID: 1,
- },
- }
+ return mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "1.1.1.1", UserAgent: ua, APIKeyID: 1})
}
h1 := svc.GenerateSessionHash(base("Mozilla/5.0 codex_cli_rs/0.1.0"))
@@ -527,18 +298,9 @@ func TestGenerateSessionHash_SessionContext_UAVersionNoiseIgnored(t *testing.T)
func TestGenerateSessionHash_SessionContext_FreeformUAVersionNoiseIgnored(t *testing.T) {
svc := &GatewayService{}
-
+ body := anthropicSessionBody(nil, []any{msg("user", "test")}, "")
base := func(ua string) *ParsedRequest {
- return &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "test"},
- },
- SessionContext: &SessionContext{
- ClientIP: "1.1.1.1",
- UserAgent: ua,
- APIKeyID: 1,
- },
- }
+ return mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "1.1.1.1", UserAgent: ua, APIKeyID: 1})
}
h1 := svc.GenerateSessionHash(base("Codex CLI 0.1.0"))
@@ -548,18 +310,9 @@ func TestGenerateSessionHash_SessionContext_FreeformUAVersionNoiseIgnored(t *tes
func TestGenerateSessionHash_SessionContext_APIKeyIDDifference(t *testing.T) {
svc := &GatewayService{}
-
+ body := anthropicSessionBody(nil, []any{msg("user", "test")}, "")
base := func(keyID int64) *ParsedRequest {
- return &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "test"},
- },
- SessionContext: &SessionContext{
- ClientIP: "1.1.1.1",
- UserAgent: "same-ua",
- APIKeyID: keyID,
- },
- }
+ return mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "1.1.1.1", UserAgent: "same-ua", APIKeyID: keyID})
}
h1 := svc.GenerateSessionHash(base(1))
@@ -567,24 +320,12 @@ func TestGenerateSessionHash_SessionContext_APIKeyIDDifference(t *testing.T) {
require.NotEqual(t, h1, h2, "different APIKeyID should produce different hash")
}
-// ============ 多用户并发相同消息场景 ============
-
func TestGenerateSessionHash_MultipleUsersSameFirstMessage(t *testing.T) {
svc := &GatewayService{}
-
- // 模拟 5 个不同用户同时发送 "hello" → 应该产生 5 个不同的 hash
hashes := make(map[string]bool)
+ body := anthropicSessionBody(nil, []any{msg("user", "hello")}, "")
for i := 0; i < 5; i++ {
- parsed := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: &SessionContext{
- ClientIP: "192.168.1." + string(rune('1'+i)),
- UserAgent: "client-" + string(rune('A'+i)),
- APIKeyID: int64(i + 1),
- },
- }
+ parsed := mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "192.168.1." + string(rune('1'+i)), UserAgent: "client-" + string(rune('A'+i)), APIKeyID: int64(i + 1)})
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h)
require.False(t, hashes[h], "hash collision detected for user %d", i)
@@ -593,134 +334,56 @@ func TestGenerateSessionHash_MultipleUsersSameFirstMessage(t *testing.T) {
require.Len(t, hashes, 5, "5 different users should produce 5 unique hashes")
}
-// ============ 连续会话粘性:多轮对话同一用户 ============
-
func TestGenerateSessionHash_SameUserGrowingConversation(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "browser", APIKeyID: 42}
-
- // 模拟同一用户的连续会话,每轮 hash 不同但同用户重试保持一致
- messages := []map[string]any{
- {"role": "user", "content": "msg1"},
- {"role": "assistant", "content": "reply1"},
- {"role": "user", "content": "msg2"},
- {"role": "assistant", "content": "reply2"},
- {"role": "user", "content": "msg3"},
- {"role": "assistant", "content": "reply3"},
- {"role": "user", "content": "msg4"},
+ messages := []any{
+ msg("user", "msg1"), msg("assistant", "reply1"), msg("user", "msg2"), msg("assistant", "reply2"),
+ msg("user", "msg3"), msg("assistant", "reply3"), msg("user", "msg4"),
}
prevHash := ""
for round := 1; round <= len(messages); round += 2 {
- // 构建前 round 条消息
- msgs := make([]any, round)
- for j := 0; j < round; j++ {
- msgs[j] = messages[j]
- }
- parsed := &ParsedRequest{
- System: "System",
- HasSystem: true,
- Messages: msgs,
- SessionContext: ctx,
- }
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody("System", messages[:round], ""), ctx)
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h, "round %d hash should not be empty", round)
-
if prevHash != "" {
require.NotEqual(t, prevHash, h, "round %d hash should differ from previous round", round)
}
prevHash = h
-
- // 同一轮重试应该相同
h2 := svc.GenerateSessionHash(parsed)
require.Equal(t, h, h2, "retry of round %d should produce same hash", round)
}
}
-// ============ 多轮消息内容结构化测试 ============
-
func TestGenerateSessionHash_MultipleUserMessages(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1}
-
- // 5 条用户消息(无 assistant 回复)
- parsed := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "first"},
- map[string]any{"role": "user", "content": "second"},
- map[string]any{"role": "user", "content": "third"},
- map[string]any{"role": "user", "content": "fourth"},
- map[string]any{"role": "user", "content": "fifth"},
- },
- SessionContext: ctx,
- }
-
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody(nil, []any{msg("user", "first"), msg("user", "second"), msg("user", "third"), msg("user", "fourth"), msg("user", "fifth")}, ""), ctx)
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h)
- // 修改中间一条消息应该改变 hash
- parsed2 := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "first"},
- map[string]any{"role": "user", "content": "CHANGED"},
- map[string]any{"role": "user", "content": "third"},
- map[string]any{"role": "user", "content": "fourth"},
- map[string]any{"role": "user", "content": "fifth"},
- },
- SessionContext: ctx,
- }
-
+ parsed2 := mustParseSessionHashRequest(t, anthropicSessionBody(nil, []any{msg("user", "first"), msg("user", "CHANGED"), msg("user", "third"), msg("user", "fourth"), msg("user", "fifth")}, ""), ctx)
h2 := svc.GenerateSessionHash(parsed2)
require.NotEqual(t, h, h2, "changing any message should change the hash")
}
func TestGenerateSessionHash_MessageOrderMatters(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1}
-
- parsed1 := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "alpha"},
- map[string]any{"role": "user", "content": "beta"},
- },
- SessionContext: ctx,
- }
- parsed2 := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "beta"},
- map[string]any{"role": "user", "content": "alpha"},
- },
- SessionContext: ctx,
- }
+ parsed1 := mustParseSessionHashRequest(t, anthropicSessionBody(nil, []any{msg("user", "alpha"), msg("user", "beta")}, ""), ctx)
+ parsed2 := mustParseSessionHashRequest(t, anthropicSessionBody(nil, []any{msg("user", "beta"), msg("user", "alpha")}, ""), ctx)
h1 := svc.GenerateSessionHash(parsed1)
h2 := svc.GenerateSessionHash(parsed2)
require.NotEqual(t, h1, h2, "message order should affect the hash")
}
-// ============ 复杂内容格式测试 ============
-
func TestGenerateSessionHash_StructuredContent(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1}
-
- // 结构化 content(数组形式)
- parsed := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "content": []any{
- map[string]any{"type": "text", "text": "Look at this"},
- map[string]any{"type": "text", "text": "And this too"},
- },
- },
- },
- SessionContext: ctx,
- }
+ content := []any{map[string]any{"type": "text", "text": "Look at this"}, map[string]any{"type": "text", "text": "And this too"}}
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody(nil, []any{msg("user", content)}, ""), ctx)
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h, "structured content should produce a hash")
@@ -728,100 +391,37 @@ func TestGenerateSessionHash_StructuredContent(t *testing.T) {
func TestGenerateSessionHash_ArraySystemPrompt(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1}
-
- // 数组格式的 system prompt
- parsed := &ParsedRequest{
- System: []any{
- map[string]any{"type": "text", "text": "You are a helpful assistant."},
- map[string]any{"type": "text", "text": "Be concise."},
- },
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: ctx,
- }
+ system := []any{map[string]any{"type": "text", "text": "You are a helpful assistant."}, map[string]any{"type": "text", "text": "Be concise."}}
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody(system, []any{msg("user", "hello")}, ""), ctx)
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h, "array system prompt should produce a hash")
}
-// ============ SessionContext 与 cache_control 优先级 ============
-
func TestGenerateSessionHash_CacheControlOverridesSessionContext(t *testing.T) {
svc := &GatewayService{}
-
- // 当有 cache_control: ephemeral 时,使用第 2 级优先级
- // SessionContext 不应影响结果
- parsed1 := &ParsedRequest{
- System: []any{
- map[string]any{
- "type": "text",
- "text": "You are a tool-specific assistant.",
- "cache_control": map[string]any{"type": "ephemeral"},
- },
- },
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: &SessionContext{
- ClientIP: "1.1.1.1",
- UserAgent: "ua1",
- APIKeyID: 100,
- },
- }
- parsed2 := &ParsedRequest{
- System: []any{
- map[string]any{
- "type": "text",
- "text": "You are a tool-specific assistant.",
- "cache_control": map[string]any{"type": "ephemeral"},
- },
- },
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: &SessionContext{
- ClientIP: "2.2.2.2",
- UserAgent: "ua2",
- APIKeyID: 200,
- },
- }
+ system := []any{map[string]any{"type": "text", "text": "You are a tool-specific assistant.", "cache_control": map[string]any{"type": "ephemeral"}}}
+ body := anthropicSessionBody(system, []any{msg("user", "hello")}, "")
+ parsed1 := mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "1.1.1.1", UserAgent: "ua1", APIKeyID: 100})
+ parsed2 := mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "2.2.2.2", UserAgent: "ua2", APIKeyID: 200})
h1 := svc.GenerateSessionHash(parsed1)
h2 := svc.GenerateSessionHash(parsed2)
require.Equal(t, h1, h2, "cache_control ephemeral has higher priority, SessionContext should not affect result")
}
-// ============ 边界情况 ============
-
func TestGenerateSessionHash_EmptyMessages(t *testing.T) {
svc := &GatewayService{}
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody(nil, []any{}, ""), &SessionContext{ClientIP: "1.1.1.1", UserAgent: "test", APIKeyID: 1})
- parsed := &ParsedRequest{
- Messages: []any{},
- SessionContext: &SessionContext{
- ClientIP: "1.1.1.1",
- UserAgent: "test",
- APIKeyID: 1,
- },
- }
-
- // 空 messages + 只有 SessionContext 时,combined.Len() > 0 因为有 context 写入
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h, "empty messages with SessionContext should still produce a hash from context")
}
func TestGenerateSessionHash_EmptyMessagesNoContext(t *testing.T) {
svc := &GatewayService{}
-
- parsed := &ParsedRequest{
- Messages: []any{},
- }
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody(nil, []any{}, ""), nil)
h := svc.GenerateSessionHash(parsed)
require.Empty(t, h, "empty messages without SessionContext should produce empty hash")
@@ -829,98 +429,37 @@ func TestGenerateSessionHash_EmptyMessagesNoContext(t *testing.T) {
func TestGenerateSessionHash_SessionContextWithEmptyFields(t *testing.T) {
svc := &GatewayService{}
-
- // SessionContext 字段为空字符串和零值时仍应影响 hash
- withEmptyCtx := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "test"},
- },
- SessionContext: &SessionContext{
- ClientIP: "",
- UserAgent: "",
- APIKeyID: 0,
- },
- }
- withoutCtx := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "test"},
- },
- }
+ body := anthropicSessionBody(nil, []any{msg("user", "test")}, "")
+ withEmptyCtx := mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "", UserAgent: "", APIKeyID: 0})
+ withoutCtx := mustParseSessionHashRequest(t, body, nil)
h1 := svc.GenerateSessionHash(withEmptyCtx)
h2 := svc.GenerateSessionHash(withoutCtx)
- // 有 SessionContext(即使字段为空)仍然会写入分隔符 "::" 等
require.NotEqual(t, h1, h2, "empty-field SessionContext should still differ from nil SessionContext")
}
-// ============ 长对话历史测试 ============
-
func TestGenerateSessionHash_LongConversation(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1}
-
- // 构建 20 轮对话
messages := make([]any, 0, 40)
for i := 0; i < 20; i++ {
- messages = append(messages, map[string]any{
- "role": "user",
- "content": "user message " + string(rune('A'+i)),
- })
- messages = append(messages, map[string]any{
- "role": "assistant",
- "content": "assistant reply " + string(rune('A'+i)),
- })
- }
-
- parsed := &ParsedRequest{
- System: "System prompt",
- HasSystem: true,
- Messages: messages,
- SessionContext: ctx,
+ messages = append(messages, msg("user", "user message "+string(rune('A'+i))))
+ messages = append(messages, msg("assistant", "assistant reply "+string(rune('A'+i))))
}
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody("System prompt", messages, ""), ctx)
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h)
- // 再加一轮应该不同
- moreMessages := make([]any, len(messages)+2)
- copy(moreMessages, messages)
- moreMessages[len(messages)] = map[string]any{"role": "user", "content": "one more"}
- moreMessages[len(messages)+1] = map[string]any{"role": "assistant", "content": "ok"}
-
- parsed2 := &ParsedRequest{
- System: "System prompt",
- HasSystem: true,
- Messages: moreMessages,
- SessionContext: ctx,
- }
-
+ moreMessages := append(append([]any{}, messages...), msg("user", "one more"), msg("assistant", "ok"))
+ parsed2 := mustParseSessionHashRequest(t, anthropicSessionBody("System prompt", moreMessages, ""), ctx)
h2 := svc.GenerateSessionHash(parsed2)
require.NotEqual(t, h, h2, "adding more messages to long conversation should change hash")
}
-// ============ Gemini 原生格式 session hash 测试 ============
-
func TestGenerateSessionHash_GeminiContentsProducesHash(t *testing.T) {
svc := &GatewayService{}
-
- // Gemini 格式: contents[].parts[].text
- parsed := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{
- map[string]any{"text": "Hello from Gemini"},
- },
- },
- },
- SessionContext: &SessionContext{
- ClientIP: "1.2.3.4",
- UserAgent: "gemini-cli",
- APIKeyID: 1,
- },
- }
+ parsed := mustParseGeminiSessionHashRequest(t, geminiSessionBody(nil, []any{geminiMsg("user", "Hello from Gemini")}), &SessionContext{ClientIP: "1.2.3.4", UserAgent: "gemini-cli", APIKeyID: 1})
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h, "Gemini contents with parts should produce a non-empty hash")
@@ -928,31 +467,9 @@ func TestGenerateSessionHash_GeminiContentsProducesHash(t *testing.T) {
func TestGenerateSessionHash_GeminiDifferentContentsDifferentHash(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "gemini-cli", APIKeyID: 1}
-
- parsed1 := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{
- map[string]any{"text": "Hello"},
- },
- },
- },
- SessionContext: ctx,
- }
- parsed2 := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{
- map[string]any{"text": "Goodbye"},
- },
- },
- },
- SessionContext: ctx,
- }
+ parsed1 := mustParseGeminiSessionHashRequest(t, geminiSessionBody(nil, []any{geminiMsg("user", "Hello")}), ctx)
+ parsed2 := mustParseGeminiSessionHashRequest(t, geminiSessionBody(nil, []any{geminiMsg("user", "Goodbye")}), ctx)
h1 := svc.GenerateSessionHash(parsed1)
h2 := svc.GenerateSessionHash(parsed2)
@@ -961,28 +478,9 @@ func TestGenerateSessionHash_GeminiDifferentContentsDifferentHash(t *testing.T)
func TestGenerateSessionHash_GeminiSameContentsSameHash(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "gemini-cli", APIKeyID: 1}
-
- mk := func() *ParsedRequest {
- return &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{
- map[string]any{"text": "Hello"},
- },
- },
- map[string]any{
- "role": "model",
- "parts": []any{
- map[string]any{"text": "Hi there!"},
- },
- },
- },
- SessionContext: ctx,
- }
- }
+ body := geminiSessionBody(nil, []any{geminiMsg("user", "Hello"), geminiMsg("model", "Hi there!")})
+ mk := func() *ParsedRequest { return mustParseGeminiSessionHashRequest(t, body, ctx) }
h1 := svc.GenerateSessionHash(mk())
h2 := svc.GenerateSessionHash(mk())
@@ -991,36 +489,9 @@ func TestGenerateSessionHash_GeminiSameContentsSameHash(t *testing.T) {
func TestGenerateSessionHash_GeminiMultiTurnHashChanges(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "gemini-cli", APIKeyID: 1}
-
- round1 := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{map[string]any{"text": "hello"}},
- },
- },
- SessionContext: ctx,
- }
-
- round2 := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{map[string]any{"text": "hello"}},
- },
- map[string]any{
- "role": "model",
- "parts": []any{map[string]any{"text": "Hi!"}},
- },
- map[string]any{
- "role": "user",
- "parts": []any{map[string]any{"text": "How are you?"}},
- },
- },
- SessionContext: ctx,
- }
+ round1 := mustParseGeminiSessionHashRequest(t, geminiSessionBody(nil, []any{geminiMsg("user", "hello")}), ctx)
+ round2 := mustParseGeminiSessionHashRequest(t, geminiSessionBody(nil, []any{geminiMsg("user", "hello"), geminiMsg("model", "Hi!"), geminiMsg("user", "How are you?")}), ctx)
h1 := svc.GenerateSessionHash(round1)
h2 := svc.GenerateSessionHash(round2)
@@ -1031,34 +502,9 @@ func TestGenerateSessionHash_GeminiMultiTurnHashChanges(t *testing.T) {
func TestGenerateSessionHash_GeminiDifferentUsersSameContentDifferentHash(t *testing.T) {
svc := &GatewayService{}
-
- // 核心场景:两个不同用户发送相同 Gemini 格式消息应得到不同 hash
- user1 := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{map[string]any{"text": "hello"}},
- },
- },
- SessionContext: &SessionContext{
- ClientIP: "1.1.1.1",
- UserAgent: "gemini-cli",
- APIKeyID: 10,
- },
- }
- user2 := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{map[string]any{"text": "hello"}},
- },
- },
- SessionContext: &SessionContext{
- ClientIP: "2.2.2.2",
- UserAgent: "gemini-cli",
- APIKeyID: 20,
- },
- }
+ body := geminiSessionBody(nil, []any{geminiMsg("user", "hello")})
+ user1 := mustParseGeminiSessionHashRequest(t, body, &SessionContext{ClientIP: "1.1.1.1", UserAgent: "gemini-cli", APIKeyID: 10})
+ user2 := mustParseGeminiSessionHashRequest(t, body, &SessionContext{ClientIP: "2.2.2.2", UserAgent: "gemini-cli", APIKeyID: 20})
h1 := svc.GenerateSessionHash(user1)
h2 := svc.GenerateSessionHash(user2)
@@ -1067,31 +513,9 @@ func TestGenerateSessionHash_GeminiDifferentUsersSameContentDifferentHash(t *tes
func TestGenerateSessionHash_GeminiSystemInstructionAffectsHash(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "gemini-cli", APIKeyID: 1}
-
- // systemInstruction 经 ParseGatewayRequest 解析后存入 parsed.System
- withSys := &ParsedRequest{
- System: []any{
- map[string]any{"text": "You are a coding assistant."},
- },
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{map[string]any{"text": "hello"}},
- },
- },
- SessionContext: ctx,
- }
- withoutSys := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{map[string]any{"text": "hello"}},
- },
- },
- SessionContext: ctx,
- }
+ withSys := mustParseGeminiSessionHashRequest(t, geminiSessionBody([]any{map[string]any{"text": "You are a coding assistant."}}, []any{geminiMsg("user", "hello")}), ctx)
+ withoutSys := mustParseGeminiSessionHashRequest(t, geminiSessionBody(nil, []any{geminiMsg("user", "hello")}), ctx)
h1 := svc.GenerateSessionHash(withSys)
h2 := svc.GenerateSessionHash(withoutSys)
@@ -1100,64 +524,21 @@ func TestGenerateSessionHash_GeminiSystemInstructionAffectsHash(t *testing.T) {
func TestGenerateSessionHash_GeminiMultiPartMessage(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "gemini-cli", APIKeyID: 1}
-
- // 多 parts 的消息
- parsed := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{
- map[string]any{"text": "Part 1"},
- map[string]any{"text": "Part 2"},
- map[string]any{"text": "Part 3"},
- },
- },
- },
- SessionContext: ctx,
- }
-
+ parsed := mustParseGeminiSessionHashRequest(t, geminiSessionBody(nil, []any{geminiMsg("user", "Part 1", "Part 2", "Part 3")}), ctx)
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h, "multi-part Gemini message should produce a hash")
- // 不同内容的多 parts
- parsed2 := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{
- map[string]any{"text": "Part 1"},
- map[string]any{"text": "CHANGED"},
- map[string]any{"text": "Part 3"},
- },
- },
- },
- SessionContext: ctx,
- }
-
+ parsed2 := mustParseGeminiSessionHashRequest(t, geminiSessionBody(nil, []any{geminiMsg("user", "Part 1", "CHANGED", "Part 3")}), ctx)
h2 := svc.GenerateSessionHash(parsed2)
require.NotEqual(t, h, h2, "changing a part should change the hash")
}
func TestGenerateSessionHash_GeminiNonTextPartsIgnored(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "gemini-cli", APIKeyID: 1}
-
- // 含非 text 类型 parts(如 inline_data),应被跳过但不报错
- parsed := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{
- map[string]any{"text": "Describe this image"},
- map[string]any{"inline_data": map[string]any{"mime_type": "image/png", "data": "base64..."}},
- },
- },
- },
- SessionContext: ctx,
- }
+ content := []any{map[string]any{"role": "user", "parts": []any{map[string]any{"text": "Describe this image"}, map[string]any{"inline_data": map[string]any{"mime_type": "image/png", "data": "base64..."}}}}}
+ parsed := mustParseGeminiSessionHashRequest(t, geminiSessionBody(nil, content), ctx)
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h, "Gemini message with mixed parts should still produce a hash from text parts")
@@ -1165,107 +546,41 @@ func TestGenerateSessionHash_GeminiNonTextPartsIgnored(t *testing.T) {
func TestGenerateSessionHash_GeminiMultiTurnHashNotSticky(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "10.0.0.1", UserAgent: "gemini-cli", APIKeyID: 42}
+ rounds := []string{
+ geminiSessionBody([]any{map[string]any{"text": "You are a coding assistant."}}, []any{geminiMsg("user", "Write a Go function")}),
+ geminiSessionBody([]any{map[string]any{"text": "You are a coding assistant."}}, []any{geminiMsg("user", "Write a Go function"), geminiMsg("model", "func hello() {}"), geminiMsg("user", "Add error handling")}),
+ geminiSessionBody([]any{map[string]any{"text": "You are a coding assistant."}}, []any{geminiMsg("user", "Write a Go function"), geminiMsg("model", "func hello() {}"), geminiMsg("user", "Add error handling"), geminiMsg("model", "func hello() error { return nil }"), geminiMsg("user", "Now add tests")}),
+ }
- // 模拟同一 Gemini 会话的三轮请求,每轮 contents 累积增长。
- // 验证预期行为:每轮 hash 都不同,即 GenerateSessionHash 不具备跨轮粘性。
- // 这是 by-design 的——Gemini 的跨轮粘性由 Digest Fallback(BuildGeminiDigestChain)负责。
- round1Body := []byte(`{
- "systemInstruction": {"parts": [{"text": "You are a coding assistant."}]},
- "contents": [
- {"role": "user", "parts": [{"text": "Write a Go function"}]}
- ]
- }`)
- round2Body := []byte(`{
- "systemInstruction": {"parts": [{"text": "You are a coding assistant."}]},
- "contents": [
- {"role": "user", "parts": [{"text": "Write a Go function"}]},
- {"role": "model", "parts": [{"text": "func hello() {}"}]},
- {"role": "user", "parts": [{"text": "Add error handling"}]}
- ]
- }`)
- round3Body := []byte(`{
- "systemInstruction": {"parts": [{"text": "You are a coding assistant."}]},
- "contents": [
- {"role": "user", "parts": [{"text": "Write a Go function"}]},
- {"role": "model", "parts": [{"text": "func hello() {}"}]},
- {"role": "user", "parts": [{"text": "Add error handling"}]},
- {"role": "model", "parts": [{"text": "func hello() error { return nil }"}]},
- {"role": "user", "parts": [{"text": "Now add tests"}]}
- ]
- }`)
-
- hashes := make([]string, 3)
- for i, body := range [][]byte{round1Body, round2Body, round3Body} {
- parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "gemini")
- require.NoError(t, err)
- parsed.SessionContext = ctx
+ hashes := make([]string, len(rounds))
+ for i, body := range rounds {
+ parsed := mustParseGeminiSessionHashRequest(t, body, ctx)
hashes[i] = svc.GenerateSessionHash(parsed)
require.NotEmpty(t, hashes[i], "round %d hash should not be empty", i+1)
}
-
- // 每轮 hash 都不同——这是预期行为
require.NotEqual(t, hashes[0], hashes[1], "round 1 vs 2 hash should differ (contents grow)")
require.NotEqual(t, hashes[1], hashes[2], "round 2 vs 3 hash should differ (contents grow)")
require.NotEqual(t, hashes[0], hashes[2], "round 1 vs 3 hash should differ")
- // 同一轮重试应产生相同 hash
- parsed1Again, err := ParseGatewayRequest(NewRequestBodyRef(round2Body), "gemini")
- require.NoError(t, err)
- parsed1Again.SessionContext = ctx
- h2Again := svc.GenerateSessionHash(parsed1Again)
+ parsedAgain := mustParseGeminiSessionHashRequest(t, rounds[1], ctx)
+ h2Again := svc.GenerateSessionHash(parsedAgain)
require.Equal(t, hashes[1], h2Again, "retry of same round should produce same hash")
}
func TestGenerateSessionHash_GeminiEndToEnd(t *testing.T) {
svc := &GatewayService{}
-
- // 端到端测试:模拟 ParseGatewayRequest + GenerateSessionHash 完整流程
- body := []byte(`{
- "model": "gemini-2.5-pro",
- "systemInstruction": {
- "parts": [{"text": "You are a coding assistant."}]
- },
- "contents": [
- {"role": "user", "parts": [{"text": "Write a Go function"}]},
- {"role": "model", "parts": [{"text": "Here is a function..."}]},
- {"role": "user", "parts": [{"text": "Now add error handling"}]}
- ]
- }`)
-
- parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "gemini")
- require.NoError(t, err)
- parsed.SessionContext = &SessionContext{
- ClientIP: "10.0.0.1",
- UserAgent: "gemini-cli/1.0",
- APIKeyID: 42,
- }
+ body := geminiSessionBody([]any{map[string]any{"text": "You are a coding assistant."}}, []any{geminiMsg("user", "Write a Go function"), geminiMsg("model", "Here is a function..."), geminiMsg("user", "Now add error handling")})
+ parsed := mustParseGeminiSessionHashRequest(t, body, &SessionContext{ClientIP: "10.0.0.1", UserAgent: "gemini-cli/1.0", APIKeyID: 42})
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h, "end-to-end Gemini flow should produce a hash")
- // 同一请求再次解析应产生相同 hash
- parsed2, err := ParseGatewayRequest(NewRequestBodyRef(body), "gemini")
- require.NoError(t, err)
- parsed2.SessionContext = &SessionContext{
- ClientIP: "10.0.0.1",
- UserAgent: "gemini-cli/1.0",
- APIKeyID: 42,
- }
-
+ parsed2 := mustParseGeminiSessionHashRequest(t, body, &SessionContext{ClientIP: "10.0.0.1", UserAgent: "gemini-cli/1.0", APIKeyID: 42})
h2 := svc.GenerateSessionHash(parsed2)
require.Equal(t, h, h2, "same request should produce same hash")
- // 不同用户发送相同请求应产生不同 hash
- parsed3, err := ParseGatewayRequest(NewRequestBodyRef(body), "gemini")
- require.NoError(t, err)
- parsed3.SessionContext = &SessionContext{
- ClientIP: "10.0.0.2",
- UserAgent: "gemini-cli/1.0",
- APIKeyID: 99,
- }
-
+ parsed3 := mustParseGeminiSessionHashRequest(t, body, &SessionContext{ClientIP: "10.0.0.2", UserAgent: "gemini-cli/1.0", APIKeyID: 99})
h3 := svc.GenerateSessionHash(parsed3)
require.NotEqual(t, h, h3, "different user with same Gemini request should get different hash")
}
diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go
index cd5a4015..f8c34031 100644
--- a/backend/internal/service/openai_gateway_service.go
+++ b/backend/internal/service/openai_gateway_service.go
@@ -36,6 +36,13 @@ import (
"go.uber.org/zap"
)
+// openAIParsedRequestBodyCache 绑定 body 指纹,避免 handler 预解析的旧 body 污染后续 forwardBody。
+type openAIParsedRequestBodyCache struct {
+ bodyHash uint64
+ bodyLen int
+ reqBody map[string]any
+}
+
const (
// ChatGPT internal API for OAuth accounts
chatgptCodexURL = "https://chatgpt.com/backend-api/codex/responses"
@@ -2284,6 +2291,20 @@ func (s *OpenAIGatewayService) shouldFailoverOpenAIUpstreamResponse(statusCode i
return isOpenAITransientProcessingError(statusCode, upstreamMsg, upstreamBody)
}
+func marshalOpenAIUpstreamJSON(v any) ([]byte, error) {
+ var buf bytes.Buffer
+ enc := json.NewEncoder(&buf)
+ enc.SetEscapeHTML(false)
+ if err := enc.Encode(v); err != nil {
+ return nil, err
+ }
+ out := buf.Bytes()
+ if len(out) > 0 && out[len(out)-1] == '\n' {
+ out = out[:len(out)-1]
+ }
+ return out, nil
+}
+
func (s *OpenAIGatewayService) handleFailoverSideEffects(ctx context.Context, resp *http.Response, account *Account, requestedModel ...string) {
body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
if len(requestedModel) > 0 {
@@ -2561,7 +2582,11 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
// 确保高版本模型向低版本模型映射不报错
if !SupportsVerbosity(upstreamModel) {
if text, ok := reqBody["text"].(map[string]any); ok {
- delete(text, "verbosity")
+ if _, exists := text["verbosity"]; exists {
+ delete(text, "verbosity")
+ bodyModified = true
+ markPatchDelete("text.verbosity")
+ }
}
}
}
@@ -2762,7 +2787,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
}
if !serializedByPatch {
var marshalErr error
- body, marshalErr = json.Marshal(reqBody)
+ body, marshalErr = marshalOpenAIUpstreamJSON(reqBody)
if marshalErr != nil {
return nil, fmt.Errorf("serialize request body: %w", marshalErr)
}
@@ -3043,7 +3068,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
upstreamCode := extractUpstreamErrorCode(respBody)
if !httpInvalidEncryptedContentRetryTried && resp.StatusCode == http.StatusBadRequest && upstreamCode == "invalid_encrypted_content" {
if trimOpenAIEncryptedReasoningItems(reqBody) {
- body, err = json.Marshal(reqBody)
+ body, err = marshalOpenAIUpstreamJSON(reqBody)
if err != nil {
return nil, fmt.Errorf("serialize invalid_encrypted_content retry body: %w", err)
}
@@ -3074,6 +3099,8 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
})
s.handleFailoverSideEffects(ctx, resp, account, upstreamModel)
+ // reqBody 会被本次账号尝试原地修改,failover 前必须释放,避免下一账号复用脏 map。
+ releaseOpenAIParsedRequestBody(c)
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
@@ -6711,7 +6738,7 @@ func sanitizeEmptyBase64InputImagesInOpenAIBody(body []byte) ([]byte, bool, erro
if !sanitizeEmptyBase64InputImagesInOpenAIRequestBodyMap(reqBody) {
return body, false, nil
}
- normalized, err := json.Marshal(reqBody)
+ normalized, err := marshalOpenAIUpstreamJSON(reqBody)
if err != nil {
return body, false, fmt.Errorf("serialize sanitized request body: %w", err)
}
@@ -6817,10 +6844,13 @@ func isEmptyBase64DataURI(raw string) bool {
}
func getOpenAIRequestBodyMap(c *gin.Context, body []byte) (map[string]any, error) {
+ // 同一个 gin.Context 内 failover/渠道映射可能传入新 body,缓存必须先校验 body 指纹。
+ bodyHash := xxhash.Sum64(body)
+ bodyLen := len(body)
if c != nil {
if cached, ok := c.Get(OpenAIParsedRequestBodyKey); ok {
- if reqBody, ok := cached.(map[string]any); ok && reqBody != nil {
- return reqBody, nil
+ if cache, ok := cached.(openAIParsedRequestBodyCache); ok && cache.reqBody != nil && cache.bodyLen == bodyLen && cache.bodyHash == bodyHash {
+ return cache.reqBody, nil
}
}
}
@@ -6830,11 +6860,36 @@ func getOpenAIRequestBodyMap(c *gin.Context, body []byte) (map[string]any, error
return nil, fmt.Errorf("parse request: %w", err)
}
if c != nil {
- c.Set(OpenAIParsedRequestBodyKey, reqBody)
+ c.Set(OpenAIParsedRequestBodyKey, openAIParsedRequestBodyCache{bodyHash: bodyHash, bodyLen: bodyLen, reqBody: reqBody})
}
return reqBody, nil
}
+// CacheOpenAIParsedRequestBody 仅缓存与当前 body 绑定的解析结果。
+func CacheOpenAIParsedRequestBody(c *gin.Context, body []byte, reqBody map[string]any) {
+ if c == nil || reqBody == nil {
+ return
+ }
+ c.Set(OpenAIParsedRequestBodyKey, openAIParsedRequestBodyCache{
+ bodyHash: xxhash.Sum64(body),
+ bodyLen: len(body),
+ reqBody: reqBody,
+ })
+}
+
+// CachedOpenAIParsedRequestBody 只给同请求内不关心 body 参数的轻量识别逻辑使用。
+func CachedOpenAIParsedRequestBody(c *gin.Context) map[string]any {
+ if c == nil {
+ return nil
+ }
+ if cached, ok := c.Get(OpenAIParsedRequestBodyKey); ok {
+ if cache, ok := cached.(openAIParsedRequestBodyCache); ok {
+ return cache.reqBody
+ }
+ }
+ return nil
+}
+
func releaseOpenAIParsedRequestBody(c *gin.Context) {
if c == nil {
return
diff --git a/backend/internal/service/openai_gateway_service_hotpath_test.go b/backend/internal/service/openai_gateway_service_hotpath_test.go
index 234dee00..2ff72e2f 100644
--- a/backend/internal/service/openai_gateway_service_hotpath_test.go
+++ b/backend/internal/service/openai_gateway_service_hotpath_test.go
@@ -112,7 +112,7 @@ func TestGetOpenAIRequestBodyMap_UsesContextCache(t *testing.T) {
c, _ := gin.CreateTestContext(rec)
cached := map[string]any{"model": "cached-model", "stream": true}
- c.Set(OpenAIParsedRequestBodyKey, cached)
+ CacheOpenAIParsedRequestBody(c, []byte(`{invalid-json`), cached)
got, err := getOpenAIRequestBodyMap(c, []byte(`{invalid-json`))
require.NoError(t, err)
@@ -134,11 +134,19 @@ func TestGetOpenAIRequestBodyMap_WriteBackContextCache(t *testing.T) {
require.NoError(t, err)
require.Equal(t, "gpt-5", got["model"])
- cached, ok := c.Get(OpenAIParsedRequestBodyKey)
- require.True(t, ok)
- cachedMap, ok := cached.(map[string]any)
- require.True(t, ok)
- require.Equal(t, got, cachedMap)
+ require.Equal(t, got, CachedOpenAIParsedRequestBody(c))
+}
+
+func TestGetOpenAIRequestBodyMap_IgnoresCacheForDifferentBody(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+
+ CacheOpenAIParsedRequestBody(c, []byte(`{"model":"cached-model"}`), map[string]any{"model": "cached-model"})
+
+ got, err := getOpenAIRequestBodyMap(c, []byte(`{"model":"forward-model"}`))
+ require.NoError(t, err)
+ require.Equal(t, "forward-model", got["model"])
}
func TestSanitizeEmptyBase64InputImagesInOpenAIRequestBodyMap(t *testing.T) {
diff --git a/backend/internal/service/openai_oauth_passthrough_test.go b/backend/internal/service/openai_oauth_passthrough_test.go
index 398cbb85..2710c696 100644
--- a/backend/internal/service/openai_oauth_passthrough_test.go
+++ b/backend/internal/service/openai_oauth_passthrough_test.go
@@ -14,6 +14,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/openai_compat"
"github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
@@ -101,6 +102,57 @@ func TestOpenAIGatewayService_ResponsesUnknownModelDoesNotFallbackToGPT54(t *tes
require.True(t, rec.Code >= http.StatusBadRequest)
}
+func TestOpenAIGatewayService_NativeResponsesBodyModificationPreservesHTMLChars(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ payloadText := strings.Repeat(`&value`, 128)
+ originalBody := []byte(fmt.Sprintf(`{"model":"gpt-5.5","stream":false,"max_output_tokens":100,"previous_response_id":"resp_prev","input":[{"type":"message","role":"user","content":[{"type":"input_text","text":%q}]}]}`, payloadText))
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(originalBody))
+ c.Request.Header.Set("Content-Type", "application/json")
+
+ upstream := &httpUpstreamRecorder{resp: &http.Response{
+ StatusCode: http.StatusBadRequest,
+ Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid_native_reencode"}},
+ Body: io.NopCloser(strings.NewReader(`{"error":{"type":"invalid_request_error","message":"stop after capture"}}`)),
+ }}
+ svc := &OpenAIGatewayService{
+ cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{
+ Enabled: false,
+ AllowInsecureHTTP: true,
+ }}},
+ httpUpstream: upstream,
+ }
+ account := &Account{
+ ID: 456,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "http://upstream.example",
+ },
+ Extra: map[string]any{
+ openai_compat.ExtraKeyResponsesMode: string(openai_compat.ResponsesSupportModeAuto),
+ openai_compat.ExtraKeyResponsesSupported: true,
+ },
+ Status: StatusActive,
+ Schedulable: true,
+ }
+
+ result, err := svc.Forward(context.Background(), c, account, originalBody)
+ require.Error(t, err)
+ require.Nil(t, result)
+ require.NotNil(t, upstream.lastReq)
+ require.Equal(t, "http://upstream.example/v1/responses", upstream.lastReq.URL.String())
+ require.Contains(t, string(upstream.lastBody), payloadText)
+ require.NotContains(t, string(upstream.lastBody), `\\u003c`)
+ require.NotContains(t, string(upstream.lastBody), `\\u003e`)
+ require.NotContains(t, string(upstream.lastBody), `\\u0026`)
+}
+
func TestOpenAIGatewayService_OAuthMessagesBridgeDoesNotInjectDefaultInstructions(t *testing.T) {
gin.SetMode(gin.TestMode)
diff --git a/backend/internal/service/user_msg_queue_service.go b/backend/internal/service/user_msg_queue_service.go
index a0ce95a8..f3f105ac 100644
--- a/backend/internal/service/user_msg_queue_service.go
+++ b/backend/internal/service/user_msg_queue_service.go
@@ -12,6 +12,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
+ "github.com/tidwall/gjson"
)
// UserMsgQueueCache 用户消息串行队列 Redis 缓存接口
@@ -62,43 +63,48 @@ func NewUserMessageQueueService(cache UserMsgQueueCache, rpmCache RPMCache, cfg
// 2. 最后一条消息 role == "user"
// 3. 最后一条消息 content(如果是数组)中不含 type:"tool_result" / "tool_use_result"
func IsRealUserMessage(parsed *ParsedRequest) bool {
- if parsed == nil || len(parsed.Messages) == 0 {
+ if parsed == nil {
+ return false
+ }
+ messagesRaw := parsed.MessagesRaw()
+ if len(messagesRaw) == 0 {
return false
}
- lastMsg := parsed.Messages[len(parsed.Messages)-1]
- msgMap, ok := lastMsg.(map[string]any)
- if !ok {
+ messages := gjson.ParseBytes(messagesRaw)
+ if !messages.IsArray() {
+ return false
+ }
+ lastMsg := gjson.Result{}
+ messages.ForEach(func(_, msg gjson.Result) bool {
+ lastMsg = msg
+ return true
+ })
+ if !lastMsg.Exists() || !lastMsg.IsObject() {
+ return false
+ }
+ if lastMsg.Get("role").String() != "user" {
return false
}
- role, _ := msgMap["role"].(string)
- if role != "user" {
- return false
+ content := lastMsg.Get("content")
+ if !content.Exists() {
+ return true
+ }
+ if !content.IsArray() {
+ return true
}
- // 检查 content 是否包含 tool_result 类型
- content, ok := msgMap["content"]
- if !ok {
- return true // 没有 content 字段,视为普通用户消息
- }
-
- contentArr, ok := content.([]any)
- if !ok {
- return true // content 不是数组(可能是 string),视为普通用户消息
- }
-
- for _, item := range contentArr {
- itemMap, ok := item.(map[string]any)
- if !ok {
- continue
- }
- itemType, _ := itemMap["type"].(string)
+ isReal := true
+ content.ForEach(func(_, item gjson.Result) bool {
+ itemType := item.Get("type").String()
if itemType == "tool_result" || itemType == "tool_use_result" {
+ isReal = false
return false
}
- }
- return true
+ return true
+ })
+ return isReal
}
// TryAcquire 尝试立即获取串行锁
From 619e5ae6193703465eaa65a0ada3c045c5bacb2e Mon Sep 17 00:00:00 2001
From: name <136912576+is7Qin@users.noreply.github.com>
Date: Sat, 30 May 2026 02:40:52 +0800
Subject: [PATCH 50/79] refactor(gateway): isolate anthropic body rewrites
Keep Anthropic request body rewrites attempt-local and synchronize the accepted wire body only after upstream success so failover and retry paths do not reuse stale parsed state.
---
backend/internal/handler/gateway_handler.go | 41 ++--
...teway_anthropic_apikey_passthrough_test.go | 4 +-
...y_anthropic_vertex_service_account_test.go | 6 +-
.../gateway_context_management_test.go | 16 +-
.../gateway_forward_as_chat_completions.go | 2 +-
.../service/gateway_forward_as_responses.go | 2 +-
backend/internal/service/gateway_request.go | 153 +++++++-------
backend/internal/service/gateway_service.go | 194 ++++++++++++++----
8 files changed, 277 insertions(+), 141 deletions(-)
diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go
index 10d67ba3..c9b3297e 100644
--- a/backend/internal/handler/gateway_handler.go
+++ b/backend/internal/handler/gateway_handler.go
@@ -563,6 +563,12 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
retryWithFallback := false
for {
+ attemptParsedReq, err := parsedReq.CloneForBody(body)
+ if err != nil {
+ h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
+ return
+ }
+
// 选择支持该模型的账号
reqLog.Info("sticky.selecting_account",
zap.String("session_key", sessionKey),
@@ -694,7 +700,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
// ===== 用户消息串行队列 START =====
var queueRelease func()
- umqMode := h.getUserMsgQueueMode(account, parsedReq)
+ umqMode := h.getUserMsgQueueMode(account, attemptParsedReq)
switch umqMode {
case config.UMQModeSerialize:
@@ -741,20 +747,26 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
// 用 wrapReleaseOnDone 确保 context 取消时自动释放(仅 serialize 模式有 queueRelease)
queueRelease = wrapReleaseOnDone(c.Request.Context(), queueRelease)
// 注入回调到 ParsedRequest:使用外层 wrapper 以便提前清理 AfterFunc
- parsedReq.OnUpstreamAccepted = queueRelease
+ attemptParsedReq.OnUpstreamAccepted = queueRelease
// ===== 用户消息串行队列 END =====
- // 应用渠道模型映射到请求
+ // 渠道模型映射只作用于本次账号尝试,避免 failover 后污染原始 ParsedRequest。
if channelMapping.Mapped {
- parsedReq.Model = channelMapping.MappedModel
- parsedReq.ReplaceBody(h.gatewayService.ReplaceModelInBody(parsedReq.Body.Bytes(), channelMapping.MappedModel))
+ attemptParsedReq.Model = channelMapping.MappedModel
+ if err := attemptParsedReq.ReplaceBody(h.gatewayService.ReplaceModelInBody(attemptParsedReq.Body.Bytes(), channelMapping.MappedModel)); err != nil {
+ h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
+ return
+ }
}
// Bedrock CC 兼容:渠道模型映射后,清理 Anthropic API 专有字段、注入 Bedrock 必需字段
- parsedReq.ReplaceBody(h.gatewayService.ApplyBedrockCCCompat(c.Request.Context(), parsedReq.Body.Bytes(), parsedReq.Model, account, apiKey.GroupID))
- body = parsedReq.Body.Bytes()
+ if err := attemptParsedReq.ReplaceBody(h.gatewayService.ApplyBedrockCCCompat(c.Request.Context(), attemptParsedReq.Body.Bytes(), attemptParsedReq.Model, account, apiKey.GroupID)); err != nil {
+ h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
+ return
+ }
+ attemptBody := attemptParsedReq.Body.Bytes()
// 转发请求 - 根据账号平台分流
- c.Set("parsed_request", parsedReq)
+ c.Set("parsed_request", attemptParsedReq)
var result *service.ForwardResult
requestCtx := c.Request.Context()
if fs.SwitchCount > 0 {
@@ -763,9 +775,9 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
// 记录 Forward 前已写入字节数,Forward 后若增加则说明 SSE 内容已发,禁止 failover
writerSizeBeforeForward := c.Writer.Size()
if account.Platform == service.PlatformAntigravity && account.Type != service.AccountTypeAPIKey {
- result, err = h.antigravityGatewayService.Forward(requestCtx, c, account, body, hasBoundSession)
+ result, err = h.antigravityGatewayService.Forward(requestCtx, c, account, attemptBody, hasBoundSession)
} else {
- result, err = h.gatewayService.Forward(requestCtx, c, account, parsedReq)
+ result, err = h.gatewayService.Forward(requestCtx, c, account, attemptParsedReq)
}
// 兜底释放串行锁(正常情况已通过回调提前释放)
@@ -773,7 +785,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
queueRelease()
}
// 清理回调引用,防止 failover 重试时旧回调被错误调用
- parsedReq.OnUpstreamAccepted = nil
+ attemptParsedReq.OnUpstreamAccepted = nil
if accountReleaseFunc != nil {
accountReleaseFunc()
@@ -896,12 +908,13 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
// 捕获请求信息(用于异步记录,避免在 goroutine 中访问 gin.Context)
userAgent := c.GetHeader("User-Agent")
clientIP := ip.GetClientIP(c)
- requestPayloadHash := service.HashUsageRequestPayload(body)
+ // Forward 内部可能继续改写 body,usage 去重指纹必须使用最终上游接受的当前 body。
+ requestPayloadHash := service.HashUsageRequestPayload(attemptParsedReq.Body.Bytes())
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
if result.ReasoningEffort == nil {
- result.ReasoningEffort = service.NormalizeClaudeOutputEffort(parsedReq.OutputEffort)
+ result.ReasoningEffort = service.NormalizeClaudeOutputEffort(attemptParsedReq.OutputEffort)
}
// 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。
@@ -909,7 +922,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
Result: result,
- ParsedRequest: parsedReq,
+ ParsedRequest: attemptParsedReq,
QuotaPlatform: quotaPlatform,
APIKey: currentAPIKey,
User: currentAPIKey.User,
diff --git a/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go b/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go
index a67a3dc2..e2da89b5 100644
--- a/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go
+++ b/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go
@@ -713,7 +713,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_BuildRequestRejectsInvalidBas
},
}
- _, err := svc.buildUpstreamRequestAnthropicAPIKeyPassthrough(context.Background(), c, account, []byte(`{}`), "k")
+ _, _, err := svc.buildUpstreamRequestAnthropicAPIKeyPassthrough(context.Background(), c, account, []byte(`{}`), "k")
require.Error(t, err)
}
@@ -738,7 +738,7 @@ func TestGatewayService_AnthropicOAuth_NotAffectedByAPIKeyPassthroughToggle(t *t
require.False(t, account.IsAnthropicAPIKeyPassthroughEnabled())
- req, err := svc.buildUpstreamRequest(context.Background(), c, account, []byte(`{"model":"claude-3-7-sonnet-20250219"}`), "oauth-token", "oauth", "claude-3-7-sonnet-20250219", true, false)
+ req, _, err := svc.buildUpstreamRequest(context.Background(), c, account, []byte(`{"model":"claude-3-7-sonnet-20250219"}`), "oauth-token", "oauth", "claude-3-7-sonnet-20250219", true, false)
require.NoError(t, err)
require.Equal(t, "Bearer oauth-token", getHeaderRaw(req.Header, "authorization"))
require.Contains(t, getHeaderRaw(req.Header, "anthropic-beta"), claude.BetaOAuth, "OAuth 链路仍应按原逻辑补齐 oauth beta")
diff --git a/backend/internal/service/gateway_anthropic_vertex_service_account_test.go b/backend/internal/service/gateway_anthropic_vertex_service_account_test.go
index 2f42b0ab..be8c5867 100644
--- a/backend/internal/service/gateway_anthropic_vertex_service_account_test.go
+++ b/backend/internal/service/gateway_anthropic_vertex_service_account_test.go
@@ -35,7 +35,7 @@ func TestGatewayService_BuildAnthropicVertexServiceAccountRequest(t *testing.T)
body := []byte(`{"model":"claude-sonnet-4-5","stream":false,"max_tokens":32,"messages":[{"role":"user","content":"hello"}]}`)
svc := &GatewayService{}
- req, err := svc.buildUpstreamRequest(
+ req, _, err := svc.buildUpstreamRequest(
context.Background(),
c,
account,
@@ -87,7 +87,7 @@ func TestGatewayService_BuildAnthropicVertexServiceAccount_StripsContextManageme
body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015","keep":"all"}]},"messages":[{"role":"user","content":"hi"}]}`)
svc := &GatewayService{}
- req, err := svc.buildUpstreamRequest(
+ req, _, err := svc.buildUpstreamRequest(
context.Background(), c, account, body,
"vertex-token", "service_account", "claude-haiku-4-5@20251001", false, false,
)
@@ -117,7 +117,7 @@ func TestGatewayService_BuildAnthropicVertexServiceAccount_PreservesContextManag
body := []byte(`{"model":"claude-sonnet-4-6","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`)
svc := &GatewayService{}
- req, err := svc.buildUpstreamRequest(
+ req, _, err := svc.buildUpstreamRequest(
context.Background(), c, account, body,
"vertex-token", "service_account", "claude-sonnet-4-6@20260218", false, false,
)
diff --git a/backend/internal/service/gateway_context_management_test.go b/backend/internal/service/gateway_context_management_test.go
index c2263bdc..51b12809 100644
--- a/backend/internal/service/gateway_context_management_test.go
+++ b/backend/internal/service/gateway_context_management_test.go
@@ -364,7 +364,7 @@ func TestBuildUpstreamRequestAnthropicAPIKeyPassthrough_StripsContextManagementW
body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`)
svc := &GatewayService{cfg: &config.Config{}}
- req, err := svc.buildUpstreamRequestAnthropicAPIKeyPassthrough(
+ req, _, err := svc.buildUpstreamRequestAnthropicAPIKeyPassthrough(
context.Background(), c, newAnthropicAPIKeyPassthroughAccountForBetaTest(), body, "token",
)
require.NoError(t, err)
@@ -381,7 +381,7 @@ func TestBuildUpstreamRequestAnthropicAPIKeyPassthrough_PreservesContextManageme
body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`)
svc := &GatewayService{cfg: &config.Config{}}
- req, err := svc.buildUpstreamRequestAnthropicAPIKeyPassthrough(
+ req, _, err := svc.buildUpstreamRequestAnthropicAPIKeyPassthrough(
context.Background(), c, newAnthropicAPIKeyPassthroughAccountForBetaTest(), body, "token",
)
require.NoError(t, err)
@@ -427,7 +427,7 @@ func TestBuildUpstreamRequest_OAuthMimicHaiku_StripsContextManagementEndToEnd(t
// body 必须 strip。
body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`)
svc := &GatewayService{cfg: &config.Config{}}
- req, err := svc.buildUpstreamRequest(
+ req, _, err := svc.buildUpstreamRequest(
context.Background(), c, account, body,
"oauth-tok", "oauth", "claude-haiku-4-5", false, true, // mimicClaudeCode=true
)
@@ -457,7 +457,7 @@ func TestBuildUpstreamRequest_OAuthMimicNonHaiku_PreservesContextManagementEndTo
// body 保留。
body := []byte(`{"model":"claude-sonnet-4-6","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`)
svc := &GatewayService{cfg: &config.Config{}}
- req, err := svc.buildUpstreamRequest(
+ req, _, err := svc.buildUpstreamRequest(
context.Background(), c, account, body,
"oauth-tok", "oauth", "claude-sonnet-4-6", false, true,
)
@@ -488,7 +488,7 @@ func TestBuildUpstreamRequest_OAuthTransparentHaikuWithRealCCBeta_PreservesField
}
body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015","keep":"all"}]},"messages":[]}`)
svc := &GatewayService{cfg: &config.Config{}}
- req, err := svc.buildUpstreamRequest(
+ req, _, err := svc.buildUpstreamRequest(
context.Background(), c, account, body,
"oauth-tok", "oauth", "claude-haiku-4-5", false, false, // mimicClaudeCode=false(真 CC)
)
@@ -580,7 +580,7 @@ func TestBuildCountTokensRequest_OAuthMimicHaiku_PreservesContextManagementEndTo
}
body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`)
svc := &GatewayService{cfg: &config.Config{}}
- req, err := svc.buildCountTokensRequest(
+ req, _, err := svc.buildCountTokensRequest(
context.Background(), c, account, body,
"oauth-tok", "oauth", "claude-haiku-4-5", true, // mimicClaudeCode=true
)
@@ -611,7 +611,7 @@ func TestBuildCountTokensRequest_APIKeyHaiku_StripsContextManagementEndToEnd(t *
}
body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[]},"messages":[]}`)
svc := &GatewayService{cfg: &config.Config{}}
- req, err := svc.buildCountTokensRequest(
+ req, _, err := svc.buildCountTokensRequest(
context.Background(), c, account, body,
"sk-ant-xxx", "apikey", "claude-haiku-4-5", false,
)
@@ -655,7 +655,7 @@ func TestBuildUpstreamRequest_APIKeyHaikuWithContextManagement_StripsField(t *te
}
body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[]},"messages":[]}`)
svc := &GatewayService{cfg: &config.Config{}}
- req, err := svc.buildUpstreamRequest(
+ req, _, err := svc.buildUpstreamRequest(
context.Background(), c, account, body,
"sk-ant-xxx", "apikey", "claude-haiku-4-5", false, false,
)
diff --git a/backend/internal/service/gateway_forward_as_chat_completions.go b/backend/internal/service/gateway_forward_as_chat_completions.go
index eaf67fab..729483e3 100644
--- a/backend/internal/service/gateway_forward_as_chat_completions.go
+++ b/backend/internal/service/gateway_forward_as_chat_completions.go
@@ -119,7 +119,7 @@ func (s *GatewayService) ForwardAsChatCompletions(
// 10. Build upstream request
upstreamCtx, releaseUpstreamCtx := detachStreamUpstreamContext(ctx, reqStream)
- upstreamReq, err := s.buildUpstreamRequest(upstreamCtx, c, account, anthropicBody, token, tokenType, mappedModel, reqStream, shouldMimicClaudeCode)
+ upstreamReq, _, err := s.buildUpstreamRequest(upstreamCtx, c, account, anthropicBody, token, tokenType, mappedModel, reqStream, shouldMimicClaudeCode)
releaseUpstreamCtx()
if err != nil {
return nil, fmt.Errorf("build upstream request: %w", err)
diff --git a/backend/internal/service/gateway_forward_as_responses.go b/backend/internal/service/gateway_forward_as_responses.go
index c55a5a98..3baea018 100644
--- a/backend/internal/service/gateway_forward_as_responses.go
+++ b/backend/internal/service/gateway_forward_as_responses.go
@@ -116,7 +116,7 @@ func (s *GatewayService) ForwardAsResponses(
// 10. Build upstream request
upstreamCtx, releaseUpstreamCtx := detachStreamUpstreamContext(ctx, reqStream)
- upstreamReq, err := s.buildUpstreamRequest(upstreamCtx, c, account, anthropicBody, token, tokenType, mappedModel, reqStream, shouldMimicClaudeCode)
+ upstreamReq, _, err := s.buildUpstreamRequest(upstreamCtx, c, account, anthropicBody, token, tokenType, mappedModel, reqStream, shouldMimicClaudeCode)
releaseUpstreamCtx()
if err != nil {
return nil, fmt.Errorf("build upstream request: %w", err)
diff --git a/backend/internal/service/gateway_request.go b/backend/internal/service/gateway_request.go
index 65718c27..dc59611c 100644
--- a/backend/internal/service/gateway_request.go
+++ b/backend/internal/service/gateway_request.go
@@ -105,6 +105,22 @@ func (r jsonRange) exists() bool {
return r.start >= 0 && r.end >= r.start
}
+// clearGatewayRequestDerivedState 清空绑定当前 body 的轻量派生字段,防止 ReplaceBody 后读到旧值。
+func clearGatewayRequestDerivedState(parsed *ParsedRequest) {
+ if parsed == nil {
+ return
+ }
+ parsed.Model = ""
+ parsed.Stream = false
+ parsed.MetadataUserID = ""
+ parsed.HasSystem = false
+ parsed.ThinkingEnabled = false
+ parsed.OutputEffort = ""
+ parsed.MaxTokens = 0
+ parsed.systemRange = missingJSONRange()
+ parsed.messagesRange = missingJSONRange()
+}
+
func clearGatewayRequestRanges(parsed *ParsedRequest) {
if parsed == nil {
return
@@ -137,12 +153,9 @@ func setGatewayRequestRanges(parsed *ParsedRequest, protocol string, jsonStr str
}
}
-func refreshGatewayRequestRanges(parsed *ParsedRequest, protocol string) error {
- if parsed == nil {
- return fmt.Errorf("empty request body")
- }
- clearGatewayRequestRanges(parsed)
- if parsed.Body == nil {
+// parseGatewayRequestCurrentBody 只做标量和 raw range 轻量解析,不恢复 system/messages 对象图。
+func parseGatewayRequestCurrentBody(parsed *ParsedRequest, protocol string) error {
+ if parsed == nil || parsed.Body == nil {
return fmt.Errorf("empty request body")
}
@@ -151,11 +164,51 @@ func refreshGatewayRequestRanges(parsed *ParsedRequest, protocol string) error {
return fmt.Errorf("invalid json")
}
+ // 只在当前函数内零拷贝读取 JSON 字段;ReplaceBody 后必须重新进入本函数刷新派生状态。
jsonStr := *(*string)(unsafe.Pointer(&bodyBytes))
+ clearGatewayRequestDerivedState(parsed)
+ parsed.protocol = protocol
+
+ modelResult := gjson.Get(jsonStr, "model")
+ if modelResult.Exists() {
+ if modelResult.Type != gjson.String {
+ return fmt.Errorf("invalid model field type")
+ }
+ parsed.Model = modelResult.String()
+ }
+
+ streamResult := gjson.Get(jsonStr, "stream")
+ if streamResult.Exists() {
+ if streamResult.Type != gjson.True && streamResult.Type != gjson.False {
+ return fmt.Errorf("invalid stream field type")
+ }
+ parsed.Stream = streamResult.Bool()
+ }
+
+ parsed.MetadataUserID = gjson.Get(jsonStr, "metadata.user_id").String()
+
+ thinkingType := gjson.Get(jsonStr, "thinking.type").String()
+ parsed.ThinkingEnabled = thinkingType == "enabled" || thinkingType == "adaptive"
+
+ parsed.OutputEffort = strings.TrimSpace(gjson.Get(jsonStr, "output_config.effort").String())
+
+ maxTokensResult := gjson.Get(jsonStr, "max_tokens")
+ if maxTokensResult.Exists() && maxTokensResult.Type == gjson.Number {
+ f := maxTokensResult.Float()
+ if !math.IsNaN(f) && !math.IsInf(f, 0) && f == math.Trunc(f) &&
+ f <= float64(math.MaxInt) && f >= float64(math.MinInt) {
+ parsed.MaxTokens = int(f)
+ }
+ }
+
setGatewayRequestRanges(parsed, protocol, jsonStr)
return nil
}
+func refreshGatewayRequestRanges(parsed *ParsedRequest, protocol string) error {
+ return parseGatewayRequestCurrentBody(parsed, protocol)
+}
+
// ParsedRequest 保存网关请求的预解析结果
//
// 性能优化说明:
@@ -238,71 +291,10 @@ func normalizeSessionUserAgentFallback(raw string) string {
// protocol 指定请求协议格式(domain.PlatformAnthropic / domain.PlatformGemini),
// 不同协议使用不同的 system/messages 字段名。
func ParseGatewayRequest(body *RequestBodyRef, protocol string) (*ParsedRequest, error) {
- bodyBytes := body.Bytes()
- // 保持与旧实现一致:请求体必须是合法 JSON。
- // 注意:gjson.GetBytes 对非法 JSON 不会报错,因此需要显式校验。
- if !gjson.ValidBytes(bodyBytes) {
- return nil, fmt.Errorf("invalid json")
+ parsed := &ParsedRequest{Body: body}
+ if err := parseGatewayRequestCurrentBody(parsed, protocol); err != nil {
+ return nil, err
}
-
- // 性能:
- // - gjson.GetBytes 会把匹配的 Raw/Str 安全复制成 string(对于巨大 messages 会产生额外拷贝)。
- // - 这里将 body 通过 unsafe 零拷贝视为 string,仅在本函数内使用,且 body 不会被修改。
- jsonStr := *(*string)(unsafe.Pointer(&bodyBytes))
-
- parsed := &ParsedRequest{
- Body: body,
- protocol: protocol,
- systemRange: missingJSONRange(),
- messagesRange: missingJSONRange(),
- }
-
- // --- gjson 提取简单字段(避免完整 Unmarshal) ---
-
- // model: 需要严格类型校验,非 string 返回错误
- modelResult := gjson.Get(jsonStr, "model")
- if modelResult.Exists() {
- if modelResult.Type != gjson.String {
- return nil, fmt.Errorf("invalid model field type")
- }
- parsed.Model = modelResult.String()
- }
-
- // stream: 需要严格类型校验,非 bool 返回错误
- streamResult := gjson.Get(jsonStr, "stream")
- if streamResult.Exists() {
- if streamResult.Type != gjson.True && streamResult.Type != gjson.False {
- return nil, fmt.Errorf("invalid stream field type")
- }
- parsed.Stream = streamResult.Bool()
- }
-
- // metadata.user_id: 直接路径提取,不需要严格类型校验
- parsed.MetadataUserID = gjson.Get(jsonStr, "metadata.user_id").String()
-
- // thinking.type: enabled/adaptive 都视为开启
- thinkingType := gjson.Get(jsonStr, "thinking.type").String()
- if thinkingType == "enabled" || thinkingType == "adaptive" {
- parsed.ThinkingEnabled = true
- }
-
- // output_config.effort: Claude API 的推理强度控制参数
- parsed.OutputEffort = strings.TrimSpace(gjson.Get(jsonStr, "output_config.effort").String())
-
- // max_tokens: 仅接受整数值
- maxTokensResult := gjson.Get(jsonStr, "max_tokens")
- if maxTokensResult.Exists() && maxTokensResult.Type == gjson.Number {
- f := maxTokensResult.Float()
- if !math.IsNaN(f) && !math.IsInf(f, 0) && f == math.Trunc(f) &&
- f <= float64(math.MaxInt) && f >= float64(math.MinInt) {
- parsed.MaxTokens = int(f)
- }
- }
-
- // --- system/messages 提取 ---
- // 只保存大字段 raw range,不默认反序列化成 []any/map[string]any 对象图。
- setGatewayRequestRanges(parsed, protocol, jsonStr)
-
return parsed, nil
}
@@ -353,9 +345,24 @@ func (p *ParsedRequest) SystemValue() (any, bool) {
return system, true
}
-func (p *ParsedRequest) ReplaceBody(data []byte) {
+// CloneForBody 为单次账号尝试创建独立 body 视图,避免 failover 复用已改写的 ParsedRequest。
+func (p *ParsedRequest) CloneForBody(body []byte) (*ParsedRequest, error) {
if p == nil {
- return
+ return nil, fmt.Errorf("parse request: empty request")
+ }
+ clone := *p
+ clone.Body = NewRequestBodyRef(body)
+ clone.OnUpstreamAccepted = nil
+ if err := refreshGatewayRequestRanges(&clone, clone.protocol); err != nil {
+ return nil, err
+ }
+ return &clone, nil
+}
+
+// ReplaceBody 统一刷新当前 body 和 raw range,保证后续 helper 读取的是最新请求体。
+func (p *ParsedRequest) ReplaceBody(data []byte) error {
+ if p == nil {
+ return fmt.Errorf("parse request: empty request")
}
if p.Body == nil {
p.Body = NewRequestBodyRef(data)
@@ -364,7 +371,9 @@ func (p *ParsedRequest) ReplaceBody(data []byte) {
}
if err := refreshGatewayRequestRanges(p, p.protocol); err != nil {
clearGatewayRequestRanges(p)
+ return err
}
+ return nil
}
// sliceRawFromBody 返回 Result.Raw 对应的原始字节切片。
diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go
index 438592e3..01b541e2 100644
--- a/backend/internal/service/gateway_service.go
+++ b/backend/internal/service/gateway_service.go
@@ -4440,6 +4440,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
}
return s.forwardAnthropicAPIKeyPassthroughWithInput(ctx, c, account, anthropicPassthroughForwardInput{
Body: passthroughBody,
+ Parsed: parsed,
RequestModel: passthroughModel,
OriginalModel: parsed.Model,
RequestStream: parsed.Stream,
@@ -4466,6 +4467,13 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
}
body := parsed.Body.Bytes()
+ replaceBody := func(next []byte) error {
+ if err := parsed.ReplaceBody(next); err != nil {
+ return fmt.Errorf("rewrite request body: %w", err)
+ }
+ body = parsed.Body.Bytes()
+ return nil
+ }
reqModel := parsed.Model
reqStream := parsed.Stream
originalModel := reqModel
@@ -4499,7 +4507,9 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
systemRewritten := false
if !strings.Contains(strings.ToLower(reqModel), "haiku") {
systemRaw, _ := parsed.SystemValue()
- body = rewriteSystemForNonClaudeCode(body, systemRaw)
+ if err := replaceBody(rewriteSystemForNonClaudeCode(body, systemRaw)); err != nil {
+ return nil, err
+ }
systemRewritten = true
}
@@ -4521,22 +4531,34 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
}
}
- body, reqModel = normalizeClaudeOAuthRequestBody(body, reqModel, normalizeOpts)
+ var normalizedBody []byte
+ normalizedBody, reqModel = normalizeClaudeOAuthRequestBody(body, reqModel, normalizeOpts)
+ if err := replaceBody(normalizedBody); err != nil {
+ return nil, err
+ }
// D/E/F: 可选 messages cache 策略 + 工具名混淆 + tools[-1] 断点
// 与 forward_as_chat_completions / forward_as_responses 路径对齐,
// 原生 /v1/messages 路径也走同一套可配置字段级改写。
- body = s.rewriteMessageCacheControlIfEnabled(ctx, body)
+ if err := replaceBody(s.rewriteMessageCacheControlIfEnabled(ctx, body)); err != nil {
+ return nil, err
+ }
if rw := buildToolNameRewriteFromBody(body); rw != nil {
- body = applyToolNameRewriteToBody(body, rw)
+ if err := replaceBody(applyToolNameRewriteToBody(body, rw)); err != nil {
+ return nil, err
+ }
c.Set(toolNameRewriteKey, rw)
} else {
- body = applyToolsLastCacheBreakpoint(body)
+ if err := replaceBody(applyToolsLastCacheBreakpoint(body)); err != nil {
+ return nil, err
+ }
}
}
// 强制执行 cache_control 块数量限制(最多 4 个)
- body = enforceCacheControlLimit(body)
+ if err := replaceBody(enforceCacheControlLimit(body)); err != nil {
+ return nil, err
+ }
// 应用模型映射:
// - APIKey 账号:使用账号级别的显式映射(如果配置),否则透传原始模型名
@@ -4570,13 +4592,18 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
}
if mappedModel != reqModel {
// 替换请求体中的模型名
- body = s.replaceModelInBody(body, mappedModel)
+ if err := replaceBody(s.replaceModelInBody(body, mappedModel)); err != nil {
+ return nil, err
+ }
reqModel = mappedModel
+ parsed.Model = mappedModel
logger.LegacyPrintf("service.gateway", "Model mapping applied: %s -> %s (account: %s, source=%s)", originalModel, mappedModel, account.Name, mappingSource)
}
if s.shouldInjectAnthropicCacheTTL1h(ctx, account) {
- body = injectAnthropicCacheControlTTL1h(body)
+ if err := replaceBody(injectAnthropicCacheControlTTL1h(body)); err != nil {
+ return nil, err
+ }
}
// 获取凭证
@@ -4600,19 +4627,24 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
logger.LegacyPrintf("service.gateway", "[Forward] Using account: ID=%d Name=%s Platform=%s Type=%s TLSFingerprint=%v Proxy=%s",
account.ID, account.Name, account.Platform, account.Type, tlsProfile, proxyURL)
// Pre-filter: strip empty text blocks (including nested in tool_result) to prevent upstream 400.
- body = StripEmptyTextBlocks(body)
+ if err := replaceBody(StripEmptyTextBlocks(body)); err != nil {
+ return nil, err
+ }
// 重试循环
var resp *http.Response
+ lastWireBody := body
retryStart := time.Now()
for attempt := 1; attempt <= maxRetryAttempts; attempt++ {
// 构建上游请求(每次重试需要重新构建,因为请求体需要重新读取)
upstreamCtx, releaseUpstreamCtx := detachStreamUpstreamContext(ctx, reqStream)
- upstreamReq, err := s.buildUpstreamRequest(upstreamCtx, c, account, body, token, tokenType, reqModel, reqStream, shouldMimicClaudeCode)
+ upstreamReq, wireBody, err := s.buildUpstreamRequest(upstreamCtx, c, account, body, token, tokenType, reqModel, reqStream, shouldMimicClaudeCode)
releaseUpstreamCtx()
if err != nil {
return nil, err
}
+ // 记录本次实际发送的 wire body;只有请求成功后才写回 ParsedRequest,避免 400 retry 基于已签名 CCH 再改写。
+ lastWireBody = wireBody
// 发送请求
resp, err = s.httpUpstream.DoWithTLS(upstreamReq, proxyURL, account.ID, account.Concurrency, tlsProfile)
@@ -4690,12 +4722,17 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
filteredBody := FilterThinkingBlocksForRetry(body)
retryCtx, releaseRetryCtx := detachStreamUpstreamContext(ctx, reqStream)
- retryReq, buildErr := s.buildUpstreamRequest(retryCtx, c, account, filteredBody, token, tokenType, reqModel, reqStream, shouldMimicClaudeCode)
+ retryReq, retryWireBody, buildErr := s.buildUpstreamRequest(retryCtx, c, account, filteredBody, token, tokenType, reqModel, reqStream, shouldMimicClaudeCode)
releaseRetryCtx()
if buildErr == nil {
retryResp, retryErr := s.httpUpstream.DoWithTLS(retryReq, proxyURL, account.ID, account.Concurrency, tlsProfile)
if retryErr == nil {
if retryResp.StatusCode < 400 {
+ // 重试请求被上游接受后同步 ParsedRequest,保证 usage/日志看到真实请求体。
+ if err := replaceBody(retryWireBody); err != nil {
+ _ = retryResp.Body.Close()
+ return nil, err
+ }
logger.LegacyPrintf("service.gateway", "Account %d: thinking block retry succeeded (blocks downgraded)", account.ID)
resp = retryResp
break
@@ -4725,11 +4762,18 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
logger.LegacyPrintf("service.gateway", "Account %d: signature retry still failing and looks tool-related, retrying with tool blocks downgraded", account.ID)
filteredBody2 := FilterSignatureSensitiveBlocksForRetry(body)
retryCtx2, releaseRetryCtx2 := detachStreamUpstreamContext(ctx, reqStream)
- retryReq2, buildErr2 := s.buildUpstreamRequest(retryCtx2, c, account, filteredBody2, token, tokenType, reqModel, reqStream, shouldMimicClaudeCode)
+ retryReq2, retryWireBody2, buildErr2 := s.buildUpstreamRequest(retryCtx2, c, account, filteredBody2, token, tokenType, reqModel, reqStream, shouldMimicClaudeCode)
releaseRetryCtx2()
if buildErr2 == nil {
retryResp2, retryErr2 := s.httpUpstream.DoWithTLS(retryReq2, proxyURL, account.ID, account.Concurrency, tlsProfile)
if retryErr2 == nil {
+ if retryResp2.StatusCode < 400 {
+ // 二阶段工具块降级成功时也必须更新当前 body。
+ if err := replaceBody(retryWireBody2); err != nil {
+ _ = retryResp2.Body.Close()
+ return nil, err
+ }
+ }
resp = retryResp2
break
}
@@ -4796,11 +4840,18 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
if applied && time.Since(retryStart) < maxRetryElapsed {
logger.LegacyPrintf("service.gateway", "Account %d: detected budget_tokens constraint error, retrying with rectified budget (budget_tokens=%d, max_tokens=%d)", account.ID, BudgetRectifyBudgetTokens, BudgetRectifyMaxTokens)
budgetRetryCtx, releaseBudgetRetryCtx := detachStreamUpstreamContext(ctx, reqStream)
- budgetRetryReq, buildErr := s.buildUpstreamRequest(budgetRetryCtx, c, account, rectifiedBody, token, tokenType, reqModel, reqStream, shouldMimicClaudeCode)
+ budgetRetryReq, budgetWireBody, buildErr := s.buildUpstreamRequest(budgetRetryCtx, c, account, rectifiedBody, token, tokenType, reqModel, reqStream, shouldMimicClaudeCode)
releaseBudgetRetryCtx()
if buildErr == nil {
budgetRetryResp, retryErr := s.httpUpstream.DoWithTLS(budgetRetryReq, proxyURL, account.ID, account.Concurrency, tlsProfile)
if retryErr == nil {
+ if budgetRetryResp.StatusCode < 400 {
+ // budget 修正请求成功后,ParsedRequest 也要描述被接受的修正版。
+ if err := replaceBody(budgetWireBody); err != nil {
+ _ = budgetRetryResp.Body.Close()
+ return nil, err
+ }
+ }
resp = budgetRetryResp
break
}
@@ -4997,6 +5048,13 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
// 处理正常响应
+ if !bytes.Equal(lastWireBody, body) {
+ // 成功后再同步最终 wire body,避免失败重试从已签名 CCH 的 body 继续派生。
+ if err := replaceBody(lastWireBody); err != nil {
+ return nil, err
+ }
+ }
+
// 触发上游接受回调(提前释放串行锁,不等流完成)
if parsed.OnUpstreamAccepted != nil {
parsed.OnUpstreamAccepted()
@@ -5039,6 +5097,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
type anthropicPassthroughForwardInput struct {
Body []byte
+ Parsed *ParsedRequest
RequestModel string
OriginalModel string
RequestStream bool
@@ -5091,16 +5150,29 @@ func (s *GatewayService) forwardAnthropicAPIKeyPassthroughWithInput(
}
// Pre-filter: strip empty text blocks (including nested in tool_result) to prevent upstream 400.
input.Body = StripEmptyTextBlocks(input.Body)
+ if input.Parsed != nil {
+ // 透传分支也会改写实际 wire body,成功 usage hash 依赖这里同步当前 body。
+ if err := input.Parsed.ReplaceBody(input.Body); err != nil {
+ return nil, err
+ }
+ }
var resp *http.Response
retryStart := time.Now()
for attempt := 1; attempt <= maxRetryAttempts; attempt++ {
upstreamCtx, releaseUpstreamCtx := detachStreamUpstreamContext(ctx, input.RequestStream)
- upstreamReq, err := s.buildUpstreamRequestAnthropicAPIKeyPassthrough(upstreamCtx, c, account, input.Body, token)
+ upstreamReq, wireBody, err := s.buildUpstreamRequestAnthropicAPIKeyPassthrough(upstreamCtx, c, account, input.Body, token)
releaseUpstreamCtx()
if err != nil {
return nil, err
}
+ if input.Parsed != nil && !bytes.Equal(wireBody, input.Body) {
+ // build 阶段会按 beta 能力清理 body,发送前同步到 ParsedRequest 当前视图。
+ if err := input.Parsed.ReplaceBody(wireBody); err != nil {
+ return nil, err
+ }
+ input.Body = input.Parsed.Body.Bytes()
+ }
resp, err = s.httpUpstream.DoWithTLS(upstreamReq, proxyURL, account.ID, account.Concurrency, s.tlsFPProfileService.ResolveTLSProfile(account))
if err != nil {
@@ -5292,13 +5364,13 @@ func (s *GatewayService) buildUpstreamRequestAnthropicAPIKeyPassthrough(
account *Account,
body []byte,
token string,
-) (*http.Request, error) {
+) (*http.Request, []byte, error) {
targetURL := claudeAPIURL
baseURL := account.GetBaseURL()
if baseURL != "" {
validatedURL, err := s.validateUpstreamBaseURL(baseURL)
if err != nil {
- return nil, err
+ return nil, nil, err
}
targetURL = validatedURL + "/v1/messages?beta=true"
}
@@ -5316,7 +5388,7 @@ func (s *GatewayService) buildUpstreamRequestAnthropicAPIKeyPassthrough(
req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body))
if err != nil {
- return nil, err
+ return nil, nil, err
}
if c != nil && c.Request != nil {
@@ -5346,7 +5418,7 @@ func (s *GatewayService) buildUpstreamRequestAnthropicAPIKeyPassthrough(
setHeaderRaw(req.Header, "anthropic-version", "2023-06-01")
}
- return req, nil
+ return req, body, nil
}
func (s *GatewayService) handleStreamingResponseAnthropicAPIKeyPassthrough(
@@ -5828,6 +5900,11 @@ func (s *GatewayService) forwardBedrock(
return s.handleBedrockUpstreamErrors(ctx, resp, c, account)
}
+ // Bedrock 分支绕过通用 Forward 成功路径,这里保持上游接受回调语义一致。
+ if parsed.OnUpstreamAccepted != nil {
+ parsed.OnUpstreamAccepted()
+ }
+
// 响应处理
var usage *ClaudeUsage
var firstTokenMs *int
@@ -6104,9 +6181,10 @@ func (s *GatewayService) handleBedrockNonStreamingResponse(
return usage, nil
}
-func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token, tokenType, modelID string, reqStream bool, mimicClaudeCode bool) (*http.Request, error) {
+func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token, tokenType, modelID string, reqStream bool, mimicClaudeCode bool) (*http.Request, []byte, error) {
if account.Platform == PlatformAnthropic && account.Type == AccountTypeServiceAccount {
- return s.buildUpstreamRequestAnthropicVertex(ctx, c, account, body, token, modelID, reqStream)
+ req, err := s.buildUpstreamRequestAnthropicVertex(ctx, c, account, body, token, modelID, reqStream)
+ return req, body, err
}
// 确定目标URL
@@ -6116,18 +6194,18 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex
if baseURL != "" {
validatedURL, err := s.validateUpstreamBaseURL(baseURL)
if err != nil {
- return nil, err
+ return nil, nil, err
}
targetURL = validatedURL + "/v1/messages?beta=true"
}
} else if account.IsCustomBaseURLEnabled() {
customURL := account.GetCustomBaseURL()
if customURL == "" {
- return nil, fmt.Errorf("custom_base_url is enabled but not configured for account %d", account.ID)
+ return nil, nil, fmt.Errorf("custom_base_url is enabled but not configured for account %d", account.ID)
}
validatedURL, err := s.validateUpstreamBaseURL(customURL)
if err != nil {
- return nil, err
+ return nil, nil, err
}
targetURL = s.buildCustomRelayURL(validatedURL, "/v1/messages", account)
}
@@ -6202,7 +6280,7 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex
req, err := http.NewRequestWithContext(ctx, "POST", targetURL, bytes.NewReader(body))
if err != nil {
- return nil, err
+ return nil, nil, err
}
// 设置认证头(保持原始大小写)
@@ -6287,7 +6365,7 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex
logClaudeMimicDebug(req, body, account, tokenType, mimicClaudeCode)
}
- return req, nil
+ return req, body, nil
}
func (s *GatewayService) buildUpstreamRequestAnthropicVertex(
@@ -9214,23 +9292,42 @@ func (s *GatewayService) ForwardCountTokens(ctx context.Context, c *gin.Context,
}
body := parsed.Body.Bytes()
+ replaceBody := func(next []byte) error {
+ if err := parsed.ReplaceBody(next); err != nil {
+ return fmt.Errorf("rewrite count_tokens body: %w", err)
+ }
+ body = parsed.Body.Bytes()
+ return nil
+ }
reqModel := parsed.Model
// Pre-filter: strip empty text blocks to prevent upstream 400.
- body = StripEmptyTextBlocks(body)
+ if err := replaceBody(StripEmptyTextBlocks(body)); err != nil {
+ return err
+ }
isClaudeCodeCT := IsClaudeCodeClient(ctx) || isClaudeCodeClient(c.GetHeader("User-Agent"), parsed.MetadataUserID)
shouldMimicClaudeCode := account.IsOAuth() && !isClaudeCodeCT
if shouldMimicClaudeCode {
normalizeOpts := claudeOAuthNormalizeOptions{stripSystemCacheControl: true}
- body, reqModel = normalizeClaudeOAuthRequestBody(body, reqModel, normalizeOpts)
+ var normalizedBody []byte
+ normalizedBody, reqModel = normalizeClaudeOAuthRequestBody(body, reqModel, normalizeOpts)
+ if err := replaceBody(normalizedBody); err != nil {
+ return err
+ }
- body = s.rewriteMessageCacheControlIfEnabled(ctx, body)
+ if err := replaceBody(s.rewriteMessageCacheControlIfEnabled(ctx, body)); err != nil {
+ return err
+ }
if rw := buildToolNameRewriteFromBody(body); rw != nil {
- body = applyToolNameRewriteToBody(body, rw)
+ if err := replaceBody(applyToolNameRewriteToBody(body, rw)); err != nil {
+ return err
+ }
} else {
- body = applyToolsLastCacheBreakpoint(body)
+ if err := replaceBody(applyToolsLastCacheBreakpoint(body)); err != nil {
+ return err
+ }
}
}
@@ -9261,9 +9358,13 @@ func (s *GatewayService) ForwardCountTokens(ctx context.Context, c *gin.Context,
}
}
if mappedModel != reqModel {
- body = s.replaceModelInBody(body, mappedModel)
+ originalReqModel := reqModel
+ if err := replaceBody(s.replaceModelInBody(body, mappedModel)); err != nil {
+ return err
+ }
reqModel = mappedModel
- logger.LegacyPrintf("service.gateway", "CountTokens model mapping applied: %s -> %s (account: %s, source=%s)", parsed.Model, mappedModel, account.Name, mappingSource)
+ parsed.Model = mappedModel
+ logger.LegacyPrintf("service.gateway", "CountTokens model mapping applied: %s -> %s (account: %s, source=%s)", originalReqModel, mappedModel, account.Name, mappingSource)
}
}
@@ -9275,11 +9376,13 @@ func (s *GatewayService) ForwardCountTokens(ctx context.Context, c *gin.Context,
}
// 构建上游请求
- upstreamReq, err := s.buildCountTokensRequest(ctx, c, account, body, token, tokenType, reqModel, shouldMimicClaudeCode)
+ upstreamReq, wireBody, err := s.buildCountTokensRequest(ctx, c, account, body, token, tokenType, reqModel, shouldMimicClaudeCode)
if err != nil {
s.countTokensError(c, http.StatusInternalServerError, "api_error", "Failed to build request")
return err
}
+ // 先记录首发 wire body;如果后面进入 400 retry,retry 会基于未签名的逻辑 body 重新构建。
+ acceptedWireBody := wireBody
// 获取代理URL(自定义 base URL 模式下,proxy 通过 buildCustomRelayURL 作为查询参数传递)
proxyURL := ""
@@ -9315,10 +9418,14 @@ func (s *GatewayService) ForwardCountTokens(ctx context.Context, c *gin.Context,
logger.LegacyPrintf("service.gateway", "Account %d: detected thinking block signature error on count_tokens, retrying with filtered thinking blocks", account.ID)
filteredBody := FilterThinkingBlocksForRetry(body)
- retryReq, buildErr := s.buildCountTokensRequest(ctx, c, account, filteredBody, token, tokenType, reqModel, shouldMimicClaudeCode)
+ retryReq, retryWireBody, buildErr := s.buildCountTokensRequest(ctx, c, account, filteredBody, token, tokenType, reqModel, shouldMimicClaudeCode)
if buildErr == nil {
retryResp, retryErr := s.httpUpstream.DoWithTLS(retryReq, proxyURL, account.ID, account.Concurrency, s.tlsFPProfileService.ResolveTLSProfile(account))
if retryErr == nil {
+ if retryResp.StatusCode < 400 {
+ // count_tokens 签名重试成功后记录最终 wire body,错误响应仍保留原 body 便于后续处理。
+ acceptedWireBody = retryWireBody
+ }
resp = retryResp
respBody, err = ReadUpstreamResponseBody(resp.Body, s.cfg, c, countTokensTooLarge)
_ = resp.Body.Close()
@@ -9332,6 +9439,13 @@ func (s *GatewayService) ForwardCountTokens(ctx context.Context, c *gin.Context,
}
}
+ if resp.StatusCode < 400 && !bytes.Equal(acceptedWireBody, body) {
+ // count_tokens 成功后再同步最终 wire body,避免 retry 从已签名 body 派生。
+ if err := replaceBody(acceptedWireBody); err != nil {
+ return err
+ }
+ }
+
// 处理错误响应
if resp.StatusCode >= 400 {
// 标记账号状态(429/529等)
@@ -9558,7 +9672,7 @@ func (s *GatewayService) buildCountTokensRequestAnthropicAPIKeyPassthrough(
}
// buildCountTokensRequest 构建 count_tokens 上游请求
-func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token, tokenType, modelID string, mimicClaudeCode bool) (*http.Request, error) {
+func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token, tokenType, modelID string, mimicClaudeCode bool) (*http.Request, []byte, error) {
// 确定目标 URL
targetURL := claudeAPICountTokensURL
if account.Type == AccountTypeAPIKey {
@@ -9566,18 +9680,18 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con
if baseURL != "" {
validatedURL, err := s.validateUpstreamBaseURL(baseURL)
if err != nil {
- return nil, err
+ return nil, nil, err
}
targetURL = validatedURL + "/v1/messages/count_tokens?beta=true"
}
} else if account.IsCustomBaseURLEnabled() {
customURL := account.GetCustomBaseURL()
if customURL == "" {
- return nil, fmt.Errorf("custom_base_url is enabled but not configured for account %d", account.ID)
+ return nil, nil, fmt.Errorf("custom_base_url is enabled but not configured for account %d", account.ID)
}
validatedURL, err := s.validateUpstreamBaseURL(customURL)
if err != nil {
- return nil, err
+ return nil, nil, err
}
targetURL = s.buildCustomRelayURL(validatedURL, "/v1/messages/count_tokens", account)
}
@@ -9633,7 +9747,7 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con
req, err := http.NewRequestWithContext(ctx, "POST", targetURL, bytes.NewReader(body))
if err != nil {
- return nil, err
+ return nil, nil, err
}
// 设置认证头(保持原始大小写)
@@ -9697,7 +9811,7 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con
logClaudeMimicDebug(req, body, account, tokenType, mimicClaudeCode)
}
- return req, nil
+ return req, body, nil
}
func sanitizeCountTokensRequestBody(body []byte) []byte {
From 2caee9d884224c43e5c2af1c992ccde99b51d2c7 Mon Sep 17 00:00:00 2001
From: name <136912576+is7Qin@users.noreply.github.com>
Date: Sat, 30 May 2026 15:52:33 +0800
Subject: [PATCH 51/79] refactor(gateway): snapshot usage worker inputs
---
backend/internal/handler/gateway_handler.go | 10 ++++----
.../internal/handler/gemini_v1beta_handler.go | 4 +++-
.../handler/openai_gateway_handler.go | 6 ++++-
backend/internal/handler/openai_images.go | 17 ++++++-------
backend/internal/service/gateway_service.go | 24 +++++++------------
5 files changed, 31 insertions(+), 30 deletions(-)
diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go
index c9b3297e..3bd5e82e 100644
--- a/backend/internal/handler/gateway_handler.go
+++ b/backend/internal/handler/gateway_handler.go
@@ -510,11 +510,12 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
}
// 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。
+ // ForceCacheBilling 提前拍成标量,避免 worker 闭包保活 failover 状态里的响应体。
+ forceCacheBilling := fs.ForceCacheBilling
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
Result: result,
- ParsedRequest: parsedReq,
QuotaPlatform: quotaPlatform,
APIKey: apiKey,
User: apiKey.User,
@@ -525,7 +526,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
UserAgent: userAgent,
IPAddress: clientIP,
RequestPayloadHash: requestPayloadHash,
- ForceCacheBilling: fs.ForceCacheBilling,
+ ForceCacheBilling: forceCacheBilling,
APIKeyService: h.apiKeyService,
ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel),
}); err != nil {
@@ -918,11 +919,12 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
}
// 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。
+ // ForceCacheBilling 提前拍成标量,避免 worker 闭包保活 failover 状态里的响应体。
+ forceCacheBilling := fs.ForceCacheBilling
quotaPlatform := service.QuotaPlatform(c.Request.Context(), currentAPIKey)
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
Result: result,
- ParsedRequest: attemptParsedReq,
QuotaPlatform: quotaPlatform,
APIKey: currentAPIKey,
User: currentAPIKey.User,
@@ -933,7 +935,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
UserAgent: userAgent,
IPAddress: clientIP,
RequestPayloadHash: requestPayloadHash,
- ForceCacheBilling: fs.ForceCacheBilling,
+ ForceCacheBilling: forceCacheBilling,
APIKeyService: h.apiKeyService,
ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel),
}); err != nil {
diff --git a/backend/internal/handler/gemini_v1beta_handler.go b/backend/internal/handler/gemini_v1beta_handler.go
index 5d8e6fa8..024f8727 100644
--- a/backend/internal/handler/gemini_v1beta_handler.go
+++ b/backend/internal/handler/gemini_v1beta_handler.go
@@ -527,6 +527,8 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
requestPayloadHash := service.HashUsageRequestPayload(body)
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
+ // ForceCacheBilling 提前拍成标量,避免 worker 闭包保活 failover 状态里的响应体。
+ forceCacheBilling := fs.ForceCacheBilling
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsageWithLongContext(ctx, &service.RecordUsageLongContextInput{
@@ -543,7 +545,7 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
RequestPayloadHash: requestPayloadHash,
LongContextThreshold: 200000, // Gemini 200K 阈值
LongContextMultiplier: 2.0, // 超出部分双倍计费
- ForceCacheBilling: fs.ForceCacheBilling,
+ ForceCacheBilling: forceCacheBilling,
APIKeyService: h.apiKeyService,
ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel),
}); err != nil {
diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go
index a131cbd9..ac312dbb 100644
--- a/backend/internal/handler/openai_gateway_handler.go
+++ b/backend/internal/handler/openai_gateway_handler.go
@@ -1379,6 +1379,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
zap.Int("candidate_count", scheduleDecision.CandidateCount),
)
+ var requestPayloadHash string
hooks := &service.OpenAIWSIngressHooks{
InitialRequestModel: reqModel,
BeforeRequest: func(turn int, payload []byte, originalModel string) error {
@@ -1464,7 +1465,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
UpstreamEndpoint: upstreamEndpoint,
UserAgent: userAgent,
IPAddress: clientIP,
- RequestPayloadHash: service.HashUsageRequestPayload(firstMessage),
+ RequestPayloadHash: requestPayloadHash,
APIKeyService: h.apiKeyService,
ChannelUsageFields: channelMappingWS.ToUsageFields(reqModel, result.UpstreamModel),
}); err != nil {
@@ -1484,6 +1485,9 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
wsFirstMessage = h.gatewayService.ReplaceModelInBody(firstMessage, channelMappingWS.MappedModel)
}
+ // WebSocket 首包可能很大,hash 必须在 hooks 外算成字符串,避免 AfterTurn 闭包保活请求体。
+ requestPayloadHash = service.HashUsageRequestPayload(wsFirstMessage)
+
if err := h.gatewayService.ProxyResponsesWebSocketFromClient(ctx, c, wsConn, account, token, wsFirstMessage, hooks); err != nil {
var failoverErr *service.UpstreamFailoverError
if errors.As(err, &failoverErr) {
diff --git a/backend/internal/handler/openai_images.go b/backend/internal/handler/openai_images.go
index 36339d4b..1e3b5306 100644
--- a/backend/internal/handler/openai_images.go
+++ b/backend/internal/handler/openai_images.go
@@ -73,9 +73,10 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
+ requestModel := parsed.Model
reqLog = reqLog.With(
- zap.String("model", parsed.Model),
+ zap.String("model", requestModel),
zap.Bool("stream", parsed.Stream),
zap.Bool("multipart", parsed.Multipart),
zap.String("capability", string(parsed.RequiredCapability)),
@@ -85,7 +86,7 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
h.errorResponse(c, http.StatusForbidden, "permission_error", service.ImageGenerationPermissionMessage())
return
}
- if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIImages, parsed.Model, parsed.ModerationBody()); decision != nil && decision.Blocked {
+ if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIImages, requestModel, parsed.ModerationBody()); decision != nil && decision.Blocked {
h.errorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
return
}
@@ -98,13 +99,13 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
}
if parsed.Multipart {
- setOpsRequestContext(c, parsed.Model, parsed.Stream)
+ setOpsRequestContext(c, requestModel, parsed.Stream)
} else {
- setOpsRequestContext(c, parsed.Model, parsed.Stream)
+ setOpsRequestContext(c, requestModel, parsed.Stream)
}
setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(parsed.Stream, false)))
- channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, parsed.Model)
+ channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, requestModel)
if h.errorPassthroughService != nil {
service.BindErrorPassthroughService(c, h.errorPassthroughService)
@@ -147,7 +148,7 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
c.Request.Context(),
apiKey.GroupID,
sessionHash,
- parsed.Model,
+ requestModel,
failedAccountIDs,
parsed.RequiredCapability,
)
@@ -324,14 +325,14 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
IPAddress: clientIP,
RequestPayloadHash: requestPayloadHash,
APIKeyService: h.apiKeyService,
- ChannelUsageFields: channelMapping.ToUsageFields(parsed.Model, upstreamModel),
+ ChannelUsageFields: channelMapping.ToUsageFields(requestModel, upstreamModel),
}); err != nil {
logger.L().With(
zap.String("component", "handler.openai_gateway.images"),
zap.Int64("user_id", subject.UserID),
zap.Int64("api_key_id", apiKey.ID),
zap.Any("group_id", apiKey.GroupID),
- zap.String("model", parsed.Model),
+ zap.String("model", requestModel),
zap.Int64("account_id", account.ID),
).Error("openai.images.record_usage_failed", zap.Error(err))
}
diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go
index 01b541e2..78a1f4bd 100644
--- a/backend/internal/service/gateway_service.go
+++ b/backend/internal/service/gateway_service.go
@@ -4729,6 +4729,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
if retryErr == nil {
if retryResp.StatusCode < 400 {
// 重试请求被上游接受后同步 ParsedRequest,保证 usage/日志看到真实请求体。
+ lastWireBody = retryWireBody
if err := replaceBody(retryWireBody); err != nil {
_ = retryResp.Body.Close()
return nil, err
@@ -4769,6 +4770,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
if retryErr2 == nil {
if retryResp2.StatusCode < 400 {
// 二阶段工具块降级成功时也必须更新当前 body。
+ lastWireBody = retryWireBody2
if err := replaceBody(retryWireBody2); err != nil {
_ = retryResp2.Body.Close()
return nil, err
@@ -4847,6 +4849,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
if retryErr == nil {
if budgetRetryResp.StatusCode < 400 {
// budget 修正请求成功后,ParsedRequest 也要描述被接受的修正版。
+ lastWireBody = budgetWireBody
if err := replaceBody(budgetWireBody); err != nil {
_ = budgetRetryResp.Body.Close()
return nil, err
@@ -8228,10 +8231,10 @@ func (s *GatewayService) getUserGroupRateMultiplier(ctx context.Context, userID,
return resolver.Resolve(ctx, userID, groupID, groupDefaultMultiplier)
}
-// RecordUsageInput 记录使用量的输入参数
+// RecordUsageInput 记录使用量的输入参数。
+// 异步 worker 只接收计费所需快照,不能持有 ParsedRequest/RequestBodyRef 这类大请求体引用。
type RecordUsageInput struct {
Result *ForwardResult
- ParsedRequest *ParsedRequest
APIKey *APIKey
User *User
Account *Account
@@ -8709,15 +8712,8 @@ func writeUsageLogBestEffort(ctx context.Context, repo UsageLogRepository, usage
}
}
-// recordUsageOpts 内部选项,参数化 RecordUsage 与 RecordUsageWithLongContext 的差异点。
+// recordUsageOpts 内部选项,参数化普通计费与长上下文计费的差异点。
type recordUsageOpts struct {
- // Claude Max 策略所需的 ParsedRequest(可选,仅 Claude 路径传入)
- ParsedRequest *ParsedRequest
-
- // EnableClaudePath 启用 Claude 路径特有逻辑:
- // - Claude Max 缓存计费策略
- EnableClaudePath bool
-
// 长上下文计费(仅 Gemini 路径需要)
LongContextThreshold int
LongContextMultiplier float64
@@ -8740,9 +8736,7 @@ func (s *GatewayService) RecordUsage(ctx context.Context, input *RecordUsageInpu
APIKeyService: input.APIKeyService,
QuotaPlatform: input.QuotaPlatform,
ChannelUsageFields: input.ChannelUsageFields,
- }, &recordUsageOpts{
- EnableClaudePath: true,
- })
+ }, &recordUsageOpts{})
}
// RecordUsageLongContextInput 记录使用量的输入参数(支持长上下文双倍计费)
@@ -8808,9 +8802,7 @@ type recordUsageCoreInput struct {
}
// recordUsageCore 是 RecordUsage 和 RecordUsageWithLongContext 的统一实现。
-// opts 中的字段控制两者之间的差异行为:
-// - ParsedRequest != nil → 启用 Claude Max 缓存计费策略
-// - LongContextThreshold > 0 → Token 计费回退走 CalculateCostWithLongContext
+// LongContextThreshold > 0 时 Token 计费回退走 CalculateCostWithLongContext。
func (s *GatewayService) recordUsageCore(ctx context.Context, input *recordUsageCoreInput, opts *recordUsageOpts) error {
result := input.Result
apiKey := input.APIKey
From 5a3e193b5371b7fd89f1b23c30c69357fe9e8d8c Mon Sep 17 00:00:00 2001
From: name <136912576+is7Qin@users.noreply.github.com>
Date: Sat, 30 May 2026 16:52:15 +0800
Subject: [PATCH 52/79] refactor(gateway): shorten OpenAI request body
retention
---
backend/internal/handler/gateway_helper.go | 27 --------
.../handler/gateway_helper_hotpath_test.go | 20 +-----
.../handler/openai_gateway_handler.go | 11 +---
.../service/openai_gateway_service.go | 62 +----------------
.../openai_gateway_service_hotpath_test.go | 32 +--------
.../service/openai_tool_continuation.go | 63 +++++++++++++++++-
.../service/openai_tool_continuation_test.go | 66 +++++++++++++++++++
7 files changed, 137 insertions(+), 144 deletions(-)
diff --git a/backend/internal/handler/gateway_helper.go b/backend/internal/handler/gateway_helper.go
index 52362dd1..4b6a47eb 100644
--- a/backend/internal/handler/gateway_helper.go
+++ b/backend/internal/handler/gateway_helper.go
@@ -18,18 +18,12 @@ import (
// claudeCodeValidator is a singleton validator for Claude Code client detection
var claudeCodeValidator = service.NewClaudeCodeValidator()
-const claudeCodeParsedRequestContextKey = "claude_code_parsed_request"
-
// SetClaudeCodeClientContext 检查请求是否来自 Claude Code 客户端,并设置到 context 中
// 返回更新后的 context
func SetClaudeCodeClientContext(c *gin.Context, body []byte, parsedReq *service.ParsedRequest) {
if c == nil || c.Request == nil {
return
}
- if parsedReq != nil {
- c.Set(claudeCodeParsedRequestContextKey, parsedReq)
- }
-
ua := c.GetHeader("User-Agent")
// Fast path:非 Claude CLI UA 直接判定 false,避免热路径二次 JSON 反序列化。
if !claudeCodeValidator.ValidateUserAgent(ua) {
@@ -45,9 +39,6 @@ func SetClaudeCodeClientContext(c *gin.Context, body []byte, parsedReq *service.
} else {
// 仅在确认为 Claude CLI 且 messages 路径时再做 body 解析。
bodyMap := claudeCodeBodyMapFromParsedRequest(parsedReq)
- if bodyMap == nil {
- bodyMap = claudeCodeBodyMapFromContextCache(c)
- }
if bodyMap == nil && len(body) > 0 {
_ = json.Unmarshal(body, &bodyMap)
}
@@ -87,24 +78,6 @@ func claudeCodeBodyMapFromParsedRequest(parsedReq *service.ParsedRequest) map[st
return bodyMap
}
-func claudeCodeBodyMapFromContextCache(c *gin.Context) map[string]any {
- if c == nil {
- return nil
- }
- if bodyMap := service.CachedOpenAIParsedRequestBody(c); bodyMap != nil {
- return bodyMap
- }
- if cached, ok := c.Get(claudeCodeParsedRequestContextKey); ok {
- switch v := cached.(type) {
- case *service.ParsedRequest:
- return claudeCodeBodyMapFromParsedRequest(v)
- case service.ParsedRequest:
- return claudeCodeBodyMapFromParsedRequest(&v)
- }
- }
- return nil
-}
-
// 并发槽位等待相关常量
//
// 性能优化说明:
diff --git a/backend/internal/handler/gateway_helper_hotpath_test.go b/backend/internal/handler/gateway_helper_hotpath_test.go
index d973bedb..a6b6a429 100644
--- a/backend/internal/handler/gateway_helper_hotpath_test.go
+++ b/backend/internal/handler/gateway_helper_hotpath_test.go
@@ -177,7 +177,7 @@ func TestSetClaudeCodeClientContext_FastPathAndStrictPath(t *testing.T) {
})
}
-func TestSetClaudeCodeClientContext_ReuseParsedRequestAndContextCache(t *testing.T) {
+func TestSetClaudeCodeClientContext_ReuseParsedRequest(t *testing.T) {
t.Run("reuse parsed request without body unmarshal", func(t *testing.T) {
c, _ := newHelperTestContext(http.MethodPost, "/v1/messages")
c.Request.Header.Set("User-Agent", "claude-cli/1.0.1")
@@ -192,24 +192,6 @@ func TestSetClaudeCodeClientContext_ReuseParsedRequestAndContextCache(t *testing
SetClaudeCodeClientContext(c, []byte(`{invalid`), parsedReq)
require.True(t, service.IsClaudeCodeClient(c.Request.Context()))
})
-
- t.Run("reuse context cache without body unmarshal", func(t *testing.T) {
- c, _ := newHelperTestContext(http.MethodPost, "/v1/messages")
- c.Request.Header.Set("User-Agent", "claude-cli/1.0.1")
- c.Request.Header.Set("X-App", "claude-code")
- c.Request.Header.Set("anthropic-beta", "message-batches-2024-09-24")
- c.Request.Header.Set("anthropic-version", "2023-06-01")
- service.CacheOpenAIParsedRequestBody(c, []byte(`{invalid`), map[string]any{
- "model": "claude-3-5-sonnet-20241022",
- "system": []any{
- map[string]any{"text": "You are Claude Code, Anthropic's official CLI for Claude."},
- },
- "metadata": map[string]any{"user_id": "user_aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa_account__session_aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa"},
- })
-
- SetClaudeCodeClientContext(c, []byte(`{invalid`), nil)
- require.True(t, service.IsClaudeCodeClient(c.Request.Context()))
- })
}
func TestWaitForSlotWithPingTimeout_AccountAndUserAcquire(t *testing.T) {
diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go
index ac312dbb..0aa477b0 100644
--- a/backend/internal/handler/openai_gateway_handler.go
+++ b/backend/internal/handler/openai_gateway_handler.go
@@ -948,19 +948,12 @@ func (h *OpenAIGatewayHandler) validateFunctionCallOutputRequest(c *gin.Context,
return true
}
- var reqBody map[string]any
- if err := json.Unmarshal(body, &reqBody); err != nil {
- // 保持原有容错语义:解析失败时跳过预校验,沿用后续上游校验结果。
- return true
- }
-
- service.CacheOpenAIParsedRequestBody(c, body, reqBody)
- validation := service.ValidateFunctionCallOutputContext(reqBody)
+ validation := service.ValidateFunctionCallOutputContextBytes(body)
if !validation.HasFunctionCallOutput {
return true
}
- previousResponseID, _ := reqBody["previous_response_id"].(string)
+ previousResponseID := gjson.GetBytes(body, "previous_response_id").String()
if strings.TrimSpace(previousResponseID) != "" || validation.HasToolCallContext {
return true
}
diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go
index f8c34031..62182655 100644
--- a/backend/internal/service/openai_gateway_service.go
+++ b/backend/internal/service/openai_gateway_service.go
@@ -36,13 +36,6 @@ import (
"go.uber.org/zap"
)
-// openAIParsedRequestBodyCache 绑定 body 指纹,避免 handler 预解析的旧 body 污染后续 forwardBody。
-type openAIParsedRequestBodyCache struct {
- bodyHash uint64
- bodyLen int
- reqBody map[string]any
-}
-
const (
// ChatGPT internal API for OAuth accounts
chatgptCodexURL = "https://chatgpt.com/backend-api/codex/responses"
@@ -53,8 +46,6 @@ const (
// codex_cli_only 拒绝时单个请求头日志长度上限(字符)
codexCLIOnlyHeaderValueMaxBytes = 256
- // OpenAIParsedRequestBodyKey 缓存 handler 侧已解析的请求体,避免重复解析。
- OpenAIParsedRequestBodyKey = "openai_parsed_request_body"
// OpenAI WS Mode 失败后的重连次数上限(不含首次尝试)。
// 与 Codex 客户端保持一致:失败后最多重连 5 次。
openAIWSReconnectRetryLimit = 5
@@ -3099,8 +3090,6 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
})
s.handleFailoverSideEffects(ctx, resp, account, upstreamModel)
- // reqBody 会被本次账号尝试原地修改,failover 前必须释放,避免下一账号复用脏 map。
- releaseOpenAIParsedRequestBody(c)
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
@@ -3113,7 +3102,8 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
reasoningEffort := extractOpenAIReasoningEffort(reqBody, originalModel)
serviceTier := extractOpenAIServiceTier(reqBody)
- releaseOpenAIParsedRequestBody(c)
+ // 上游接受后只保留计费需要的标量,避免响应处理期间继续保活完整 input/tools map。
+ reqBody = nil
// Handle normal response
var usage *OpenAIUsage
@@ -6843,60 +6833,14 @@ func isEmptyBase64DataURI(raw string) bool {
return strings.TrimSpace(strings.TrimPrefix(rest, "base64,")) == ""
}
-func getOpenAIRequestBodyMap(c *gin.Context, body []byte) (map[string]any, error) {
- // 同一个 gin.Context 内 failover/渠道映射可能传入新 body,缓存必须先校验 body 指纹。
- bodyHash := xxhash.Sum64(body)
- bodyLen := len(body)
- if c != nil {
- if cached, ok := c.Get(OpenAIParsedRequestBodyKey); ok {
- if cache, ok := cached.(openAIParsedRequestBodyCache); ok && cache.reqBody != nil && cache.bodyLen == bodyLen && cache.bodyHash == bodyHash {
- return cache.reqBody, nil
- }
- }
- }
-
+func getOpenAIRequestBodyMap(_ *gin.Context, body []byte) (map[string]any, error) {
var reqBody map[string]any
if err := json.Unmarshal(body, &reqBody); err != nil {
return nil, fmt.Errorf("parse request: %w", err)
}
- if c != nil {
- c.Set(OpenAIParsedRequestBodyKey, openAIParsedRequestBodyCache{bodyHash: bodyHash, bodyLen: bodyLen, reqBody: reqBody})
- }
return reqBody, nil
}
-// CacheOpenAIParsedRequestBody 仅缓存与当前 body 绑定的解析结果。
-func CacheOpenAIParsedRequestBody(c *gin.Context, body []byte, reqBody map[string]any) {
- if c == nil || reqBody == nil {
- return
- }
- c.Set(OpenAIParsedRequestBodyKey, openAIParsedRequestBodyCache{
- bodyHash: xxhash.Sum64(body),
- bodyLen: len(body),
- reqBody: reqBody,
- })
-}
-
-// CachedOpenAIParsedRequestBody 只给同请求内不关心 body 参数的轻量识别逻辑使用。
-func CachedOpenAIParsedRequestBody(c *gin.Context) map[string]any {
- if c == nil {
- return nil
- }
- if cached, ok := c.Get(OpenAIParsedRequestBodyKey); ok {
- if cache, ok := cached.(openAIParsedRequestBodyCache); ok {
- return cache.reqBody
- }
- }
- return nil
-}
-
-func releaseOpenAIParsedRequestBody(c *gin.Context) {
- if c == nil {
- return
- }
- delete(c.Keys, OpenAIParsedRequestBodyKey)
-}
-
func extractOpenAIReasoningEffort(reqBody map[string]any, requestedModel string) *string {
if value, present := getOpenAIReasoningEffortFromReqBody(reqBody); present {
if value == "" {
diff --git a/backend/internal/service/openai_gateway_service_hotpath_test.go b/backend/internal/service/openai_gateway_service_hotpath_test.go
index 2ff72e2f..af0d21c4 100644
--- a/backend/internal/service/openai_gateway_service_hotpath_test.go
+++ b/backend/internal/service/openai_gateway_service_hotpath_test.go
@@ -106,26 +106,13 @@ func TestExtractOpenAIReasoningEffortFromBody(t *testing.T) {
}
}
-func TestGetOpenAIRequestBodyMap_UsesContextCache(t *testing.T) {
- gin.SetMode(gin.TestMode)
- rec := httptest.NewRecorder()
- c, _ := gin.CreateTestContext(rec)
-
- cached := map[string]any{"model": "cached-model", "stream": true}
- CacheOpenAIParsedRequestBody(c, []byte(`{invalid-json`), cached)
-
- got, err := getOpenAIRequestBodyMap(c, []byte(`{invalid-json`))
- require.NoError(t, err)
- require.Equal(t, cached, got)
-}
-
-func TestGetOpenAIRequestBodyMap_ParseErrorWithoutCache(t *testing.T) {
+func TestGetOpenAIRequestBodyMap_ParseError(t *testing.T) {
_, err := getOpenAIRequestBodyMap(nil, []byte(`{invalid-json`))
require.Error(t, err)
require.Contains(t, err.Error(), "parse request")
}
-func TestGetOpenAIRequestBodyMap_WriteBackContextCache(t *testing.T) {
+func TestGetOpenAIRequestBodyMap_DoesNotWriteContextCache(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
@@ -133,20 +120,7 @@ func TestGetOpenAIRequestBodyMap_WriteBackContextCache(t *testing.T) {
got, err := getOpenAIRequestBodyMap(c, []byte(`{"model":"gpt-5","stream":true}`))
require.NoError(t, err)
require.Equal(t, "gpt-5", got["model"])
-
- require.Equal(t, got, CachedOpenAIParsedRequestBody(c))
-}
-
-func TestGetOpenAIRequestBodyMap_IgnoresCacheForDifferentBody(t *testing.T) {
- gin.SetMode(gin.TestMode)
- rec := httptest.NewRecorder()
- c, _ := gin.CreateTestContext(rec)
-
- CacheOpenAIParsedRequestBody(c, []byte(`{"model":"cached-model"}`), map[string]any{"model": "cached-model"})
-
- got, err := getOpenAIRequestBodyMap(c, []byte(`{"model":"forward-model"}`))
- require.NoError(t, err)
- require.Equal(t, "forward-model", got["model"])
+ require.Empty(t, c.Keys)
}
func TestSanitizeEmptyBase64InputImagesInOpenAIRequestBodyMap(t *testing.T) {
diff --git a/backend/internal/service/openai_tool_continuation.go b/backend/internal/service/openai_tool_continuation.go
index 7d503f5a..36064549 100644
--- a/backend/internal/service/openai_tool_continuation.go
+++ b/backend/internal/service/openai_tool_continuation.go
@@ -1,6 +1,10 @@
package service
-import "strings"
+import (
+ "strings"
+
+ "github.com/tidwall/gjson"
+)
// ToolContinuationSignals 聚合工具续链相关信号,避免重复遍历 input。
type ToolContinuationSignals struct {
@@ -150,6 +154,63 @@ func AnalyzeToolContinuationSignals(reqBody map[string]any) ToolContinuationSign
return signals
}
+// ValidateFunctionCallOutputContextBytes 基于 raw JSON 校验工具输出续链,避免 handler 预校验阶段全量解码大 input。
+func ValidateFunctionCallOutputContextBytes(body []byte) FunctionCallOutputValidation {
+ result := FunctionCallOutputValidation{}
+ input := gjson.GetBytes(body, "input")
+ if !input.IsArray() {
+ return result
+ }
+
+ var callIDs map[string]struct{}
+ var referenceIDs map[string]struct{}
+ input.ForEach(func(_, item gjson.Result) bool {
+ if !item.IsObject() {
+ return true
+ }
+ itemType := item.Get("type").String()
+ switch {
+ case isCodexToolCallOutputItemType(itemType):
+ result.HasFunctionCallOutput = true
+ callID := strings.TrimSpace(item.Get("call_id").String())
+ if callID == "" {
+ result.HasFunctionCallOutputMissingCallID = true
+ return true
+ }
+ if callIDs == nil {
+ callIDs = make(map[string]struct{})
+ }
+ callIDs[callID] = struct{}{}
+ case isCodexToolCallContextItemType(itemType):
+ if strings.TrimSpace(item.Get("call_id").String()) != "" {
+ result.HasToolCallContext = true
+ }
+ case itemType == "item_reference":
+ idValue := strings.TrimSpace(item.Get("id").String())
+ if idValue == "" {
+ return true
+ }
+ if referenceIDs == nil {
+ referenceIDs = make(map[string]struct{})
+ }
+ referenceIDs[idValue] = struct{}{}
+ }
+ return !(result.HasFunctionCallOutput && result.HasToolCallContext)
+ })
+ if !result.HasFunctionCallOutput || result.HasToolCallContext || len(callIDs) == 0 || len(referenceIDs) == 0 {
+ return result
+ }
+ allReferenced := true
+ for callID := range callIDs {
+ if _, ok := referenceIDs[callID]; !ok {
+ allReferenced = false
+ break
+ }
+ }
+ result.HasItemReferenceForAllCallIDs = allReferenced
+ return result
+}
+
// ValidateFunctionCallOutputContext 为 handler 提供低开销校验结果:
// 1) 无工具输出直接返回
// 2) 若已存在工具调用上下文则提前返回
diff --git a/backend/internal/service/openai_tool_continuation_test.go b/backend/internal/service/openai_tool_continuation_test.go
index 0e0552f6..4610652b 100644
--- a/backend/internal/service/openai_tool_continuation_test.go
+++ b/backend/internal/service/openai_tool_continuation_test.go
@@ -1,6 +1,7 @@
package service
import (
+ "encoding/json"
"testing"
"github.com/stretchr/testify/require"
@@ -118,3 +119,68 @@ func TestHasItemReferenceForCallIDs(t *testing.T) {
require.True(t, HasItemReferenceForCallIDs(req, []string{"call_1", "call_2"}))
require.False(t, HasItemReferenceForCallIDs(req, []string{"call_1", "call_3"}))
}
+
+func TestValidateFunctionCallOutputContextBytesMatchesMapValidation(t *testing.T) {
+ // handler 预校验走 raw JSON 扫描,语义必须与 service 内部 map 校验保持一致。
+ cases := []struct {
+ name string
+ body map[string]any
+ }{
+ {
+ name: "no_input",
+ body: map[string]any{"model": "gpt-5.4"},
+ },
+ {
+ name: "missing_call_id",
+ body: map[string]any{"input": []any{map[string]any{"type": "function_call_output"}}},
+ },
+ {
+ name: "call_id_without_reference",
+ body: map[string]any{"input": []any{map[string]any{"type": "function_call_output", "call_id": "call_1"}}},
+ },
+ {
+ name: "matching_reference",
+ body: map[string]any{"input": []any{
+ map[string]any{"type": "function_call_output", "call_id": "call_1"},
+ map[string]any{"type": "item_reference", "id": "call_1"},
+ }},
+ },
+ {
+ name: "partial_reference",
+ body: map[string]any{"input": []any{
+ map[string]any{"type": "function_call_output", "call_id": "call_1"},
+ map[string]any{"type": "tool_search_output", "call_id": "call_2"},
+ map[string]any{"type": "item_reference", "id": "call_1"},
+ }},
+ },
+ {
+ name: "tool_context",
+ body: map[string]any{"input": []any{
+ map[string]any{"type": "function_call_output", "call_id": "call_1"},
+ map[string]any{"type": "function_call", "call_id": "call_1"},
+ }},
+ },
+ {
+ name: "all_codex_tool_outputs",
+ body: map[string]any{"input": []any{
+ map[string]any{"type": "function_call_output", "call_id": "call_function"},
+ map[string]any{"type": "tool_search_output", "call_id": "call_search"},
+ map[string]any{"type": "custom_tool_call_output", "call_id": "call_custom"},
+ map[string]any{"type": "mcp_tool_call_output", "call_id": "call_mcp"},
+ map[string]any{"type": "item_reference", "id": "call_function"},
+ map[string]any{"type": "item_reference", "id": "call_search"},
+ map[string]any{"type": "item_reference", "id": "call_custom"},
+ map[string]any{"type": "item_reference", "id": "call_mcp"},
+ }},
+ },
+ }
+
+ for _, tt := range cases {
+ t.Run(tt.name, func(t *testing.T) {
+ bodyBytes, err := json.Marshal(tt.body)
+ require.NoError(t, err)
+
+ require.Equal(t, ValidateFunctionCallOutputContext(tt.body), ValidateFunctionCallOutputContextBytes(bodyBytes))
+ })
+ }
+}
From 09af6ebd405d42f39f0b9fcf42924e59e903ba40 Mon Sep 17 00:00:00 2001
From: name <136912576+is7Qin@users.noreply.github.com>
Date: Sat, 30 May 2026 16:58:58 +0800
Subject: [PATCH 53/79] refactor(gateway): avoid extra OpenAI WS body map copy
---
backend/internal/service/openai_gateway_service.go | 7 +------
1 file changed, 1 insertion(+), 6 deletions(-)
diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go
index 62182655..2df2206d 100644
--- a/backend/internal/service/openai_gateway_service.go
+++ b/backend/internal/service/openai_gateway_service.go
@@ -2793,13 +2793,8 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
// 命中 WS 时仅走 WebSocket Mode;不再自动回退 HTTP。
if wsDecision.Transport == OpenAIUpstreamTransportResponsesWebsocketV2 {
+ // WS 分支不会再回落 HTTP;重连恢复可直接更新 reqBody,避免额外保留一份完整顶层 map。
wsReqBody := reqBody
- if len(reqBody) > 0 {
- wsReqBody = make(map[string]any, len(reqBody))
- for k, v := range reqBody {
- wsReqBody[k] = v
- }
- }
_, hasPreviousResponseID := wsReqBody["previous_response_id"]
logOpenAIWSModeDebug(
"forward_start account_id=%d account_type=%s model=%s stream=%v has_previous_response_id=%v",
From 34de99ee0e3ae69c03774ae5e354c2e6494ca67f Mon Sep 17 00:00:00 2001
From: name <136912576+is7Qin@users.noreply.github.com>
Date: Sat, 30 May 2026 19:51:03 +0800
Subject: [PATCH 54/79] refactor(gateway): cap upstream error body reads
---
.../service/antigravity_gateway_service.go | 30 ++-
.../antigravity_gateway_service_test.go | 9 +
.../gateway_forward_as_chat_completions.go | 2 +-
.../service/gateway_forward_as_responses.go | 2 +-
backend/internal/service/gateway_service.go | 62 +++--
.../service/gateway_service_benchmark_test.go | 214 +++++++++++++++++-
.../gemini_chat_completions_compat_service.go | 4 +-
.../service/gemini_messages_compat_service.go | 24 +-
backend/internal/service/openai_embeddings.go | 2 +-
.../openai_gateway_chat_completions.go | 2 +-
.../openai_gateway_chat_completions_raw.go | 2 +-
.../service/openai_gateway_messages.go | 2 +-
.../openai_gateway_responses_chat_fallback.go | 2 +-
.../service/openai_gateway_service.go | 34 ++-
backend/internal/service/openai_images.go | 23 +-
.../service/openai_images_responses.go | 2 +-
.../internal/service/openai_images_test.go | 31 ++-
.../service/openai_tool_continuation.go | 6 +-
18 files changed, 392 insertions(+), 61 deletions(-)
diff --git a/backend/internal/service/antigravity_gateway_service.go b/backend/internal/service/antigravity_gateway_service.go
index 2b849bdd..de62cfab 100644
--- a/backend/internal/service/antigravity_gateway_service.go
+++ b/backend/internal/service/antigravity_gateway_service.go
@@ -662,7 +662,7 @@ urlFallbackLoop:
// 统一处理错误响应
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
if overagesInjected && shouldMarkCreditsExhausted(resp, respBody, nil) {
@@ -875,6 +875,22 @@ type AntigravityGatewayService struct {
internal500Cache Internal500CounterCache // INTERNAL 500 渐进惩罚计数器
}
+func (s *AntigravityGatewayService) upstreamErrorBodyReadLimit() int64 {
+ limit := gatewayUpstreamErrorBodyReadLimit
+ if s != nil && s.settingService != nil && s.settingService.cfg != nil && s.settingService.cfg.Gateway.LogUpstreamErrorBody && s.settingService.cfg.Gateway.LogUpstreamErrorBodyMaxBytes > int(limit) {
+ limit = int64(s.settingService.cfg.Gateway.LogUpstreamErrorBodyMaxBytes)
+ }
+ return limit
+}
+
+func (s *AntigravityGatewayService) readUpstreamErrorBody(resp *http.Response) []byte {
+ if resp == nil || resp.Body == nil {
+ return nil
+ }
+ body, _ := io.ReadAll(io.LimitReader(resp.Body, s.upstreamErrorBodyReadLimit()))
+ return body
+}
+
func NewAntigravityGatewayService(
accountRepo AccountRepository,
cache GatewayCache,
@@ -1090,7 +1106,7 @@ func (s *AntigravityGatewayService) TestConnection(ctx context.Context, account
}
defer func() { _ = result.resp.Body.Close() }()
- respBody, err := io.ReadAll(io.LimitReader(result.resp.Body, 2<<20))
+ respBody, err := io.ReadAll(io.LimitReader(result.resp.Body, s.upstreamErrorBodyReadLimit()))
if err != nil {
return nil, fmt.Errorf("读取响应失败: %w", err)
}
@@ -1427,7 +1443,7 @@ func (s *AntigravityGatewayService) Forward(ctx context.Context, c *gin.Context,
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
// 优先检测 thinking block 的 signature 相关错误(400)并重试一次:
// Antigravity /v1internal 链路在部分场景会对 thought/thinking signature 做严格校验,
@@ -1622,7 +1638,7 @@ func (s *AntigravityGatewayService) Forward(ctx context.Context, c *gin.Context,
resp = retryResp
respBody = nil
} else {
- retryBody, _ := io.ReadAll(io.LimitReader(retryResp.Body, 2<<20))
+ retryBody := s.readUpstreamErrorBody(retryResp)
_ = retryResp.Body.Close()
respBody = retryBody
resp = &http.Response{
@@ -2189,7 +2205,7 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co
// 处理错误响应
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
contentType := resp.Header.Get("Content-Type")
// 尽早关闭原始响应体,释放连接;后续逻辑仍可能需要读取 body,因此用内存副本重新包装。
_ = resp.Body.Close()
@@ -2270,7 +2286,7 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co
if retryResp.StatusCode < 400 {
resp = retryResp
} else {
- retryRespBody, _ := io.ReadAll(io.LimitReader(retryResp.Body, 2<<20))
+ retryRespBody := s.readUpstreamErrorBody(retryResp)
_ = retryResp.Body.Close()
retryOpsBody := retryRespBody
if retryUnwrapped, unwrapErr := s.unwrapV1InternalResponse(retryRespBody); unwrapErr == nil && len(retryUnwrapped) > 0 {
@@ -4252,7 +4268,7 @@ func (s *AntigravityGatewayService) ForwardUpstream(ctx context.Context, c *gin.
// 处理错误响应
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
// 429 错误时标记账号限流
if resp.StatusCode == http.StatusTooManyRequests {
diff --git a/backend/internal/service/antigravity_gateway_service_test.go b/backend/internal/service/antigravity_gateway_service_test.go
index 22124374..903c8a1a 100644
--- a/backend/internal/service/antigravity_gateway_service_test.go
+++ b/backend/internal/service/antigravity_gateway_service_test.go
@@ -42,6 +42,15 @@ func newAntigravityTestService(cfg *config.Config) *AntigravityGatewayService {
}
}
+func TestAntigravityUpstreamErrorBodyReadLimit_RespectsDiagnosticLimit(t *testing.T) {
+ svc := newAntigravityTestService(&config.Config{Gateway: config.GatewayConfig{
+ LogUpstreamErrorBody: true,
+ LogUpstreamErrorBodyMaxBytes: int(gatewayUpstreamErrorBodyReadLimit) + 1024,
+ }})
+
+ require.Equal(t, int64(svc.settingService.cfg.Gateway.LogUpstreamErrorBodyMaxBytes), svc.upstreamErrorBodyReadLimit())
+}
+
func TestStripSignatureSensitiveBlocksFromClaudeRequest(t *testing.T) {
req := &antigravity.ClaudeRequest{
Model: "claude-sonnet-4-5",
diff --git a/backend/internal/service/gateway_forward_as_chat_completions.go b/backend/internal/service/gateway_forward_as_chat_completions.go
index 729483e3..1df450d6 100644
--- a/backend/internal/service/gateway_forward_as_chat_completions.go
+++ b/backend/internal/service/gateway_forward_as_chat_completions.go
@@ -148,7 +148,7 @@ func (s *GatewayService) ForwardAsChatCompletions(
// 12. Handle error response with failover
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
diff --git a/backend/internal/service/gateway_forward_as_responses.go b/backend/internal/service/gateway_forward_as_responses.go
index 3baea018..22951b88 100644
--- a/backend/internal/service/gateway_forward_as_responses.go
+++ b/backend/internal/service/gateway_forward_as_responses.go
@@ -145,7 +145,7 @@ func (s *GatewayService) ForwardAsResponses(
// 12. Handle error response with failover
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go
index 78a1f4bd..812780dc 100644
--- a/backend/internal/service/gateway_service.go
+++ b/backend/internal/service/gateway_service.go
@@ -23,6 +23,7 @@ import (
"sync/atomic"
"syscall"
"time"
+ "unsafe"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/pkg/claude"
@@ -56,6 +57,8 @@ const (
defaultModelsListCacheTTL = 15 * time.Second
postUsageBillingTimeout = 15 * time.Second
debugGatewayBodyEnv = "SUB2API_DEBUG_GATEWAY_BODY"
+ // 上游错误体只需要提取错误 JSON/日志摘要,默认 512KiB 避免错误风暴叠加大请求体。
+ gatewayUpstreamErrorBodyReadLimit int64 = 512 << 10
)
const (
@@ -833,8 +836,16 @@ func (s *GatewayService) extractCacheableContent(parsed *ParsedRequest) string {
return systemText
}
+func parseRawJSONView(raw []byte) gjson.Result {
+ if len(raw) == 0 {
+ return gjson.Result{}
+ }
+ // 这里只做同步只读解析,避免 gjson.ParseBytes 为大 messages/contents 复制整段 raw。
+ return gjson.Parse(*(*string)(unsafe.Pointer(&raw)))
+}
+
func extractTextFromSystemRaw(raw []byte) string {
- system := gjson.ParseBytes(raw)
+ system := parseRawJSONView(raw)
switch system.Type {
case gjson.String:
return system.String()
@@ -880,7 +891,7 @@ func appendMessageTextsFromRaw(builder *strings.Builder, raw []byte) {
if builder == nil || len(raw) == 0 {
return
}
- messages := gjson.ParseBytes(raw)
+ messages := parseRawJSONView(raw)
if !messages.IsArray() {
return
}
@@ -903,7 +914,7 @@ func appendMessageTextsFromRaw(builder *strings.Builder, raw []byte) {
}
func extractCacheableTextFromSystemRaw(raw []byte) string {
- system := gjson.ParseBytes(raw)
+ system := parseRawJSONView(raw)
if !system.IsArray() {
return ""
}
@@ -920,7 +931,7 @@ func extractCacheableTextFromSystemRaw(raw []byte) string {
}
func extractCacheableTextFromMessagesRaw(raw []byte) string {
- messages := gjson.ParseBytes(raw)
+ messages := parseRawJSONView(raw)
if !messages.IsArray() {
return ""
}
@@ -4676,7 +4687,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
// 优先检测thinking block签名错误(400)并重试一次
if resp.StatusCode == 400 {
- respBody, readErr := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, readErr := s.readUpstreamErrorBody(resp)
if readErr == nil {
_ = resp.Body.Close()
@@ -4739,7 +4750,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
break
}
- retryRespBody, retryReadErr := io.ReadAll(io.LimitReader(retryResp.Body, 2<<20))
+ retryRespBody, retryReadErr := s.readUpstreamErrorBody(retryResp)
_ = retryResp.Body.Close()
if retryReadErr == nil && retryResp.StatusCode == 400 && s.isSignatureErrorPattern(ctx, account, retryRespBody) {
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
@@ -4889,7 +4900,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
break
}
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
@@ -4936,7 +4947,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
// 处理重试耗尽的情况
if resp.StatusCode >= 400 && s.shouldRetryUpstreamError(account, resp.StatusCode) {
if s.shouldFailoverUpstreamError(resp.StatusCode) {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -4971,7 +4982,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
// 处理可切换账号的错误
if resp.StatusCode >= 400 && s.shouldFailoverUpstreamError(resp.StatusCode) {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -5003,7 +5014,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
if resp.StatusCode >= 400 {
// 可选:对部分 400 触发 failover(默认关闭以保持语义)
if resp.StatusCode == 400 && s.cfg != nil && s.cfg.Gateway.FailoverOn400 {
- respBody, readErr := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, readErr := s.readUpstreamErrorBody(resp)
if readErr != nil {
// ReadAll failed, fall back to normal error handling without consuming the stream
return s.handleErrorResponse(ctx, resp, c, account, reqModel)
@@ -5221,7 +5232,7 @@ func (s *GatewayService) forwardAnthropicAPIKeyPassthroughWithInput(
break
}
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
@@ -5259,7 +5270,7 @@ func (s *GatewayService) forwardAnthropicAPIKeyPassthroughWithInput(
if resp.StatusCode >= 400 && s.shouldRetryUpstreamError(account, resp.StatusCode) {
if s.shouldFailoverUpstreamError(resp.StatusCode) {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -5293,7 +5304,7 @@ func (s *GatewayService) forwardAnthropicAPIKeyPassthroughWithInput(
}
if resp.StatusCode >= 400 && s.shouldFailoverUpstreamError(resp.StatusCode) {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -6011,7 +6022,7 @@ func (s *GatewayService) executeBedrockUpstream(
break
}
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
@@ -6056,7 +6067,7 @@ func (s *GatewayService) handleBedrockUpstreamErrors(
// retry exhausted + failover
if s.shouldRetryUpstreamError(account, resp.StatusCode) {
if s.shouldFailoverUpstreamError(resp.StatusCode) {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -6083,7 +6094,7 @@ func (s *GatewayService) handleBedrockUpstreamErrors(
// non-retryable failover
if s.shouldFailoverUpstreamError(resp.StatusCode) {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -7226,8 +7237,19 @@ func isCountTokensUnsupported404(statusCode int, body []byte) bool {
return strings.Contains(msg, "count_tokens") && strings.Contains(msg, "not found")
}
+func (s *GatewayService) readUpstreamErrorBody(resp *http.Response) ([]byte, error) {
+ if resp == nil || resp.Body == nil {
+ return nil, nil
+ }
+ limit := gatewayUpstreamErrorBodyReadLimit
+ if s != nil && s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody && s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes > int(limit) {
+ limit = int64(s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes)
+ }
+ return io.ReadAll(io.LimitReader(resp.Body, limit))
+}
+
func (s *GatewayService) handleErrorResponse(ctx context.Context, resp *http.Response, c *gin.Context, account *Account, requestedModel ...string) (*ForwardResult, error) {
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ body, _ := s.readUpstreamErrorBody(resp)
// 调试日志:打印上游错误响应
logger.LegacyPrintf("service.gateway", "[Forward] Upstream error (non-retryable): Account=%d(%s) Status=%d RequestID=%s Body=%s",
@@ -7380,7 +7402,7 @@ func (s *GatewayService) handleErrorResponse(ctx context.Context, resp *http.Res
}
func (s *GatewayService) handleRetryExhaustedSideEffects(ctx context.Context, resp *http.Response, account *Account) {
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ body, _ := s.readUpstreamErrorBody(resp)
statusCode := resp.StatusCode
// OAuth/Setup Token 账号的 403:标记账号异常
@@ -7394,7 +7416,7 @@ func (s *GatewayService) handleRetryExhaustedSideEffects(ctx context.Context, re
}
func (s *GatewayService) handleFailoverSideEffects(ctx context.Context, resp *http.Response, account *Account, requestedModel ...string) {
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ body, _ := s.readUpstreamErrorBody(resp)
if len(requestedModel) > 0 {
s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, body, requestedModel[0])
return
@@ -7407,7 +7429,7 @@ func (s *GatewayService) handleFailoverSideEffects(ctx context.Context, resp *ht
// API Key 未配置错误码:仅返回错误,不标记账号
func (s *GatewayService) handleRetryExhaustedError(ctx context.Context, resp *http.Response, c *gin.Context, account *Account) (*ForwardResult, error) {
// Capture upstream error body before side-effects consume the stream.
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
diff --git a/backend/internal/service/gateway_service_benchmark_test.go b/backend/internal/service/gateway_service_benchmark_test.go
index 8b30cb24..b2df1435 100644
--- a/backend/internal/service/gateway_service_benchmark_test.go
+++ b/backend/internal/service/gateway_service_benchmark_test.go
@@ -4,9 +4,14 @@ import (
"strconv"
"strings"
"testing"
+
+ "github.com/Wei-Shaw/sub2api/internal/domain"
)
-var benchmarkStringSink string
+var (
+ benchmarkStringSink string
+ benchmarkIntSink int
+)
// BenchmarkGenerateSessionHash_Metadata 关注 JSON 解析与正则匹配开销。
func BenchmarkGenerateSessionHash_Metadata(b *testing.B) {
@@ -23,6 +28,121 @@ func BenchmarkGenerateSessionHash_Metadata(b *testing.B) {
}
}
+func BenchmarkParseGatewayRequest_LargeAnthropicMessages(b *testing.B) {
+ for _, size := range benchmarkBodySizes() {
+ b.Run(size.name, func(b *testing.B) {
+ body := buildLargeAnthropicMessagesBody(size.bytes, false)
+
+ b.SetBytes(int64(len(body)))
+ b.ReportAllocs()
+ b.ResetTimer()
+ for i := 0; i < b.N; i++ {
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformAnthropic)
+ if err != nil {
+ b.Fatalf("解析 Anthropic 请求失败: %v", err)
+ }
+ benchmarkIntSink = len(parsed.MessagesRaw())
+ }
+ })
+ }
+}
+
+func BenchmarkParseGatewayRequest_LargeGeminiContents(b *testing.B) {
+ for _, size := range benchmarkBodySizes() {
+ b.Run(size.name, func(b *testing.B) {
+ body := buildLargeGeminiContentsBody(size.bytes)
+
+ b.SetBytes(int64(len(body)))
+ b.ReportAllocs()
+ b.ResetTimer()
+ for i := 0; i < b.N; i++ {
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
+ if err != nil {
+ b.Fatalf("解析 Gemini 请求失败: %v", err)
+ }
+ benchmarkIntSink = len(parsed.MessagesRaw())
+ }
+ })
+ }
+}
+
+func BenchmarkGenerateSessionHash_LargeAnthropicMessages(b *testing.B) {
+ svc := &GatewayService{}
+ for _, size := range benchmarkBodySizes() {
+ b.Run(size.name, func(b *testing.B) {
+ body := buildLargeAnthropicMessagesBody(size.bytes, true)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformAnthropic)
+ if err != nil {
+ b.Fatalf("解析请求失败: %v", err)
+ }
+
+ b.SetBytes(int64(len(body)))
+ b.ReportAllocs()
+ b.ResetTimer()
+ for i := 0; i < b.N; i++ {
+ benchmarkStringSink = svc.GenerateSessionHash(parsed)
+ }
+ })
+ }
+}
+
+func BenchmarkOpenAIResponses_LargeInputMeta(b *testing.B) {
+ for _, size := range benchmarkBodySizes() {
+ b.Run(size.name, func(b *testing.B) {
+ body := buildLargeOpenAIResponsesBody(size.bytes)
+
+ b.SetBytes(int64(len(body)))
+ b.ReportAllocs()
+ b.ResetTimer()
+ for i := 0; i < b.N; i++ {
+ model, stream, promptCacheKey := extractOpenAIRequestMetaFromBody(body)
+ benchmarkStringSink = model + promptCacheKey
+ if stream {
+ benchmarkIntSink++
+ }
+ }
+ })
+ }
+}
+
+func BenchmarkOpenAIResponses_LargeInputDecodeMap(b *testing.B) {
+ for _, size := range benchmarkBodySizes() {
+ b.Run(size.name, func(b *testing.B) {
+ body := buildLargeOpenAIResponsesBody(size.bytes)
+
+ b.SetBytes(int64(len(body)))
+ b.ReportAllocs()
+ b.ResetTimer()
+ for i := 0; i < b.N; i++ {
+ reqBody, err := getOpenAIRequestBodyMap(nil, body)
+ if err != nil {
+ b.Fatalf("解析 OpenAI 请求失败: %v", err)
+ }
+ benchmarkIntSink = len(reqBody)
+ }
+ })
+ }
+}
+
+func BenchmarkOpenAIResponses_LargeInputFunctionCallValidation(b *testing.B) {
+ for _, size := range benchmarkBodySizes() {
+ b.Run(size.name, func(b *testing.B) {
+ body := buildLargeOpenAIResponsesToolContinuationBody(size.bytes)
+
+ b.SetBytes(int64(len(body)))
+ b.ReportAllocs()
+ b.ResetTimer()
+ for i := 0; i < b.N; i++ {
+ validation := ValidateFunctionCallOutputContextBytes(body)
+ if !validation.HasFunctionCallOutput || !validation.HasItemReferenceForAllCallIDs {
+ b.Fatalf("工具续链校验结果异常: %+v", validation)
+ }
+ benchmarkIntSink++
+ }
+ })
+ }
+}
+
// BenchmarkExtractCacheableContent_System 关注字符串拼接路径的性能。
func BenchmarkExtractCacheableContent_System(b *testing.B) {
svc := &GatewayService{}
@@ -34,6 +154,21 @@ func BenchmarkExtractCacheableContent_System(b *testing.B) {
}
}
+func benchmarkBodySizes() []struct {
+ name string
+ bytes int
+} {
+ return []struct {
+ name string
+ bytes int
+ }{
+ {name: "4MB", bytes: 4 << 20},
+ {name: "8MB", bytes: 8 << 20},
+ {name: "16MB", bytes: 16 << 20},
+ {name: "32MB", bytes: 32 << 20},
+ }
+}
+
func buildSystemCacheableRequest(parts int) *ParsedRequest {
var builder strings.Builder
builder.WriteString(`{"system":[`)
@@ -52,3 +187,80 @@ func buildSystemCacheableRequest(parts int) *ParsedRequest {
}
return parsed
}
+
+func buildLargeAnthropicMessagesBody(targetBytes int, includeCacheControl bool) []byte {
+ var builder strings.Builder
+ builder.Grow(targetBytes + 1024)
+ builder.WriteString(`{"model":"claude-sonnet-4-5","stream":true,"system":[{"type":"text","text":"system seed"}],"messages":[`)
+ for i := 0; builder.Len() < targetBytes; i++ {
+ if i > 0 {
+ builder.WriteByte(',')
+ }
+ builder.WriteString(`{"role":"user","content":[{"type":"text","text":"`)
+ builder.WriteString(strings.Repeat("anthropic payload ", 64))
+ builder.WriteString(strconv.Itoa(i))
+ builder.WriteByte('"')
+ if includeCacheControl && i%32 == 0 {
+ builder.WriteString(`,"cache_control":{"type":"ephemeral"}`)
+ }
+ builder.WriteString(`}]}`)
+ }
+ builder.WriteString(`]}`)
+ return []byte(builder.String())
+}
+
+func buildLargeGeminiContentsBody(targetBytes int) []byte {
+ var builder strings.Builder
+ builder.Grow(targetBytes + 1024)
+ builder.WriteString(`{"model":"gemini-2.5-pro","systemInstruction":{"parts":[{"text":"system seed"}]},"contents":[`)
+ for i := 0; builder.Len() < targetBytes; i++ {
+ if i > 0 {
+ builder.WriteByte(',')
+ }
+ builder.WriteString(`{"role":"user","parts":[{"text":"`)
+ builder.WriteString(strings.Repeat("gemini payload ", 64))
+ builder.WriteString(strconv.Itoa(i))
+ builder.WriteString(`"}]}`)
+ }
+ builder.WriteString(`]}`)
+ return []byte(builder.String())
+}
+
+func buildLargeOpenAIResponsesBody(targetBytes int) []byte {
+ var builder strings.Builder
+ builder.Grow(targetBytes + 1024)
+ builder.WriteString(`{"model":"gpt-5.4","stream":true,"prompt_cache_key":"session-benchmark","input":[`)
+ for i := 0; builder.Len() < targetBytes; i++ {
+ if i > 0 {
+ builder.WriteByte(',')
+ }
+ builder.WriteString(`{"type":"message","role":"user","content":[{"type":"input_text","text":"`)
+ builder.WriteString(strings.Repeat("openai responses payload ", 48))
+ builder.WriteString(strconv.Itoa(i))
+ builder.WriteString(`"}]}`)
+ }
+ builder.WriteString(`],"tools":[{"type":"function","name":"lookup","parameters":{"type":"object","properties":{"query":{"type":"string"}}}}]}`)
+ return []byte(builder.String())
+}
+
+func buildLargeOpenAIResponsesToolContinuationBody(targetBytes int) []byte {
+ var builder strings.Builder
+ builder.Grow(targetBytes + 1024)
+ builder.WriteString(`{"model":"gpt-5.4","stream":true,"previous_response_id":"resp_benchmark","input":[`)
+ for i := 0; builder.Len() < targetBytes; i++ {
+ if i > 0 {
+ builder.WriteByte(',')
+ }
+ callID := "call_" + strconv.Itoa(i)
+ builder.WriteString(`{"type":"item_reference","id":"`)
+ builder.WriteString(callID)
+ builder.WriteString(`"},{"type":"function_call_output","call_id":"`)
+ builder.WriteString(callID)
+ builder.WriteString(`","output":"`)
+ builder.WriteString(strings.Repeat("tool output payload ", 48))
+ builder.WriteString(strconv.Itoa(i))
+ builder.WriteString(`"}`)
+ }
+ builder.WriteString(`]}`)
+ return []byte(builder.String())
+}
diff --git a/backend/internal/service/gemini_chat_completions_compat_service.go b/backend/internal/service/gemini_chat_completions_compat_service.go
index dcc3213b..ffea1595 100644
--- a/backend/internal/service/gemini_chat_completions_compat_service.go
+++ b/backend/internal/service/gemini_chat_completions_compat_service.go
@@ -151,7 +151,7 @@ func (s *GeminiMessagesCompatService) forwardClaudeBodyAsChatCompletions(
}
if resp.StatusCode >= 400 && s.shouldRetryGeminiUpstreamError(account, resp.StatusCode) {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
if resp.StatusCode == http.StatusForbidden && isGeminiInsufficientScope(resp.Header, respBody) {
resp = &http.Response{
@@ -207,7 +207,7 @@ func (s *GeminiMessagesCompatService) forwardClaudeBodyAsChatCompletions(
reasoningEffort := extractCCReasoningEffortFromBody(originalChatBody)
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
s.handleGeminiUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
evBody := unwrapIfNeeded(account.Type == AccountTypeOAuth, respBody)
diff --git a/backend/internal/service/gemini_messages_compat_service.go b/backend/internal/service/gemini_messages_compat_service.go
index 64f19b2e..86073d9c 100644
--- a/backend/internal/service/gemini_messages_compat_service.go
+++ b/backend/internal/service/gemini_messages_compat_service.go
@@ -56,6 +56,18 @@ type GeminiMessagesCompatService struct {
responseHeaderFilter *responseheaders.CompiledHeaderFilter
}
+func (s *GeminiMessagesCompatService) readUpstreamErrorBody(resp *http.Response) []byte {
+ if resp == nil || resp.Body == nil {
+ return nil
+ }
+ limit := gatewayUpstreamErrorBodyReadLimit
+ if s != nil && s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody && s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes > int(limit) {
+ limit = int64(s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes)
+ }
+ body, _ := io.ReadAll(io.LimitReader(resp.Body, limit))
+ return body
+}
+
func NewGeminiMessagesCompatService(
accountRepo AccountRepository,
groupRepo GroupRepository,
@@ -789,7 +801,7 @@ func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Contex
// Special-case: signature/thought_signature validation errors are not transient, but may be fixed by
// downgrading Claude thinking/tool history to plain text (conservative two-stage retry).
if resp.StatusCode == http.StatusBadRequest && signatureRetryStage < 2 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
if isGeminiSignatureRelatedError(respBody) {
@@ -860,7 +872,7 @@ func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Contex
}
if resp.StatusCode >= 400 && s.shouldRetryGeminiUpstreamError(account, resp.StatusCode) {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
// Don't treat insufficient-scope as transient.
if resp.StatusCode == 403 && isGeminiInsufficientScope(resp.Header, respBody) {
@@ -919,7 +931,7 @@ func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Contex
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
// 统一错误策略:自定义错误码 + 临时不可调度
if s.rateLimitService != nil {
switch s.rateLimitService.CheckErrorPolicy(ctx, account, resp.StatusCode, respBody) {
@@ -1329,7 +1341,7 @@ func (s *GeminiMessagesCompatService) ForwardNative(ctx context.Context, c *gin.
}
if resp.StatusCode >= 400 && s.shouldRetryGeminiUpstreamError(account, resp.StatusCode) {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
// Don't treat insufficient-scope as transient.
if resp.StatusCode == 403 && isGeminiInsufficientScope(resp.Header, respBody) {
@@ -1410,7 +1422,7 @@ func (s *GeminiMessagesCompatService) ForwardNative(ctx context.Context, c *gin.
isOAuth := account.Type == AccountTypeOAuth
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
// Best-effort fallback for OAuth tokens missing AI Studio scopes when calling countTokens.
// This avoids Gemini SDKs failing hard during preflight token counting.
// Checked before error policy so it always works regardless of custom error codes.
@@ -1619,7 +1631,7 @@ func (s *GeminiMessagesCompatService) checkErrorPolicyInLoop(
if resp.StatusCode < 400 || s.rateLimitService == nil {
return false, resp
}
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ body := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
rebuilt = &http.Response{
StatusCode: resp.StatusCode,
diff --git a/backend/internal/service/openai_embeddings.go b/backend/internal/service/openai_embeddings.go
index 359df3bb..7c710259 100644
--- a/backend/internal/service/openai_embeddings.go
+++ b/backend/internal/service/openai_embeddings.go
@@ -104,7 +104,7 @@ func (s *OpenAIGatewayService) ForwardEmbeddings(
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
diff --git a/backend/internal/service/openai_gateway_chat_completions.go b/backend/internal/service/openai_gateway_chat_completions.go
index 807ff43a..6e91d85c 100644
--- a/backend/internal/service/openai_gateway_chat_completions.go
+++ b/backend/internal/service/openai_gateway_chat_completions.go
@@ -241,7 +241,7 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions(
// 8. Handle error response with failover
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
diff --git a/backend/internal/service/openai_gateway_chat_completions_raw.go b/backend/internal/service/openai_gateway_chat_completions_raw.go
index e351fa75..3ff6fac4 100644
--- a/backend/internal/service/openai_gateway_chat_completions_raw.go
+++ b/backend/internal/service/openai_gateway_chat_completions_raw.go
@@ -183,7 +183,7 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions(
// 7. Handle error response with failover
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go
index 291c217e..4398bd27 100644
--- a/backend/internal/service/openai_gateway_messages.go
+++ b/backend/internal/service/openai_gateway_messages.go
@@ -300,7 +300,7 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
// 8. Handle error response with failover
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
diff --git a/backend/internal/service/openai_gateway_responses_chat_fallback.go b/backend/internal/service/openai_gateway_responses_chat_fallback.go
index cfab389a..205b27f7 100644
--- a/backend/internal/service/openai_gateway_responses_chat_fallback.go
+++ b/backend/internal/service/openai_gateway_responses_chat_fallback.go
@@ -163,7 +163,7 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions(
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go
index 2df2206d..d1460e5a 100644
--- a/backend/internal/service/openai_gateway_service.go
+++ b/backend/internal/service/openai_gateway_service.go
@@ -49,6 +49,8 @@ const (
// OpenAI WS Mode 失败后的重连次数上限(不含首次尝试)。
// 与 Codex 客户端保持一致:失败后最多重连 5 次。
openAIWSReconnectRetryLimit = 5
+ // 上游错误体只需要提取错误 JSON/日志摘要,默认 512KiB 避免错误风暴叠加大请求体。
+ openAIUpstreamErrorBodyReadLimit int64 = 512 << 10
// OpenAI WS Mode 重连退避默认值(可由配置覆盖)。
openAIWSRetryBackoffInitialDefault = 120 * time.Millisecond
openAIWSRetryBackoffMaxDefault = 2 * time.Second
@@ -2296,8 +2298,28 @@ func marshalOpenAIUpstreamJSON(v any) ([]byte, error) {
return out, nil
}
+func openAIUpstreamErrorBodyReadLimitForConfig(cfg *config.Config) int64 {
+ limit := openAIUpstreamErrorBodyReadLimit
+ if cfg != nil && cfg.Gateway.LogUpstreamErrorBody && cfg.Gateway.LogUpstreamErrorBodyMaxBytes > int(limit) {
+ limit = int64(cfg.Gateway.LogUpstreamErrorBodyMaxBytes)
+ }
+ return limit
+}
+
+func (s *OpenAIGatewayService) readUpstreamErrorBody(resp *http.Response) []byte {
+ if resp == nil || resp.Body == nil {
+ return nil
+ }
+ cfg := (*config.Config)(nil)
+ if s != nil {
+ cfg = s.cfg
+ }
+ body, _ := io.ReadAll(io.LimitReader(resp.Body, openAIUpstreamErrorBodyReadLimitForConfig(cfg)))
+ return body
+}
+
func (s *OpenAIGatewayService) handleFailoverSideEffects(ctx context.Context, resp *http.Response, account *Account, requestedModel ...string) {
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ body := s.readUpstreamErrorBody(resp)
if len(requestedModel) > 0 {
s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body, requestedModel[0])
return
@@ -3045,7 +3067,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
// Handle error response
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -3563,7 +3585,7 @@ func (s *OpenAIGatewayService) handleFailoverErrorResponsePassthrough(
account *Account,
requestBody []byte,
) error {
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ body := s.readUpstreamErrorBody(resp)
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body))
upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg)
@@ -3605,7 +3627,7 @@ func (s *OpenAIGatewayService) handleErrorResponsePassthrough(
account *Account,
requestBody []byte,
) error {
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ body := s.readUpstreamErrorBody(resp)
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body))
upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg)
@@ -4298,7 +4320,7 @@ func (s *OpenAIGatewayService) handleErrorResponse(
requestBody []byte,
requestedModel ...string,
) (*OpenAIForwardResult, error) {
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ body := s.readUpstreamErrorBody(resp)
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body))
upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg)
@@ -4459,7 +4481,7 @@ func (s *OpenAIGatewayService) handleCompatErrorResponse(
writeError compatErrorWriter,
requestedModel ...string,
) (*OpenAIForwardResult, error) {
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ body := s.readUpstreamErrorBody(resp)
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body))
if upstreamMsg == "" {
diff --git a/backend/internal/service/openai_images.go b/backend/internal/service/openai_images.go
index 1bcd947c..beb34780 100644
--- a/backend/internal/service/openai_images.go
+++ b/backend/internal/service/openai_images.go
@@ -622,7 +622,7 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesAPIKey(
return nil, fmt.Errorf("upstream request failed: %s", safeErr)
}
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(respBody))
@@ -1276,21 +1276,22 @@ func resolveOpenAIImageBytes(
headers http.Header,
conversationID string,
pointer openAIImagePointerInfo,
+ errorBodyReadLimit int64,
) ([]byte, error) {
if normalized := normalizeOpenAIImageBase64(pointer.B64JSON); normalized != "" {
return base64.StdEncoding.DecodeString(normalized)
}
if downloadURL := strings.TrimSpace(pointer.DownloadURL); downloadURL != "" {
- return downloadOpenAIImageBytes(ctx, client, headers, downloadURL)
+ return downloadOpenAIImageBytes(ctx, client, headers, downloadURL, errorBodyReadLimit)
}
if strings.TrimSpace(pointer.Pointer) == "" {
return nil, fmt.Errorf("image asset is missing pointer, url, and base64 data")
}
- downloadURL, err := fetchOpenAIImageDownloadURL(ctx, client, headers, conversationID, pointer.Pointer)
+ downloadURL, err := fetchOpenAIImageDownloadURL(ctx, client, headers, conversationID, pointer.Pointer, errorBodyReadLimit)
if err != nil {
return nil, err
}
- return downloadOpenAIImageBytes(ctx, client, headers, downloadURL)
+ return downloadOpenAIImageBytes(ctx, client, headers, downloadURL, errorBodyReadLimit)
}
func normalizeOpenAIImageBase64(raw string) string {
@@ -1395,6 +1396,7 @@ func fetchOpenAIImageDownloadURL(
headers http.Header,
conversationID string,
pointer string,
+ errorBodyReadLimit int64,
) (string, error) {
url := ""
allowConversationRetry := false
@@ -1425,7 +1427,7 @@ func fetchOpenAIImageDownloadURL(
} else if resp.IsSuccessState() && strings.TrimSpace(result.DownloadURL) != "" {
return strings.TrimSpace(result.DownloadURL), nil
} else {
- statusErr := newOpenAIImageStatusError(resp, "fetch image download url failed")
+ statusErr := newOpenAIImageStatusError(resp, "fetch image download url failed", errorBodyReadLimit)
if !allowConversationRetry || !isOpenAIImageTransientConversationNotFoundError(statusErr) {
return "", statusErr
}
@@ -1450,7 +1452,7 @@ func fetchOpenAIImageDownloadURL(
return "", lastErr
}
-func downloadOpenAIImageBytes(ctx context.Context, client *req.Client, headers http.Header, downloadURL string) ([]byte, error) {
+func downloadOpenAIImageBytes(ctx context.Context, client *req.Client, headers http.Header, downloadURL string, errorBodyReadLimit int64) ([]byte, error) {
request := client.R().
SetContext(ctx).
DisableAutoReadResponse()
@@ -1478,7 +1480,7 @@ func downloadOpenAIImageBytes(ctx context.Context, client *req.Client, headers h
}
}()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
- return nil, newOpenAIImageStatusError(resp, "download image bytes failed")
+ return nil, newOpenAIImageStatusError(resp, "download image bytes failed", errorBodyReadLimit)
}
return io.ReadAll(io.LimitReader(resp.Body, openAIImageMaxDownloadBytes))
}
@@ -1505,7 +1507,7 @@ func (e *openAIImageStatusError) Error() string {
return "openai image backend request failed"
}
-func newOpenAIImageStatusError(resp *req.Response, fallback string) error {
+func newOpenAIImageStatusError(resp *req.Response, fallback string, errorBodyReadLimit int64) error {
if resp == nil {
if strings.TrimSpace(fallback) == "" {
fallback = "openai image backend request failed"
@@ -1526,7 +1528,10 @@ func newOpenAIImageStatusError(resp *req.Response, fallback string) error {
requestURL = resp.Request.URL.String()
}
if resp.Body != nil {
- body, _ = io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ if errorBodyReadLimit <= 0 {
+ errorBodyReadLimit = openAIUpstreamErrorBodyReadLimit
+ }
+ body, _ = io.ReadAll(io.LimitReader(resp.Body, errorBodyReadLimit))
_ = resp.Body.Close()
}
}
diff --git a/backend/internal/service/openai_images_responses.go b/backend/internal/service/openai_images_responses.go
index 849ad792..db9c7b16 100644
--- a/backend/internal/service/openai_images_responses.go
+++ b/backend/internal/service/openai_images_responses.go
@@ -1172,7 +1172,7 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesOAuth(
return nil, fmt.Errorf("upstream request failed: %s", safeErr)
}
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(respBody))
diff --git a/backend/internal/service/openai_images_test.go b/backend/internal/service/openai_images_test.go
index a87e96c1..c3efdc93 100644
--- a/backend/internal/service/openai_images_test.go
+++ b/backend/internal/service/openai_images_test.go
@@ -4,6 +4,7 @@ import (
"bytes"
"context"
"errors"
+ "fmt"
"io"
"mime/multipart"
"net/http"
@@ -14,6 +15,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/gin-gonic/gin"
+ "github.com/imroc/req/v3"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
@@ -398,11 +400,38 @@ func TestCollectOpenAIImagePointers_RecognizesDirectAssets(t *testing.T) {
func TestResolveOpenAIImageBytes_PrefersInlineBase64(t *testing.T) {
data, err := resolveOpenAIImageBytes(context.Background(), nil, nil, "", openAIImagePointerInfo{
B64JSON: "data:image/png;base64,QUJD",
- })
+ }, openAIUpstreamErrorBodyReadLimit)
require.NoError(t, err)
require.Equal(t, []byte("ABC"), data)
}
+func TestNewOpenAIImageStatusError_UsesProvidedReadLimit(t *testing.T) {
+ padding := strings.Repeat("x", int(openAIUpstreamErrorBodyReadLimit)+1024)
+ body := fmt.Sprintf(`{"error":{"padding":"%s","message":"diagnostic-marker"}}`, padding)
+ resp := &req.Response{Response: &http.Response{
+ StatusCode: http.StatusBadGateway,
+ Header: http.Header{},
+ Body: io.NopCloser(strings.NewReader(body)),
+ }}
+
+ err := newOpenAIImageStatusError(resp, "download image bytes failed", int64(len(body)))
+ require.Error(t, err)
+ require.Equal(t, "diagnostic-marker", err.Error())
+
+ var statusErr *openAIImageStatusError
+ require.ErrorAs(t, err, &statusErr)
+ require.Len(t, statusErr.ResponseBody, len(body))
+}
+
+func TestOpenAIUpstreamErrorBodyReadLimitForConfig_RespectsDiagnosticLimit(t *testing.T) {
+ cfg := &config.Config{Gateway: config.GatewayConfig{
+ LogUpstreamErrorBody: true,
+ LogUpstreamErrorBodyMaxBytes: int(openAIUpstreamErrorBodyReadLimit) + 1024,
+ }}
+
+ require.Equal(t, int64(cfg.Gateway.LogUpstreamErrorBodyMaxBytes), openAIUpstreamErrorBodyReadLimitForConfig(cfg))
+}
+
func TestAccountSupportsOpenAIImageCapability_OAuthSupportsNative(t *testing.T) {
account := &Account{
Platform: PlatformOpenAI,
diff --git a/backend/internal/service/openai_tool_continuation.go b/backend/internal/service/openai_tool_continuation.go
index 36064549..92bc2613 100644
--- a/backend/internal/service/openai_tool_continuation.go
+++ b/backend/internal/service/openai_tool_continuation.go
@@ -157,7 +157,11 @@ func AnalyzeToolContinuationSignals(reqBody map[string]any) ToolContinuationSign
// ValidateFunctionCallOutputContextBytes 基于 raw JSON 校验工具输出续链,避免 handler 预校验阶段全量解码大 input。
func ValidateFunctionCallOutputContextBytes(body []byte) FunctionCallOutputValidation {
result := FunctionCallOutputValidation{}
- input := gjson.GetBytes(body, "input")
+ if len(body) == 0 {
+ return result
+ }
+ // handler 热路径只读扫描 input,避免 GetBytes 为大 Responses body 复制整段 JSON。
+ input := parseRawJSONView(body).Get("input")
if !input.IsArray() {
return result
}
From 6a5f6b96b6d438bd785ec14bda8840ba80ff6c29 Mon Sep 17 00:00:00 2001
From: name <136912576+is7Qin@users.noreply.github.com>
Date: Sun, 31 May 2026 01:01:58 +0800
Subject: [PATCH 55/79] refactor(gateway): introduce OpenAI request view
Cache hot-path request scalars before full body decoding so later branches can avoid repeated map work while preserving current decode behavior.
---
.../service/openai_gateway_service.go | 45 ++++++++++++++-----
.../openai_gateway_service_hotpath_test.go | 20 +++++++++
2 files changed, 55 insertions(+), 10 deletions(-)
diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go
index d1460e5a..87ffb2cb 100644
--- a/backend/internal/service/openai_gateway_service.go
+++ b/backend/internal/service/openai_gateway_service.go
@@ -2346,7 +2346,8 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
}
originalBody := body
- reqModel, reqStream, promptCacheKey := extractOpenAIRequestMetaFromBody(body)
+ requestView := newOpenAIRequestView(body)
+ reqModel, reqStream, promptCacheKey := requestView.Model, requestView.Stream, requestView.PromptCacheKey
originalModel := reqModel
if account.Type == AccountTypeAPIKey && !openai_compat.ShouldUseResponsesAPI(account.Extra) {
@@ -2396,7 +2397,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
return s.forwardOpenAIPassthrough(ctx, c, account, originalBody, reqModel, reasoningEffort, reqStream, startTime)
}
- reqBody, err := getOpenAIRequestBodyMap(c, body)
+ reqBody, err := requestView.Decode(c)
if err != nil {
return nil, err
}
@@ -6274,15 +6275,39 @@ func deriveOpenAIReasoningEffortFromModel(model string) string {
return normalizeOpenAIReasoningEffort(parts[len(parts)-1])
}
-func extractOpenAIRequestMetaFromBody(body []byte) (model string, stream bool, promptCacheKey string) {
- if len(body) == 0 {
- return "", false, ""
- }
+type openAIRequestView struct {
+ body []byte
+ Model string
+ Stream bool
+ PromptCacheKey string
+ PreviousResponseID string
+ ServiceTier string
+ ReasoningEffort string
+}
- model = strings.TrimSpace(gjson.GetBytes(body, "model").String())
- stream = gjson.GetBytes(body, "stream").Bool()
- promptCacheKey = strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String())
- return model, stream, promptCacheKey
+func newOpenAIRequestView(body []byte) openAIRequestView {
+ if len(body) == 0 {
+ return openAIRequestView{}
+ }
+ return openAIRequestView{
+ body: body,
+ Model: strings.TrimSpace(gjson.GetBytes(body, "model").String()),
+ Stream: gjson.GetBytes(body, "stream").Bool(),
+ PromptCacheKey: strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String()),
+ PreviousResponseID: strings.TrimSpace(gjson.GetBytes(body, "previous_response_id").String()),
+ ServiceTier: strings.TrimSpace(gjson.GetBytes(body, "service_tier").String()),
+ ReasoningEffort: strings.TrimSpace(gjson.GetBytes(body, "reasoning.effort").String()),
+ }
+}
+
+// Decode 保留阶段一既有 full-map 行为;后续阶段会把调用点下沉到复杂分支。
+func (v openAIRequestView) Decode(c *gin.Context) (map[string]any, error) {
+ return getOpenAIRequestBodyMap(c, v.body)
+}
+
+func extractOpenAIRequestMetaFromBody(body []byte) (model string, stream bool, promptCacheKey string) {
+ view := newOpenAIRequestView(body)
+ return view.Model, view.Stream, view.PromptCacheKey
}
// normalizeOpenAIPassthroughOAuthBody 将透传 OAuth 请求体收敛为旧链路关键行为:
diff --git a/backend/internal/service/openai_gateway_service_hotpath_test.go b/backend/internal/service/openai_gateway_service_hotpath_test.go
index af0d21c4..df17ed2e 100644
--- a/backend/internal/service/openai_gateway_service_hotpath_test.go
+++ b/backend/internal/service/openai_gateway_service_hotpath_test.go
@@ -9,6 +9,26 @@ import (
"github.com/stretchr/testify/require"
)
+func TestOpenAIRequestView_ExtractsRawScalars(t *testing.T) {
+ view := newOpenAIRequestView([]byte(`{"model":" gpt-5 ","stream":true,"prompt_cache_key":" ses-1 ","previous_response_id":" resp-1 ","service_tier":" fast ","reasoning":{"effort":" medium "}}`))
+
+ require.Equal(t, "gpt-5", view.Model)
+ require.True(t, view.Stream)
+ require.Equal(t, "ses-1", view.PromptCacheKey)
+ require.Equal(t, "resp-1", view.PreviousResponseID)
+ require.Equal(t, "fast", view.ServiceTier)
+ require.Equal(t, "medium", view.ReasoningEffort)
+}
+
+func TestOpenAIRequestView_DecodeKeepsFullMapBehavior(t *testing.T) {
+ view := newOpenAIRequestView([]byte(`{"model":"gpt-5","stream":true,"input":[{"type":"message","content":"hi"}]}`))
+
+ reqBody, err := view.Decode(nil)
+ require.NoError(t, err)
+ require.Equal(t, "gpt-5", reqBody["model"])
+ require.IsType(t, []any{}, reqBody["input"])
+}
+
func TestExtractOpenAIRequestMetaFromBody(t *testing.T) {
tests := []struct {
name string
From b65dde634bf7a5338f22b46ad0f9d98a197526fc Mon Sep 17 00:00:00 2001
From: wucm667
Date: Sun, 31 May 2026 08:39:37 +0800
Subject: [PATCH 56/79] =?UTF-8?q?fix(usage):=20=E4=BF=AE=E6=AD=A3=20OpenAI?=
=?UTF-8?q?=205h=20=E7=94=A8=E9=87=8F=E7=99=BE=E5=88=86=E6=AF=94=E8=AF=AD?=
=?UTF-8?q?=E4=B9=89?=
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit
---
.../account_test_service_openai_test.go | 4 +--
.../service/account_usage_service_test.go | 29 +++++++++++++++-
.../service/openai_gateway_service.go | 17 ++++++++--
...nai_gateway_service_codex_snapshot_test.go | 34 +++++++++++++++++++
.../service/openai_gateway_service_test.go | 2 +-
.../service/ratelimit_service_openai_test.go | 32 ++++++++---------
6 files changed, 96 insertions(+), 22 deletions(-)
diff --git a/backend/internal/service/account_test_service_openai_test.go b/backend/internal/service/account_test_service_openai_test.go
index 910567fb..970c723a 100644
--- a/backend/internal/service/account_test_service_openai_test.go
+++ b/backend/internal/service/account_test_service_openai_test.go
@@ -132,7 +132,7 @@ func TestAccountTestService_OpenAISuccessPersistsSnapshotFromHeaders(t *testing.
require.Len(t, upstream.requests, 1)
require.Equal(t, HTTPUpstreamProfileOpenAI, HTTPUpstreamProfileFromContext(upstream.requests[0].Context()))
require.NotEmpty(t, repo.updatedExtra)
- require.Equal(t, 42.0, repo.updatedExtra["codex_5h_used_percent"])
+ require.Equal(t, 58.0, repo.updatedExtra["codex_5h_used_percent"])
require.Equal(t, 88.0, repo.updatedExtra["codex_7d_used_percent"])
require.Contains(t, recorder.Body.String(), "test_complete")
}
@@ -170,7 +170,7 @@ func TestAccountTestService_OpenAI429PersistsSnapshotAndRateLimitState(t *testin
resp.Header.Set("x-codex-primary-used-percent", "100")
resp.Header.Set("x-codex-primary-reset-after-seconds", "604800")
resp.Header.Set("x-codex-primary-window-minutes", "10080")
- resp.Header.Set("x-codex-secondary-used-percent", "100")
+ resp.Header.Set("x-codex-secondary-used-percent", "0")
resp.Header.Set("x-codex-secondary-reset-after-seconds", "18000")
resp.Header.Set("x-codex-secondary-window-minutes", "300")
diff --git a/backend/internal/service/account_usage_service_test.go b/backend/internal/service/account_usage_service_test.go
index e0390c4c..5f37aadb 100644
--- a/backend/internal/service/account_usage_service_test.go
+++ b/backend/internal/service/account_usage_service_test.go
@@ -73,7 +73,7 @@ func TestExtractOpenAICodexProbeUpdatesAccepts429WithCodexHeaders(t *testing.T)
headers.Set("x-codex-primary-used-percent", "100")
headers.Set("x-codex-primary-reset-after-seconds", "604800")
headers.Set("x-codex-primary-window-minutes", "10080")
- headers.Set("x-codex-secondary-used-percent", "100")
+ headers.Set("x-codex-secondary-used-percent", "0")
headers.Set("x-codex-secondary-reset-after-seconds", "18000")
headers.Set("x-codex-secondary-window-minutes", "300")
@@ -92,6 +92,33 @@ func TestExtractOpenAICodexProbeUpdatesAccepts429WithCodexHeaders(t *testing.T)
}
}
+func TestBuildCodexUsageProgressFromExtra_UsesCanonicalUsedPercent(t *testing.T) {
+ t.Parallel()
+ now := time.Date(2026, 5, 30, 7, 4, 9, 0, time.UTC)
+ extra := map[string]any{
+ "codex_5h_used_percent": 94.0,
+ "codex_5h_reset_at": now.Add(2 * time.Hour).Format(time.RFC3339),
+ "codex_7d_used_percent": 93.0,
+ "codex_7d_reset_at": now.Add(5 * 24 * time.Hour).Format(time.RFC3339),
+ }
+
+ fiveHour := buildCodexUsageProgressFromExtra(extra, "5h", now)
+ if fiveHour == nil {
+ t.Fatal("expected non-nil 5h progress")
+ }
+ if fiveHour.Utilization != 94.0 {
+ t.Fatalf("5h Utilization = %v, want 94", fiveHour.Utilization)
+ }
+
+ sevenDay := buildCodexUsageProgressFromExtra(extra, "7d", now)
+ if sevenDay == nil {
+ t.Fatal("expected non-nil 7d progress")
+ }
+ if sevenDay.Utilization != 93.0 {
+ t.Fatalf("7d Utilization = %v, want 93", sevenDay.Utilization)
+ }
+}
+
func TestAccountUsageService_PersistOpenAICodexProbeSnapshotOnlyUpdatesExtra(t *testing.T) {
t.Parallel()
diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go
index cd5a4015..c9e88bde 100644
--- a/backend/internal/service/openai_gateway_service.go
+++ b/backend/internal/service/openai_gateway_service.go
@@ -126,6 +126,19 @@ type NormalizedCodexLimits struct {
Window7dMinutes *int
}
+func normalizeCodexFiveHourUsedPercent(raw *float64) *float64 {
+ if raw == nil {
+ return nil
+ }
+ // OpenAI's 5h Codex quota header is remaining%, despite the upstream header
+ // name saying "used"; the canonical codex_5h_used_percent field stores used%.
+ used := 100 - *raw
+ if used < 0 {
+ used = 0
+ }
+ return &used
+}
+
// Normalize converts primary/secondary fields to canonical 5h/7d fields.
// Strategy: Compare window_minutes to determine which is 5h vs 7d.
// Returns nil if snapshot is nil or has no useful data.
@@ -184,7 +197,7 @@ func (s *OpenAICodexUsageSnapshot) Normalize() *NormalizedCodexLimits {
// Assign values
if use5hFromPrimary {
- result.Used5hPercent = s.PrimaryUsedPercent
+ result.Used5hPercent = normalizeCodexFiveHourUsedPercent(s.PrimaryUsedPercent)
result.Reset5hSeconds = s.PrimaryResetAfterSeconds
result.Window5hMinutes = s.PrimaryWindowMinutes
result.Used7dPercent = s.SecondaryUsedPercent
@@ -194,7 +207,7 @@ func (s *OpenAICodexUsageSnapshot) Normalize() *NormalizedCodexLimits {
result.Used7dPercent = s.PrimaryUsedPercent
result.Reset7dSeconds = s.PrimaryResetAfterSeconds
result.Window7dMinutes = s.PrimaryWindowMinutes
- result.Used5hPercent = s.SecondaryUsedPercent
+ result.Used5hPercent = normalizeCodexFiveHourUsedPercent(s.SecondaryUsedPercent)
result.Reset5hSeconds = s.SecondaryResetAfterSeconds
result.Window5hMinutes = s.SecondaryWindowMinutes
}
diff --git a/backend/internal/service/openai_gateway_service_codex_snapshot_test.go b/backend/internal/service/openai_gateway_service_codex_snapshot_test.go
index 654dd4ca..22f5fa74 100644
--- a/backend/internal/service/openai_gateway_service_codex_snapshot_test.go
+++ b/backend/internal/service/openai_gateway_service_codex_snapshot_test.go
@@ -104,6 +104,40 @@ func TestBuildCodexUsageExtraUpdates_UsesSnapshotUpdatedAt(t *testing.T) {
}
}
+func TestBuildCodexUsageExtraUpdates_NormalizesFiveHourRemainingToUsedPercent(t *testing.T) {
+ primaryUsed := 93.0
+ primaryReset := 86400
+ primaryWindow := 10080
+ secondaryRemaining := 6.0
+ secondaryReset := 3600
+ secondaryWindow := 300
+
+ snapshot := &OpenAICodexUsageSnapshot{
+ PrimaryUsedPercent: &primaryUsed,
+ PrimaryResetAfterSeconds: &primaryReset,
+ PrimaryWindowMinutes: &primaryWindow,
+ SecondaryUsedPercent: &secondaryRemaining,
+ SecondaryResetAfterSeconds: &secondaryReset,
+ SecondaryWindowMinutes: &secondaryWindow,
+ UpdatedAt: "2026-05-30T07:04:09Z",
+ }
+
+ updates := buildCodexUsageExtraUpdates(snapshot, time.Time{})
+ if updates == nil {
+ t.Fatal("expected non-nil updates")
+ }
+
+ if got := updates["codex_secondary_used_percent"]; got != 6.0 {
+ t.Fatalf("codex_secondary_used_percent = %v, want raw upstream value 6", got)
+ }
+ if got := updates["codex_5h_used_percent"]; got != 94.0 {
+ t.Fatalf("codex_5h_used_percent = %v, want 94", got)
+ }
+ if got := updates["codex_7d_used_percent"]; got != 93.0 {
+ t.Fatalf("codex_7d_used_percent = %v, want 93", got)
+ }
+}
+
func TestBuildCodexUsageExtraUpdates_FallbackToNowWhenUpdatedAtInvalid(t *testing.T) {
primaryUsed := 15.0
primaryReset := 30
diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go
index 8aad2fa6..5c4e979d 100644
--- a/backend/internal/service/openai_gateway_service_test.go
+++ b/backend/internal/service/openai_gateway_service_test.go
@@ -1774,7 +1774,7 @@ func TestOpenAIUpdateCodexUsageSnapshotFromHeaders(t *testing.T) {
select {
case updates := <-repo.updateExtraCalls:
- require.Equal(t, 12.0, updates["codex_5h_used_percent"])
+ require.Equal(t, 88.0, updates["codex_5h_used_percent"])
require.Equal(t, 34.0, updates["codex_7d_used_percent"])
require.Equal(t, 600, updates["codex_5h_reset_after_seconds"])
require.Equal(t, 86400, updates["codex_7d_reset_after_seconds"])
diff --git a/backend/internal/service/ratelimit_service_openai_test.go b/backend/internal/service/ratelimit_service_openai_test.go
index aa5a070c..107ac27e 100644
--- a/backend/internal/service/ratelimit_service_openai_test.go
+++ b/backend/internal/service/ratelimit_service_openai_test.go
@@ -51,7 +51,7 @@ func TestCalculateOpenAI429ResetTime_5hExhausted(t *testing.T) {
headers.Set("x-codex-primary-used-percent", "50")
headers.Set("x-codex-primary-reset-after-seconds", "500000")
headers.Set("x-codex-primary-window-minutes", "10080") // 7 days
- headers.Set("x-codex-secondary-used-percent", "100")
+ headers.Set("x-codex-secondary-used-percent", "0")
headers.Set("x-codex-secondary-reset-after-seconds", "3600") // 1 hour
headers.Set("x-codex-secondary-window-minutes", "300") // 5 hours
@@ -122,7 +122,7 @@ func TestCalculateOpenAI429ResetTime_ReversedWindowOrder(t *testing.T) {
// Test when OpenAI sends primary as 5h and secondary as 7d (reversed)
headers := http.Header{}
- headers.Set("x-codex-primary-used-percent", "100") // This is 5h
+ headers.Set("x-codex-primary-used-percent", "0") // This is 5h remaining%
headers.Set("x-codex-primary-reset-after-seconds", "3600") // 1 hour
headers.Set("x-codex-primary-window-minutes", "300") // 5 hours - smaller!
headers.Set("x-codex-secondary-used-percent", "50")
@@ -180,7 +180,7 @@ func TestHandle429_OpenAIPersistsCodexSnapshotImmediately(t *testing.T) {
headers.Set("x-codex-primary-used-percent", "100")
headers.Set("x-codex-primary-reset-after-seconds", "604800")
headers.Set("x-codex-primary-window-minutes", "10080")
- headers.Set("x-codex-secondary-used-percent", "100")
+ headers.Set("x-codex-secondary-used-percent", "0")
headers.Set("x-codex-secondary-reset-after-seconds", "18000")
headers.Set("x-codex-secondary-window-minutes", "300")
@@ -224,7 +224,7 @@ func TestNormalizedCodexLimits(t *testing.T) {
pUsed := 100.0
pReset := 384607
pWindow := 10080
- sUsed := 3.0
+ sRemaining := 3.0
sReset := 17369
sWindow := 300
@@ -232,7 +232,7 @@ func TestNormalizedCodexLimits(t *testing.T) {
PrimaryUsedPercent: &pUsed,
PrimaryResetAfterSeconds: &pReset,
PrimaryWindowMinutes: &pWindow,
- SecondaryUsedPercent: &sUsed,
+ SecondaryUsedPercent: &sRemaining,
SecondaryResetAfterSeconds: &sReset,
SecondaryWindowMinutes: &sWindow,
}
@@ -249,8 +249,8 @@ func TestNormalizedCodexLimits(t *testing.T) {
if normalized.Reset7dSeconds == nil || *normalized.Reset7dSeconds != 384607 {
t.Errorf("expected Reset7dSeconds=384607, got %v", normalized.Reset7dSeconds)
}
- if normalized.Used5hPercent == nil || *normalized.Used5hPercent != 3.0 {
- t.Errorf("expected Used5hPercent=3, got %v", normalized.Used5hPercent)
+ if normalized.Used5hPercent == nil || *normalized.Used5hPercent != 97.0 {
+ t.Errorf("expected Used5hPercent=97, got %v", normalized.Used5hPercent)
}
if normalized.Reset5hSeconds == nil || *normalized.Reset5hSeconds != 17369 {
t.Errorf("expected Reset5hSeconds=17369, got %v", normalized.Reset5hSeconds)
@@ -338,11 +338,11 @@ func TestRateLimitService_HandleUpstreamError_403FallsBackToRawBody(t *testing.T
func TestNormalizedCodexLimits_OnlySecondaryData(t *testing.T) {
// Test when only secondary has data, no window_minutes
- sUsed := 60.0
+ sRemaining := 60.0
sReset := 3000
snapshot := &OpenAICodexUsageSnapshot{
- SecondaryUsedPercent: &sUsed,
+ SecondaryUsedPercent: &sRemaining,
SecondaryResetAfterSeconds: &sReset,
// No window_minutes, no primary data
}
@@ -354,8 +354,8 @@ func TestNormalizedCodexLimits_OnlySecondaryData(t *testing.T) {
// Legacy assumption: primary=7d, secondary=5h
// So secondary goes to 5h
- if normalized.Used5hPercent == nil || *normalized.Used5hPercent != 60.0 {
- t.Errorf("expected Used5hPercent=60, got %v", normalized.Used5hPercent)
+ if normalized.Used5hPercent == nil || *normalized.Used5hPercent != 40.0 {
+ t.Errorf("expected Used5hPercent=40, got %v", normalized.Used5hPercent)
}
if normalized.Reset5hSeconds == nil || *normalized.Reset5hSeconds != 3000 {
t.Errorf("expected Reset5hSeconds=3000, got %v", normalized.Reset5hSeconds)
@@ -370,13 +370,13 @@ func TestNormalizedCodexLimits_BothDataNoWindowMinutes(t *testing.T) {
// Test when both have data but no window_minutes
pUsed := 100.0
pReset := 400000
- sUsed := 50.0
+ sRemaining := 30.0
sReset := 10000
snapshot := &OpenAICodexUsageSnapshot{
PrimaryUsedPercent: &pUsed,
PrimaryResetAfterSeconds: &pReset,
- SecondaryUsedPercent: &sUsed,
+ SecondaryUsedPercent: &sRemaining,
SecondaryResetAfterSeconds: &sReset,
// No window_minutes
}
@@ -393,8 +393,8 @@ func TestNormalizedCodexLimits_BothDataNoWindowMinutes(t *testing.T) {
if normalized.Reset7dSeconds == nil || *normalized.Reset7dSeconds != 400000 {
t.Errorf("expected Reset7dSeconds=400000, got %v", normalized.Reset7dSeconds)
}
- if normalized.Used5hPercent == nil || *normalized.Used5hPercent != 50.0 {
- t.Errorf("expected Used5hPercent=50, got %v", normalized.Used5hPercent)
+ if normalized.Used5hPercent == nil || *normalized.Used5hPercent != 70.0 {
+ t.Errorf("expected Used5hPercent=70, got %v", normalized.Used5hPercent)
}
if normalized.Reset5hSeconds == nil || *normalized.Reset5hSeconds != 10000 {
t.Errorf("expected Reset5hSeconds=10000, got %v", normalized.Reset5hSeconds)
@@ -425,7 +425,7 @@ func TestCalculateOpenAI429ResetTime_UserProvidedScenario(t *testing.T) {
// This is the exact scenario from the user:
// codex_7d_used_percent: 100
// codex_7d_reset_after_seconds: 384607 (约4.5天后重置)
- // codex_5h_used_percent: 3
+ // codex_5h_used_percent: 97 (from upstream 3% remaining)
// codex_5h_reset_after_seconds: 17369 (约4.8小时后重置)
svc := &RateLimitService{}
From bf3787de1f5842d869a74bdb59af92d57288a35d Mon Sep 17 00:00:00 2001
From: wucm667
Date: Sun, 31 May 2026 08:43:20 +0800
Subject: [PATCH 57/79] fix(gateway): allow Claude Code count_tokens
---
.../internal/service/claude_code_validator.go | 13 ++-
.../service/claude_code_validator_test.go | 91 +++++++++++++++++++
2 files changed, 102 insertions(+), 2 deletions(-)
diff --git a/backend/internal/service/claude_code_validator.go b/backend/internal/service/claude_code_validator.go
index 4e8ced67..2c5ded6f 100644
--- a/backend/internal/service/claude_code_validator.go
+++ b/backend/internal/service/claude_code_validator.go
@@ -56,7 +56,7 @@ func NewClaudeCodeValidator() *ClaudeCodeValidator {
// 采用与 claude-relay-service 完全一致的验证策略:
//
// Step 1: User-Agent 检查 (必需) - 必须是 claude-cli/x.x.x
-// Step 2: 对于非 messages 路径,只要 UA 匹配就通过
+// Step 2: 对于非 messages 路径和 /messages/count_tokens,只要 UA 匹配就通过
// Step 3: 检查 max_tokens=1 + haiku 探测请求绕过(UA 已验证)
// Step 4: 对于 messages 路径,进行严格验证:
// - System prompt 相似度检查
@@ -71,12 +71,17 @@ func (v *ClaudeCodeValidator) Validate(r *http.Request, body map[string]any) boo
return false
}
- // Step 2: 非 messages 路径,只要 UA 匹配就通过
+ // Step 2: 非 messages 路径只要 UA 匹配就通过
path := r.URL.Path
if !strings.Contains(path, "messages") {
return true
}
+ // count_tokens 是 Claude Code 官方辅助请求,通常不携带完整 messages system prompt。
+ if isMessagesCountTokensPath(path) {
+ return true
+ }
+
// Step 3: 检查 max_tokens=1 + haiku 探测请求绕过
// 这类请求用于 Claude Code 验证 API 连通性,不携带 system prompt
if isMaxTokensOneHaiku, ok := IsMaxTokensOneHaikuRequestFromContext(r.Context()); ok && isMaxTokensOneHaiku {
@@ -128,6 +133,10 @@ func (v *ClaudeCodeValidator) Validate(r *http.Request, body map[string]any) boo
return true
}
+func isMessagesCountTokensPath(path string) bool {
+ return strings.HasSuffix(path, "/messages/count_tokens")
+}
+
// hasClaudeCodeSystemPrompt 检查请求是否包含 Claude Code 系统提示词
// 使用字符串相似度匹配(Dice coefficient)
func (v *ClaudeCodeValidator) hasClaudeCodeSystemPrompt(body map[string]any) bool {
diff --git a/backend/internal/service/claude_code_validator_test.go b/backend/internal/service/claude_code_validator_test.go
index f87c56e8..a4b30505 100644
--- a/backend/internal/service/claude_code_validator_test.go
+++ b/backend/internal/service/claude_code_validator_test.go
@@ -48,6 +48,97 @@ func TestClaudeCodeValidator_MessagesWithoutProbeStillNeedStrictValidation(t *te
require.False(t, ok)
}
+func TestClaudeCodeValidator_CountTokensPathUAOnly(t *testing.T) {
+ validator := NewClaudeCodeValidator()
+ req := httptest.NewRequest(http.MethodPost, "http://example.com/v1/messages/count_tokens", nil)
+ req.Header.Set("User-Agent", "claude-cli/2.1.156 (Claude Code)")
+
+ ok := validator.Validate(req, map[string]any{
+ "model": "claude-opus-4-8",
+ })
+ require.True(t, ok)
+}
+
+func TestClaudeCodeValidator_CountTokensPathRequiresUA(t *testing.T) {
+ validator := NewClaudeCodeValidator()
+ req := httptest.NewRequest(http.MethodPost, "http://example.com/v1/messages/count_tokens", nil)
+ req.Header.Set("User-Agent", "curl/8.0.0")
+
+ ok := validator.Validate(req, map[string]any{
+ "model": "claude-opus-4-8",
+ })
+ require.False(t, ok)
+}
+
+func TestClaudeCodeValidator_MessagesPathFullValid(t *testing.T) {
+ validator := NewClaudeCodeValidator()
+ req := httptest.NewRequest(http.MethodPost, "http://example.com/v1/messages", nil)
+ req.Header.Set("User-Agent", "claude-cli/2.1.156 (Claude Code)")
+ req.Header.Set("X-App", "claude-code")
+ req.Header.Set("anthropic-beta", "claude-code-20250219")
+ req.Header.Set("anthropic-version", "2023-06-01")
+
+ ok := validator.Validate(req, map[string]any{
+ "model": "claude-opus-4-8",
+ "stream": true,
+ "system": []any{
+ map[string]any{
+ "type": "text",
+ "text": "You are Claude Code, Anthropic's official CLI for Claude.",
+ },
+ },
+ "metadata": map[string]any{
+ "user_id": "user_aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa_account__session_aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa",
+ },
+ })
+ require.True(t, ok)
+}
+
+func TestClaudeCodeValidator_MessagesPathRejectsNonClaudeCodeUA(t *testing.T) {
+ validator := NewClaudeCodeValidator()
+ req := httptest.NewRequest(http.MethodPost, "http://example.com/v1/messages", nil)
+ req.Header.Set("User-Agent", "curl/8.0.0")
+ req.Header.Set("X-App", "claude-code")
+ req.Header.Set("anthropic-beta", "claude-code-20250219")
+ req.Header.Set("anthropic-version", "2023-06-01")
+
+ ok := validator.Validate(req, map[string]any{
+ "model": "claude-opus-4-8",
+ "stream": true,
+ "system": []any{
+ map[string]any{
+ "type": "text",
+ "text": "You are Claude Code, Anthropic's official CLI for Claude.",
+ },
+ },
+ "metadata": map[string]any{
+ "user_id": "user_aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa_account__session_aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa",
+ },
+ })
+ require.False(t, ok)
+}
+
+func TestClaudeCodeValidator_MessagesPathWithoutSystemPromptStillRejected(t *testing.T) {
+ validator := NewClaudeCodeValidator()
+ req := httptest.NewRequest(http.MethodPost, "http://example.com/v1/messages", nil)
+ req.Header.Set("User-Agent", "claude-cli/2.1.156 (Claude Code)")
+ req.Header.Set("X-App", "claude-code")
+ req.Header.Set("anthropic-beta", "claude-code-20250219")
+ req.Header.Set("anthropic-version", "2023-06-01")
+
+ ok := validator.Validate(req, map[string]any{
+ "model": "claude-opus-4-8",
+ "stream": true,
+ "messages": []any{
+ map[string]any{"role": "user", "content": "hello"},
+ },
+ "metadata": map[string]any{
+ "user_id": "user_aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa_account__session_aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa",
+ },
+ })
+ require.False(t, ok)
+}
+
func TestClaudeCodeValidator_NonMessagesPathUAOnly(t *testing.T) {
validator := NewClaudeCodeValidator()
req := httptest.NewRequest(http.MethodPost, "http://example.com/v1/models", nil)
From 1e2193c3d27b7d770158fc7e133772b218bbe4dd Mon Sep 17 00:00:00 2001
From: gsh
Date: Sun, 31 May 2026 15:09:06 +0800
Subject: [PATCH 58/79] fix: avoid websocket usage dedup conflicts
---
.../openai_gateway_record_usage_test.go | 31 +++++++++++++++++++
.../service/openai_gateway_service.go | 5 +++
2 files changed, 36 insertions(+)
diff --git a/backend/internal/service/openai_gateway_record_usage_test.go b/backend/internal/service/openai_gateway_record_usage_test.go
index 9769a82e..318c0861 100644
--- a/backend/internal/service/openai_gateway_record_usage_test.go
+++ b/backend/internal/service/openai_gateway_record_usage_test.go
@@ -721,6 +721,37 @@ func TestOpenAIGatewayServiceRecordUsage_PrefersClientRequestIDOverUpstreamReque
require.Equal(t, "client:openai-client-stable-123", usageRepo.lastLog.RequestID)
}
+func TestOpenAIGatewayServiceRecordUsage_WSModePrefersUpstreamRequestIDOverClientRequestID(t *testing.T) {
+ usageRepo := &openAIRecordUsageLogRepoStub{}
+ billingRepo := &openAIRecordUsageBillingRepoStub{result: &UsageBillingApplyResult{Applied: true}}
+ userRepo := &openAIRecordUsageUserRepoStub{}
+ subRepo := &openAIRecordUsageSubRepoStub{}
+ svc := newOpenAIRecordUsageServiceWithBillingRepoForTest(usageRepo, billingRepo, userRepo, subRepo, nil)
+
+ ctx := context.WithValue(context.Background(), ctxkey.ClientRequestID, "openai-ws-connection-123")
+ err := svc.RecordUsage(ctx, &OpenAIRecordUsageInput{
+ Result: &OpenAIForwardResult{
+ RequestID: "resp_openai_ws_turn_456",
+ OpenAIWSMode: true,
+ Usage: OpenAIUsage{
+ InputTokens: 8,
+ OutputTokens: 4,
+ },
+ Model: "gpt-5.1",
+ Duration: time.Second,
+ },
+ APIKey: &APIKey{ID: 10050},
+ User: &User{ID: 20050},
+ Account: &Account{ID: 30050},
+ })
+
+ require.NoError(t, err)
+ require.NotNil(t, billingRepo.lastCmd)
+ require.Equal(t, "resp_openai_ws_turn_456", billingRepo.lastCmd.RequestID)
+ require.NotNil(t, usageRepo.lastLog)
+ require.Equal(t, "resp_openai_ws_turn_456", usageRepo.lastLog.RequestID)
+}
+
func TestOpenAIGatewayServiceRecordUsage_GeneratesRequestIDWhenAllSourcesMissing(t *testing.T) {
usageRepo := &openAIRecordUsageLogRepoStub{}
billingRepo := &openAIRecordUsageBillingRepoStub{result: &UsageBillingApplyResult{Applied: true}}
diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go
index cd5a4015..10080f31 100644
--- a/backend/internal/service/openai_gateway_service.go
+++ b/backend/internal/service/openai_gateway_service.go
@@ -5761,6 +5761,11 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec
durationMs := int(result.Duration.Milliseconds())
accountRateMultiplier := account.BillingRateMultiplier()
requestID := resolveUsageBillingRequestID(ctx, result.RequestID)
+ if result.OpenAIWSMode {
+ if upstreamRequestID := strings.TrimSpace(result.RequestID); upstreamRequestID != "" {
+ requestID = upstreamRequestID
+ }
+ }
// 确定 RequestedModel(渠道映射前的原始模型)
requestedModel := result.Model
From 8ac2e23fb3e16323cea0a6d7dea8986467487486 Mon Sep 17 00:00:00 2001
From: name <136912576+is7Qin@users.noreply.github.com>
Date: Sun, 31 May 2026 16:00:35 +0800
Subject: [PATCH 59/79] refactor(gateway): defer OpenAI request map decoding
Keep the Responses hot path on raw request bytes until complex mutation or retry branches need a decoded map, reducing large-body retention under high concurrency.
---
.../service/gateway_service_benchmark_test.go | 75 ++
.../service/image_generation_intent.go | 90 ++-
.../service/image_generation_intent_test.go | 11 +
.../service/openai_gateway_service.go | 701 ++++++++++--------
.../openai_gateway_service_hotpath_test.go | 454 ++++++++++++
.../openai_ws_forwarder_success_test.go | 85 +++
6 files changed, 1093 insertions(+), 323 deletions(-)
diff --git a/backend/internal/service/gateway_service_benchmark_test.go b/backend/internal/service/gateway_service_benchmark_test.go
index b2df1435..f6cd1404 100644
--- a/backend/internal/service/gateway_service_benchmark_test.go
+++ b/backend/internal/service/gateway_service_benchmark_test.go
@@ -124,6 +124,64 @@ func BenchmarkOpenAIResponses_LargeInputDecodeMap(b *testing.B) {
}
}
+func BenchmarkOpenAIResponses_LargeInputRawPatch(b *testing.B) {
+ for _, size := range benchmarkBodySizes() {
+ b.Run(size.name, func(b *testing.B) {
+ body := buildLargeOpenAIResponsesBody(size.bytes)
+
+ b.SetBytes(int64(len(body)))
+ b.ReportAllocs()
+ b.ResetTimer()
+ for i := 0; i < b.N; i++ {
+ view := newOpenAIRequestView(body)
+ view.MarkPatchSet("instructions", "You are a helpful coding assistant.")
+ view.MarkPatchSet("reasoning.effort", "none")
+ patched, err := view.ApplyPatches()
+ if err != nil {
+ b.Fatalf("应用 OpenAI raw patch 失败: %v", err)
+ }
+ benchmarkIntSink = len(patched)
+ }
+ })
+ }
+}
+
+func BenchmarkOpenAIResponses_LargeInputImageBillingRaw(b *testing.B) {
+ for _, size := range benchmarkBodySizes() {
+ b.Run(size.name, func(b *testing.B) {
+ body := buildLargeOpenAIResponsesImageToolBody(size.bytes)
+
+ b.SetBytes(int64(len(body)))
+ b.ReportAllocs()
+ b.ResetTimer()
+ for i := 0; i < b.N; i++ {
+ cfg, err := resolveOpenAIResponsesImageBillingConfigDetailedFromBody(body, "gpt-5.4")
+ if err != nil {
+ b.Fatalf("解析 OpenAI 图片计费配置失败: %v", err)
+ }
+ benchmarkStringSink = cfg.Model + cfg.SizeTier + cfg.InputSize
+ }
+ })
+ }
+}
+
+func BenchmarkOpenAIResponses_LargeInputEmptyBase64Guard(b *testing.B) {
+ for _, size := range benchmarkBodySizes() {
+ b.Run(size.name, func(b *testing.B) {
+ body := buildLargeOpenAIResponsesBody(size.bytes)
+
+ b.SetBytes(int64(len(body)))
+ b.ReportAllocs()
+ b.ResetTimer()
+ for i := 0; i < b.N; i++ {
+ if openAIRequestBodyMayContainEmptyBase64InputImage(body) {
+ benchmarkIntSink++
+ }
+ }
+ })
+ }
+}
+
func BenchmarkOpenAIResponses_LargeInputFunctionCallValidation(b *testing.B) {
for _, size := range benchmarkBodySizes() {
b.Run(size.name, func(b *testing.B) {
@@ -264,3 +322,20 @@ func buildLargeOpenAIResponsesToolContinuationBody(targetBytes int) []byte {
builder.WriteString(`]}`)
return []byte(builder.String())
}
+
+func buildLargeOpenAIResponsesImageToolBody(targetBytes int) []byte {
+ var builder strings.Builder
+ builder.Grow(targetBytes + 1024)
+ builder.WriteString(`{"model":"gpt-5.4","stream":false,"tools":[{"type":"image_generation","model":"gpt-image-2","size":"2048x1152"}],"input":[`)
+ for i := 0; builder.Len() < targetBytes; i++ {
+ if i > 0 {
+ builder.WriteByte(',')
+ }
+ builder.WriteString(`{"type":"message","role":"user","content":[{"type":"input_text","text":"`)
+ builder.WriteString(strings.Repeat("openai image billing payload ", 48))
+ builder.WriteString(strconv.Itoa(i))
+ builder.WriteString(`"}]}`)
+ }
+ builder.WriteString(`]}`)
+ return []byte(builder.String())
+}
diff --git a/backend/internal/service/image_generation_intent.go b/backend/internal/service/image_generation_intent.go
index 4aca1239..80b6c66d 100644
--- a/backend/internal/service/image_generation_intent.go
+++ b/backend/internal/service/image_generation_intent.go
@@ -1,7 +1,6 @@
package service
import (
- "encoding/json"
"strings"
"github.com/tidwall/gjson"
@@ -91,7 +90,7 @@ func openAIJSONToolsContainImageGeneration(tools gjson.Result) bool {
}
found := false
tools.ForEach(func(_, item gjson.Result) bool {
- if strings.TrimSpace(item.Get("type").String()) == "image_generation" {
+ if openAIJSONString(item.Get("type")) == "image_generation" {
found = true
return false
}
@@ -100,6 +99,36 @@ func openAIJSONToolsContainImageGeneration(tools gjson.Result) bool {
return found
}
+func openAIRequestBodyHasImageGenerationTool(body []byte) bool {
+ if len(body) == 0 || !gjson.ValidBytes(body) {
+ return false
+ }
+ return openAIJSONToolsContainImageGeneration(gjson.GetBytes(body, "tools"))
+}
+
+func openAIRequestBodyImageGenerationToolNeedsNormalization(body []byte) bool {
+ if len(body) == 0 || !gjson.ValidBytes(body) {
+ return false
+ }
+ tools := gjson.GetBytes(body, "tools")
+ if !tools.IsArray() {
+ return false
+ }
+ needsNormalization := false
+ tools.ForEach(func(_, item gjson.Result) bool {
+ if openAIJSONString(item.Get("type")) != "image_generation" {
+ return true
+ }
+ // 只有旧字段需要迁移时才进入 map 修改,纯计费读取保持 raw 路径。
+ if item.Get("format").Exists() || item.Get("compression").Exists() {
+ needsNormalization = true
+ return false
+ }
+ return true
+ })
+ return needsNormalization
+}
+
func openAIJSONToolChoiceSelectsImageGeneration(choice gjson.Result) bool {
if !choice.Exists() {
return false
@@ -159,17 +188,6 @@ func apiKeyGroup(apiKey *APIKey) *Group {
return apiKey.Group
}
-func cloneRequestMapForImageIntent(body []byte) map[string]any {
- if len(body) == 0 {
- return nil
- }
- var out map[string]any
- if err := json.Unmarshal(body, &out); err != nil {
- return nil
- }
- return out
-}
-
type OpenAIResponsesImageBillingConfig struct {
Model string
SizeTier string
@@ -225,8 +243,43 @@ func resolveOpenAIResponsesImageBillingConfigFromBody(body []byte, fallbackModel
}
func resolveOpenAIResponsesImageBillingConfigDetailedFromBody(body []byte, fallbackModel string) (OpenAIResponsesImageBillingConfig, error) {
- reqBody := cloneRequestMapForImageIntent(body)
- return resolveOpenAIResponsesImageBillingConfigDetailed(reqBody, fallbackModel)
+ imageModel := ""
+ imageSize := ""
+ hasImageTool := false
+ if len(body) > 0 && gjson.ValidBytes(body) {
+ tools := gjson.GetBytes(body, "tools")
+ if tools.IsArray() {
+ tools.ForEach(func(_, item gjson.Result) bool {
+ if openAIJSONString(item.Get("type")) != "image_generation" {
+ return true
+ }
+ hasImageTool = true
+ imageModel = openAIJSONString(item.Get("model"))
+ imageSize = openAIJSONString(item.Get("size"))
+ return false
+ })
+ }
+ if imageSize == "" {
+ imageSize = openAIJSONString(gjson.GetBytes(body, "size"))
+ }
+ if imageModel == "" {
+ bodyModel := openAIJSONString(gjson.GetBytes(body, "model"))
+ if isOpenAIImageBillingModelAlias(bodyModel) || !hasImageTool {
+ imageModel = bodyModel
+ }
+ }
+ }
+ if imageModel == "" && hasImageTool {
+ imageModel = "gpt-image-2"
+ }
+ if imageModel == "" {
+ imageModel = strings.TrimSpace(fallbackModel)
+ }
+ return OpenAIResponsesImageBillingConfig{
+ Model: imageModel,
+ SizeTier: normalizeOpenAIImageSizeTier(imageSize),
+ InputSize: imageSize,
+ }, nil
}
func isOpenAIImageBillingModelAlias(model string) bool {
@@ -236,3 +289,10 @@ func isOpenAIImageBillingModelAlias(model string) bool {
}
return isOpenAIImageGenerationModel(normalized) || strings.Contains(normalized, "image")
}
+
+func openAIJSONString(value gjson.Result) string {
+ if value.Type != gjson.String {
+ return ""
+ }
+ return strings.TrimSpace(value.String())
+}
diff --git a/backend/internal/service/image_generation_intent_test.go b/backend/internal/service/image_generation_intent_test.go
index 4621e9d9..59aab39c 100644
--- a/backend/internal/service/image_generation_intent_test.go
+++ b/backend/internal/service/image_generation_intent_test.go
@@ -84,6 +84,17 @@ func TestResolveOpenAIResponsesImageBillingConfigToolModelWins(t *testing.T) {
require.Equal(t, "2K", imageSize)
}
+func TestResolveOpenAIResponsesImageBillingConfigFromBodyIgnoresUnrelatedLargeInput(t *testing.T) {
+ cfg, err := resolveOpenAIResponsesImageBillingConfigDetailedFromBody(
+ []byte(`{"model":"mapped-text-model","tools":[{"type":"image_generation","model":"gpt-image-2","size":"2048x1152"}],"input":[{"type":"message","content":[{"type":"input_text","text":"hi","nonce":1e1000000}]}]}`),
+ "requested-model",
+ )
+ require.NoError(t, err)
+ require.Equal(t, "gpt-image-2", cfg.Model)
+ require.Equal(t, "2K", cfg.SizeTier)
+ require.Equal(t, "2048x1152", cfg.InputSize)
+}
+
func TestResolveOpenAIResponsesImageBillingConfigSupportsOfficialAndCustomSizes(t *testing.T) {
tests := []struct {
name string
diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go
index 87ffb2cb..d17b85d0 100644
--- a/backend/internal/service/openai_gateway_service.go
+++ b/backend/internal/service/openai_gateway_service.go
@@ -2397,172 +2397,83 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
return s.forwardOpenAIPassthrough(ctx, c, account, originalBody, reqModel, reasoningEffort, reqStream, startTime)
}
- reqBody, err := requestView.Decode(c)
- if err != nil {
- return nil, err
+ bodyModified := false
+ var reqBody map[string]any
+ ensureReqBody := func() (map[string]any, error) {
+ if requestView.HasPatches() {
+ patchedBody, patchErr := requestView.ApplyPatches()
+ if patchErr != nil {
+ return nil, patchErr
+ }
+ body = patchedBody
+ requestView = newOpenAIRequestView(body)
+ reqBody = nil
+ bodyModified = false
+ }
+ if reqBody != nil {
+ return reqBody, nil
+ }
+ decoded, decodeErr := requestView.Decode(c)
+ if decodeErr != nil {
+ return nil, decodeErr
+ }
+ reqBody = decoded
+ return reqBody, nil
+ }
+ markPatchSet := func(path string, value any) {
+ bodyModified = true
+ if requestView.patchesDisabled {
+ if reqBody != nil {
+ setOpenAIRequestMapPath(reqBody, path, value)
+ }
+ return
+ }
+ requestView.MarkPatchSet(path, value)
+ }
+ markPatchDelete := func(path string) {
+ bodyModified = true
+ if requestView.patchesDisabled {
+ if reqBody != nil {
+ deleteOpenAIRequestMapPath(reqBody, path)
+ }
+ return
+ }
+ requestView.MarkPatchDelete(path)
+ }
+ disablePatch := func() {
+ requestView.DisablePatches()
+ }
+ markDecodedModified := func() {
+ bodyModified = true
+ disablePatch()
}
- if v, ok := reqBody["model"].(string); ok {
- reqModel = v
- originalModel = reqModel
- }
- if v, ok := reqBody["stream"].(bool); ok {
- reqStream = v
- }
- if promptCacheKey == "" {
- if v, ok := reqBody["prompt_cache_key"].(string); ok {
- promptCacheKey = strings.TrimSpace(v)
- }
- }
apiKey := getAPIKeyFromContext(c)
imageGenerationAllowed := GroupAllowsImageGeneration(nil)
if apiKey != nil {
imageGenerationAllowed = GroupAllowsImageGeneration(apiKey.Group)
}
codexImageGenerationBridgeEnabled := isCodexCLI && imageGenerationAllowed && s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey)
- if IsImageGenerationIntentMap(openAIResponsesEndpoint, reqModel, reqBody) && !imageGenerationAllowed {
+ imageIntent := IsImageGenerationIntent(openAIResponsesEndpoint, reqModel, body)
+ if imageIntent && !imageGenerationAllowed {
MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalFeatureGate)
- c.JSON(http.StatusForbidden, gin.H{
- "error": gin.H{
- "type": "permission_error",
- "message": ImageGenerationPermissionMessage(),
- },
- })
+ c.JSON(http.StatusForbidden, gin.H{"error": gin.H{"type": "permission_error", "message": ImageGenerationPermissionMessage()}})
return nil, errors.New("image generation disabled for group")
}
- // Track if body needs re-serialization
- bodyModified := false
- // 单字段补丁快速路径:只要整个变更集最终可归约为同一路径的 set/delete,就避免全量 Marshal。
- patchDisabled := false
- patchHasOp := false
- patchDelete := false
- patchPath := ""
- var patchValue any
- markPatchSet := func(path string, value any) {
- if strings.TrimSpace(path) == "" {
- patchDisabled = true
- return
- }
- if patchDisabled {
- return
- }
- if !patchHasOp {
- patchHasOp = true
- patchDelete = false
- patchPath = path
- patchValue = value
- return
- }
- if patchDelete || patchPath != path {
- patchDisabled = true
- return
- }
- patchValue = value
- }
- markPatchDelete := func(path string) {
- if strings.TrimSpace(path) == "" {
- patchDisabled = true
- return
- }
- if patchDisabled {
- return
- }
- if !patchHasOp {
- patchHasOp = true
- patchDelete = true
- patchPath = path
- return
- }
- if !patchDelete || patchPath != path {
- patchDisabled = true
- }
- }
- disablePatch := func() {
- patchDisabled = true
- }
-
- // 非透传模式下,instructions 为空时注入默认指令。
- if isInstructionsEmpty(reqBody) && !compatMessagesBridge {
- reqBody["instructions"] = "You are a helpful coding assistant."
- bodyModified = true
+ instructions := gjson.GetBytes(body, "instructions")
+ instructionsEmpty := !instructions.Exists() || instructions.Type != gjson.String || strings.TrimSpace(instructions.String()) == ""
+ if instructionsEmpty && !compatMessagesBridge {
markPatchSet("instructions", "You are a helpful coding assistant.")
}
- if codexImageGenerationBridgeEnabled && ensureOpenAIResponsesImageGenerationTool(reqBody) {
- bodyModified = true
- disablePatch()
- logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Injected /responses image_generation tool for Codex client")
- }
-
- if normalizeOpenAIResponsesImageGenerationTools(reqBody) {
- bodyModified = true
- disablePatch()
- logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Normalized /responses image_generation tool payload")
- }
- if codexImageGenerationBridgeEnabled && applyCodexImageGenerationBridgeInstructions(reqBody) {
- bodyModified = true
- disablePatch()
- logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Added Codex image_generation bridge instructions")
- }
-
- // 对所有请求执行模型映射(包含 Codex CLI)。
billingModel := account.GetMappedModel(reqModel)
if billingModel != reqModel {
logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Model mapping applied: %s -> %s (account: %s, isCodexCLI: %v)", reqModel, billingModel, account.Name, isCodexCLI)
- reqBody["model"] = billingModel
- bodyModified = true
+ reqModel = billingModel
markPatchSet("model", billingModel)
}
upstreamModel := billingModel
- if imageGenerationAllowed && normalizeOpenAIResponsesImageOnlyModel(reqBody) {
- bodyModified = true
- disablePatch()
- if model, ok := reqBody["model"].(string); ok {
- upstreamModel = strings.TrimSpace(model)
- }
- logger.LegacyPrintf(
- "service.openai_gateway",
- "[OpenAI] Normalized /responses image-only model request inbound_model=%s image_model=%s upstream_model=%s",
- reqModel,
- billingModel,
- upstreamModel,
- )
- }
- if err := validateOpenAIResponsesImageModel(reqBody, upstreamModel); err != nil {
- setOpsUpstreamError(c, http.StatusBadRequest, err.Error(), "")
- c.JSON(http.StatusBadRequest, gin.H{
- "error": gin.H{
- "type": "invalid_request_error",
- "message": err.Error(),
- "param": "model",
- },
- })
- return nil, err
- }
- if hasOpenAIImageGenerationTool(reqBody) {
- logger.LegacyPrintf(
- "service.openai_gateway",
- "[OpenAI] /responses image_generation request inbound_model=%s mapped_model=%s account_type=%s",
- reqModel,
- upstreamModel,
- account.Type,
- )
- }
- if err := validateCodexSparkInput(reqBody, upstreamModel); err != nil {
- setOpsUpstreamError(c, http.StatusBadRequest, err.Error(), "")
- c.JSON(http.StatusBadRequest, gin.H{
- "error": gin.H{
- "type": "invalid_request_error",
- "message": err.Error(),
- "param": "input",
- },
- })
- return nil, err
- }
-
- // Compact-only model 映射:仅在 /responses/compact 路径生效,且优先级高于
- // OAuth 模型规范化(避免 OAuth 规范化覆盖 compact-only 自定义模型)。
isCompactRequest := isOpenAIResponsesCompactPath(c)
compactMapped := false
if isCompactRequest {
@@ -2570,69 +2481,99 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
if compactMappedModel != "" && compactMappedModel != billingModel {
compactMapped = true
upstreamModel = compactMappedModel
- reqBody["model"] = compactMappedModel
- bodyModified = true
+ reqModel = compactMappedModel
markPatchSet("model", compactMappedModel)
logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Compact model mapping applied: %s -> %s (account: %s, isCodexCLI: %v)", billingModel, compactMappedModel, account.Name, isCodexCLI)
}
}
-
- // OpenAI OAuth 账号走 ChatGPT internal Codex endpoint,需要将模型名规范化为
- // 上游可识别的 Codex/GPT 系列。API Key 账号则应保留原始/映射后的模型名,
- // 以兼容自定义 base_url 的 OpenAI-compatible 上游。
- if model, ok := reqBody["model"].(string); ok {
- if !compactMapped {
- upstreamModel = normalizeOpenAIModelForUpstream(account, model)
- if upstreamModel != "" && upstreamModel != model {
- logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Upstream model resolved: %s -> %s (account: %s, type: %s, isCodexCLI: %v)",
- model, upstreamModel, account.Name, account.Type, isCodexCLI)
- reqBody["model"] = upstreamModel
- bodyModified = true
- markPatchSet("model", upstreamModel)
- }
+ if !compactMapped {
+ modelForNormalize := reqModel
+ if modelForNormalize == "" {
+ modelForNormalize = requestView.Model
}
-
- // 移除 gpt-5.2-codex 以下的版本 verbosity 参数
- // 确保高版本模型向低版本模型映射不报错
- if !SupportsVerbosity(upstreamModel) {
- if text, ok := reqBody["text"].(map[string]any); ok {
- if _, exists := text["verbosity"]; exists {
- delete(text, "verbosity")
- bodyModified = true
- markPatchDelete("text.verbosity")
- }
- }
+ upstreamModel = normalizeOpenAIModelForUpstream(account, modelForNormalize)
+ if upstreamModel != "" && upstreamModel != modelForNormalize {
+ logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Upstream model resolved: %s -> %s (account: %s, type: %s, isCodexCLI: %v)", modelForNormalize, upstreamModel, account.Name, account.Type, isCodexCLI)
+ reqModel = upstreamModel
+ markPatchSet("model", upstreamModel)
}
}
+ if strings.TrimSpace(gjson.GetBytes(body, "reasoning.effort").String()) == "minimal" {
+ markPatchSet("reasoning.effort", "none")
+ logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Normalized reasoning.effort: minimal -> none (account: %s)", account.Name)
+ }
- // 规范化 reasoning.effort 参数(minimal -> none),与上游允许值对齐。
- if reasoning, ok := reqBody["reasoning"].(map[string]any); ok {
- if effort, ok := reasoning["effort"].(string); ok && effort == "minimal" {
- reasoning["effort"] = "none"
- bodyModified = true
- markPatchSet("reasoning.effort", "none")
- logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Normalized reasoning.effort: minimal -> none (account: %s)", account.Name)
+ imageIntent = imageIntent || IsImageGenerationIntent(openAIResponsesEndpoint, reqModel, nil) || isOpenAIImageGenerationModel(upstreamModel)
+ if imageIntent && !imageGenerationAllowed {
+ MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalFeatureGate)
+ c.JSON(http.StatusForbidden, gin.H{"error": gin.H{"type": "permission_error", "message": ImageGenerationPermissionMessage()}})
+ return nil, errors.New("image generation disabled for group")
+ }
+
+ if imageGenerationAllowed && (codexImageGenerationBridgeEnabled || isOpenAIImageGenerationModel(requestView.Model) || openAIRequestBodyImageGenerationToolNeedsNormalization(body) || isOpenAIImageGenerationModel(upstreamModel)) {
+ decoded, decodeErr := ensureReqBody()
+ if decodeErr != nil {
+ return nil, decodeErr
+ }
+ if codexImageGenerationBridgeEnabled && ensureOpenAIResponsesImageGenerationTool(decoded) {
+ markDecodedModified()
+ logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Injected /responses image_generation tool for Codex client")
+ }
+ if normalizeOpenAIResponsesImageGenerationTools(decoded) {
+ markDecodedModified()
+ logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Normalized /responses image_generation tool payload")
+ }
+ if normalizeOpenAIResponsesImageOnlyModel(decoded) {
+ markDecodedModified()
+ if model, ok := decoded["model"].(string); ok {
+ upstreamModel = strings.TrimSpace(model)
+ }
+ logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Normalized /responses image-only model request inbound_model=%s image_model=%s upstream_model=%s", requestView.Model, billingModel, upstreamModel)
+ }
+ if err := validateOpenAIResponsesImageModel(decoded, upstreamModel); err != nil {
+ setOpsUpstreamError(c, http.StatusBadRequest, err.Error(), "")
+ c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{"type": "invalid_request_error", "message": err.Error(), "param": "model"}})
+ return nil, err
+ }
+ if hasOpenAIImageGenerationTool(decoded) {
+ imageIntent = true
+ logger.LegacyPrintf("service.openai_gateway", "[OpenAI] /responses image_generation request inbound_model=%s mapped_model=%s account_type=%s", requestView.Model, upstreamModel, account.Type)
+ }
+ if codexImageGenerationBridgeEnabled && applyCodexImageGenerationBridgeInstructions(decoded) {
+ markDecodedModified()
+ logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Added Codex image_generation bridge instructions")
+ }
+ } else if imageGenerationAllowed && imageIntent && openAIRequestBodyHasImageGenerationTool(body) {
+ logger.LegacyPrintf("service.openai_gateway", "[OpenAI] /responses image_generation request inbound_model=%s mapped_model=%s account_type=%s", requestView.Model, upstreamModel, account.Type)
+ }
+
+ if isCodexSparkModel(upstreamModel) && openAIRequestBodyMayContainImageInput(body) {
+ decoded, decodeErr := ensureReqBody()
+ if decodeErr != nil {
+ return nil, decodeErr
+ }
+ if err := validateCodexSparkInput(decoded, upstreamModel); err != nil {
+ setOpsUpstreamError(c, http.StatusBadRequest, err.Error(), "")
+ c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{"type": "invalid_request_error", "message": err.Error(), "param": "input"}})
+ return nil, err
}
}
if account.Type == AccountTypeOAuth {
+ decoded, decodeErr := ensureReqBody()
+ if decodeErr != nil {
+ return nil, decodeErr
+ }
codexResult := codexTransformResult{}
if compatMessagesBridge {
- codexResult = applyCodexOAuthTransformWithOptions(reqBody, codexOAuthTransformOptions{
- IsCodexCLI: isCodexCLI,
- IsCompact: isCompactRequest,
- SkipDefaultInstructions: true,
- PreserveToolCallIDs: true,
- })
- ensureCodexOAuthInstructionsField(reqBody)
- bodyModified = true
- disablePatch()
+ codexResult = applyCodexOAuthTransformWithOptions(decoded, codexOAuthTransformOptions{IsCodexCLI: isCodexCLI, IsCompact: isCompactRequest, SkipDefaultInstructions: true, PreserveToolCallIDs: true})
+ ensureCodexOAuthInstructionsField(decoded)
+ markDecodedModified()
} else {
- codexResult = applyCodexOAuthTransform(reqBody, isCodexCLI, isCompactRequest)
+ codexResult = applyCodexOAuthTransform(decoded, isCodexCLI, isCompactRequest)
}
if codexResult.Modified {
- bodyModified = true
- disablePatch()
+ markDecodedModified()
}
if codexResult.NormalizedModel != "" {
upstreamModel = codexResult.NormalizedModel
@@ -2642,90 +2583,57 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
}
}
- // Handle max_output_tokens based on platform and account type
+ if !SupportsVerbosity(upstreamModel) && gjson.GetBytes(body, "text.verbosity").Exists() {
+ markPatchDelete("text.verbosity")
+ }
+
if !isCodexCLI {
- if maxOutputTokens, hasMaxOutputTokens := reqBody["max_output_tokens"]; hasMaxOutputTokens {
+ maxOutputTokens := gjson.GetBytes(body, "max_output_tokens")
+ if maxOutputTokens.Exists() {
switch account.Platform {
case PlatformOpenAI:
- // For OpenAI API Key, remove max_output_tokens (not supported)
- // For OpenAI OAuth (Responses API), keep it (supported)
if account.Type == AccountTypeAPIKey {
- delete(reqBody, "max_output_tokens")
- bodyModified = true
markPatchDelete("max_output_tokens")
}
case PlatformAnthropic:
- // For Anthropic (Claude), convert to max_tokens
- delete(reqBody, "max_output_tokens")
- markPatchDelete("max_output_tokens")
- if _, hasMaxTokens := reqBody["max_tokens"]; !hasMaxTokens {
- reqBody["max_tokens"] = maxOutputTokens
- disablePatch()
+ decoded, decodeErr := ensureReqBody()
+ if decodeErr != nil {
+ return nil, decodeErr
}
- bodyModified = true
+ delete(decoded, "max_output_tokens")
+ if _, hasMaxTokens := decoded["max_tokens"]; !hasMaxTokens {
+ decoded["max_tokens"] = maxOutputTokens.Value()
+ }
+ markDecodedModified()
case PlatformGemini:
- // For Gemini, remove (will be handled by Gemini-specific transform)
- delete(reqBody, "max_output_tokens")
- bodyModified = true
markPatchDelete("max_output_tokens")
default:
- // For unknown platforms, remove to be safe
- delete(reqBody, "max_output_tokens")
- bodyModified = true
markPatchDelete("max_output_tokens")
}
}
-
- // Also handle max_completion_tokens (similar logic)
- if _, hasMaxCompletionTokens := reqBody["max_completion_tokens"]; hasMaxCompletionTokens {
- if account.Type == AccountTypeAPIKey || account.Platform != PlatformOpenAI {
- delete(reqBody, "max_completion_tokens")
- bodyModified = true
- markPatchDelete("max_completion_tokens")
- }
+ if gjson.GetBytes(body, "max_completion_tokens").Exists() && (account.Type == AccountTypeAPIKey || account.Platform != PlatformOpenAI) {
+ markPatchDelete("max_completion_tokens")
}
-
- // Remove unsupported fields (not supported by upstream OpenAI API)
- unsupportedFields := []string{"prompt_cache_retention", "safety_identifier"}
- for _, unsupportedField := range unsupportedFields {
- if _, has := reqBody[unsupportedField]; has {
- delete(reqBody, unsupportedField)
- bodyModified = true
+ for _, unsupportedField := range []string{"prompt_cache_retention", "safety_identifier"} {
+ if gjson.GetBytes(body, unsupportedField).Exists() {
markPatchDelete(unsupportedField)
}
}
}
-
- // 仅在 WSv2 模式保留 previous_response_id,其他模式(HTTP/WSv1)统一过滤。
- // 注意:该规则同样适用于 Codex CLI 请求,避免 WSv1 向上游透传不支持字段。
- if wsDecision.Transport != OpenAIUpstreamTransportResponsesWebsocketV2 {
- if _, has := reqBody["previous_response_id"]; has {
- delete(reqBody, "previous_response_id")
- bodyModified = true
- markPatchDelete("previous_response_id")
+ if wsDecision.Transport != OpenAIUpstreamTransportResponsesWebsocketV2 && gjson.GetBytes(body, "previous_response_id").Exists() {
+ markPatchDelete("previous_response_id")
+ }
+ if openAIRequestBodyMayContainEmptyBase64InputImage(body) {
+ decoded, decodeErr := ensureReqBody()
+ if decodeErr != nil {
+ return nil, decodeErr
+ }
+ if sanitizeEmptyBase64InputImagesInOpenAIRequestBodyMap(decoded) {
+ markDecodedModified()
}
}
- if sanitizeEmptyBase64InputImagesInOpenAIRequestBodyMap(reqBody) {
- bodyModified = true
- disablePatch()
- }
-
- // Apply OpenAI fast policy (参照 Claude BetaPolicy 的 fast-mode 过滤):
- // 针对 body 的 service_tier 字段("priority" 即 fast,"flex"),按策略
- // 执行 filter(删除字段)或 block(拒绝请求)。对 gpt-5.5 等模型屏蔽
- // fast 时在此生效。
- //
- // 注意:
- // 1. 此处统一使用 upstreamModel(已经过 GetMappedModel +
- // normalizeOpenAIModelForUpstream + Codex OAuth normalize),与
- // chat-completions / messages 入口保持一致,避免不同入口因为模型
- // 维度不同而出现 whitelist 命中差异。
- // 2. action=pass 时也要把 raw "fast" 归一化为 "priority" 写回 body,
- // 否则 native /responses 入口透传 "fast" 给上游会被拒。chat-
- // completions 入口由 normalizeResponsesBodyServiceTier 完成同一
- // 行为,这里手工实现等效逻辑。
- if rawTier, ok := reqBody["service_tier"].(string); ok {
+ if rawTier := requestView.ServiceTier; rawTier != "" {
if normTier := normalizedOpenAIServiceTierValue(rawTier); normTier != "" {
action, errMsg := s.evaluateOpenAIFastPolicy(ctx, account, upstreamModel, normTier)
switch action {
@@ -2738,46 +2646,51 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
writeOpenAIFastPolicyBlockedResponse(c, blocked)
return nil, blocked
case BetaPolicyActionFilter:
- delete(reqBody, "service_tier")
- bodyModified = true
- disablePatch()
+ markPatchDelete("service_tier")
default:
- // pass:若客户端传的是别名 "fast",归一化为 "priority"
- // 后写回 body,确保上游收到的是其能识别的规范值。
if normTier != rawTier {
- reqBody["service_tier"] = normTier
- bodyModified = true
markPatchSet("service_tier", normTier)
}
}
}
}
- if IsImageGenerationIntentMap(openAIResponsesEndpoint, reqModel, reqBody) && !imageGenerationAllowed {
- MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalFeatureGate)
- c.JSON(http.StatusForbidden, gin.H{
- "error": gin.H{
- "type": "permission_error",
- "message": ImageGenerationPermissionMessage(),
- },
- })
- return nil, errors.New("image generation disabled for group")
+ if bodyModified {
+ if requestView.HasPatches() {
+ if patchedBody, patchErr := requestView.ApplyPatches(); patchErr == nil {
+ body = patchedBody
+ requestView = newOpenAIRequestView(body)
+ reqBody = nil
+ bodyModified = false
+ }
+ }
+ if bodyModified {
+ decoded, decodeErr := ensureReqBody()
+ if decodeErr != nil {
+ return nil, decodeErr
+ }
+ var marshalErr error
+ body, marshalErr = marshalOpenAIUpstreamJSON(decoded)
+ if marshalErr != nil {
+ return nil, fmt.Errorf("serialize request body: %w", marshalErr)
+ }
+ requestView = newOpenAIRequestView(body)
+ }
}
imageBillingModel := ""
imageSizeTier := ""
imageInputSize := ""
- if IsImageGenerationIntentMap(openAIResponsesEndpoint, reqModel, reqBody) {
+ if imageIntent {
+ var imageCfg OpenAIResponsesImageBillingConfig
var imageCfgErr error
- imageCfg, imageCfgErr := resolveOpenAIResponsesImageBillingConfigDetailed(reqBody, billingModel)
+ if reqBody != nil {
+ imageCfg, imageCfgErr = resolveOpenAIResponsesImageBillingConfigDetailed(reqBody, billingModel)
+ } else {
+ imageCfg, imageCfgErr = resolveOpenAIResponsesImageBillingConfigDetailedFromBody(body, billingModel)
+ }
if imageCfgErr != nil {
setOpsUpstreamError(c, http.StatusBadRequest, imageCfgErr.Error(), "")
- c.JSON(http.StatusBadRequest, gin.H{
- "error": gin.H{
- "type": "invalid_request_error",
- "message": imageCfgErr.Error(),
- "param": "size",
- },
- })
+ c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{"type": "invalid_request_error", "message": imageCfgErr.Error(), "param": "size"}})
return nil, imageCfgErr
}
imageBillingModel = imageCfg.Model
@@ -2785,29 +2698,6 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
imageInputSize = imageCfg.InputSize
}
- // Re-serialize body only if modified
- if bodyModified {
- serializedByPatch := false
- if !patchDisabled && patchHasOp {
- var patchErr error
- if patchDelete {
- body, patchErr = sjson.DeleteBytes(body, patchPath)
- } else {
- body, patchErr = sjson.SetBytes(body, patchPath, patchValue)
- }
- if patchErr == nil {
- serializedByPatch = true
- }
- }
- if !serializedByPatch {
- var marshalErr error
- body, marshalErr = marshalOpenAIUpstreamJSON(reqBody)
- if marshalErr != nil {
- return nil, fmt.Errorf("serialize request body: %w", marshalErr)
- }
- }
- }
-
// Get access token
token, _, err := s.GetAccessToken(ctx, account)
if err != nil {
@@ -2816,8 +2706,11 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
// 命中 WS 时仅走 WebSocket Mode;不再自动回退 HTTP。
if wsDecision.Transport == OpenAIUpstreamTransportResponsesWebsocketV2 {
- // WS 分支不会再回落 HTTP;重连恢复可直接更新 reqBody,避免额外保留一份完整顶层 map。
- wsReqBody := reqBody
+ // WS 分支需要结构化 payload 与重连恢复,命中后再触发 full-map decode。
+ wsReqBody, err := ensureReqBody()
+ if err != nil {
+ return nil, err
+ }
_, hasPreviousResponseID := wsReqBody["previous_response_id"]
logOpenAIWSModeDebug(
"forward_start account_id=%d account_type=%s model=%s stream=%v has_previous_response_id=%v",
@@ -3076,8 +2969,12 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg)
upstreamCode := extractUpstreamErrorCode(respBody)
if !httpInvalidEncryptedContentRetryTried && resp.StatusCode == http.StatusBadRequest && upstreamCode == "invalid_encrypted_content" {
- if trimOpenAIEncryptedReasoningItems(reqBody) {
- body, err = marshalOpenAIUpstreamJSON(reqBody)
+ decoded, decodeErr := ensureReqBody()
+ if decodeErr != nil {
+ return nil, decodeErr
+ }
+ if trimOpenAIEncryptedReasoningItems(decoded) {
+ body, err = marshalOpenAIUpstreamJSON(decoded)
if err != nil {
return nil, fmt.Errorf("serialize invalid_encrypted_content retry body: %w", err)
}
@@ -3118,8 +3015,8 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
}
defer func() { _ = resp.Body.Close() }()
- reasoningEffort := extractOpenAIReasoningEffort(reqBody, originalModel)
- serviceTier := extractOpenAIServiceTier(reqBody)
+ reasoningEffort := extractOpenAIReasoningEffortFromBody(body, originalModel)
+ serviceTier := extractOpenAIServiceTierFromBody(body)
// 上游接受后只保留计费需要的标量,避免响应处理期间继续保活完整 input/tools map。
reqBody = nil
@@ -6283,6 +6180,14 @@ type openAIRequestView struct {
PreviousResponseID string
ServiceTier string
ReasoningEffort string
+ patches []openAIRequestPatch
+ patchesDisabled bool
+}
+
+type openAIRequestPatch struct {
+ path string
+ delete bool
+ value any
}
func newOpenAIRequestView(body []byte) openAIRequestView {
@@ -6305,6 +6210,110 @@ func (v openAIRequestView) Decode(c *gin.Context) (map[string]any, error) {
return getOpenAIRequestBodyMap(c, v.body)
}
+func (v *openAIRequestView) MarkPatchSet(path string, value any) {
+ if v == nil || v.patchesDisabled {
+ return
+ }
+ path = strings.TrimSpace(path)
+ if path == "" {
+ v.DisablePatches()
+ return
+ }
+ v.patches = append(v.patches, openAIRequestPatch{path: path, value: value})
+}
+
+func (v *openAIRequestView) MarkPatchDelete(path string) {
+ if v == nil || v.patchesDisabled {
+ return
+ }
+ path = strings.TrimSpace(path)
+ if path == "" {
+ v.DisablePatches()
+ return
+ }
+ v.patches = append(v.patches, openAIRequestPatch{path: path, delete: true})
+}
+
+func (v *openAIRequestView) DisablePatches() {
+ if v == nil {
+ return
+ }
+ v.patchesDisabled = true
+ v.patches = nil
+}
+
+func (v openAIRequestView) HasPatches() bool {
+ return !v.patchesDisabled && len(v.patches) > 0
+}
+
+func (v openAIRequestView) ApplyPatches() ([]byte, error) {
+ if v.patchesDisabled || len(v.patches) == 0 {
+ return nil, errors.New("openai request patches disabled")
+ }
+ body := v.body
+ for _, patch := range v.patches {
+ var err error
+ if patch.delete {
+ body, err = sjson.DeleteBytes(body, patch.path)
+ } else {
+ body, err = sjson.SetBytes(body, patch.path, patch.value)
+ }
+ if err != nil {
+ return nil, err
+ }
+ }
+ return body, nil
+}
+
+func setOpenAIRequestMapPath(reqBody map[string]any, path string, value any) {
+ path = strings.TrimSpace(path)
+ if reqBody == nil || path == "" {
+ return
+ }
+ parts := strings.Split(path, ".")
+ current := reqBody
+ for _, part := range parts[:len(parts)-1] {
+ part = strings.TrimSpace(part)
+ if part == "" {
+ return
+ }
+ next, _ := current[part].(map[string]any)
+ if next == nil {
+ next = map[string]any{}
+ current[part] = next
+ }
+ current = next
+ }
+ last := strings.TrimSpace(parts[len(parts)-1])
+ if last != "" {
+ current[last] = value
+ }
+}
+
+func deleteOpenAIRequestMapPath(reqBody map[string]any, path string) {
+ path = strings.TrimSpace(path)
+ if reqBody == nil || path == "" {
+ return
+ }
+ parts := strings.Split(path, ".")
+ current := reqBody
+ for _, part := range parts[:len(parts)-1] {
+ part = strings.TrimSpace(part)
+ if part == "" {
+ return
+ }
+ next, _ := current[part].(map[string]any)
+ if next == nil {
+ return
+ }
+ current = next
+ }
+ last := strings.TrimSpace(parts[len(parts)-1])
+ if last != "" {
+ delete(current, last)
+ }
+}
+
func extractOpenAIRequestMetaFromBody(body []byte) (model string, stream bool, promptCacheKey string) {
view := newOpenAIRequestView(body)
return view.Model, view.Stream, view.PromptCacheKey
@@ -6758,8 +6767,84 @@ func buildOpenAIFastPolicyBlockedWSEvent(err *OpenAIFastBlockedError) []byte {
return payload
}
+func openAIRequestBodyMayContainImageInput(body []byte) bool {
+ if len(body) == 0 {
+ return false
+ }
+ input := gjson.GetBytes(body, "input")
+ messages := gjson.GetBytes(body, "messages")
+ return openAIJSONValueMayContainImageInput(input) || openAIJSONValueMayContainImageInput(messages)
+}
+
+func openAIJSONValueMayContainImageInput(value gjson.Result) bool {
+ if !value.Exists() {
+ return false
+ }
+ if value.IsArray() {
+ found := false
+ value.ForEach(func(_, item gjson.Result) bool {
+ if openAIJSONValueMayContainImageInput(item) {
+ found = true
+ return false
+ }
+ return true
+ })
+ return found
+ }
+ if value.IsObject() {
+ if strings.TrimSpace(value.Get("type").String()) == "input_image" || value.Get("image_url").Exists() {
+ return true
+ }
+ return openAIJSONValueMayContainImageInput(value.Get("content"))
+ }
+ return false
+}
+
+func openAIRequestBodyMayContainEmptyBase64InputImage(body []byte) bool {
+ if len(body) == 0 || !openAIRequestBodyMayContainInputImageToken(body) {
+ return false
+ }
+ input := gjson.GetBytes(body, "input")
+ if !input.Exists() {
+ return false
+ }
+ return openAIJSONValueMayContainEmptyBase64InputImage(input)
+}
+
+func openAIRequestBodyMayContainInputImageToken(body []byte) bool {
+ if bytes.Contains(body, []byte("input_image")) {
+ return true
+ }
+ // JSON 字符串任意字符都可能被 unicode escape,遇到 \u 时交给 gjson 解码后的结构扫描兜底。
+ return bytes.Contains(body, []byte("\\u"))
+}
+
+func openAIJSONValueMayContainEmptyBase64InputImage(value gjson.Result) bool {
+ if !value.Exists() {
+ return false
+ }
+ if value.IsArray() {
+ found := false
+ value.ForEach(func(_, item gjson.Result) bool {
+ if openAIJSONValueMayContainEmptyBase64InputImage(item) {
+ found = true
+ return false
+ }
+ return true
+ })
+ return found
+ }
+ if value.IsObject() {
+ if strings.TrimSpace(value.Get("type").String()) == "input_image" && isEmptyBase64DataURI(value.Get("image_url").String()) {
+ return true
+ }
+ return openAIJSONValueMayContainEmptyBase64InputImage(value.Get("content"))
+ }
+ return false
+}
+
func sanitizeEmptyBase64InputImagesInOpenAIBody(body []byte) ([]byte, bool, error) {
- if len(body) == 0 || !bytes.Contains(body, []byte(`"image_url"`)) || !bytes.Contains(body, []byte(`base64,`)) {
+ if !openAIRequestBodyMayContainEmptyBase64InputImage(body) {
return body, false, nil
}
diff --git a/backend/internal/service/openai_gateway_service_hotpath_test.go b/backend/internal/service/openai_gateway_service_hotpath_test.go
index df17ed2e..d240cec8 100644
--- a/backend/internal/service/openai_gateway_service_hotpath_test.go
+++ b/backend/internal/service/openai_gateway_service_hotpath_test.go
@@ -1,12 +1,18 @@
package service
import (
+ "context"
"encoding/json"
+ "io"
+ "net/http"
"net/http/httptest"
+ "strings"
"testing"
+ "github.com/Wei-Shaw/sub2api/internal/config"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
+ "github.com/tidwall/gjson"
)
func TestOpenAIRequestView_ExtractsRawScalars(t *testing.T) {
@@ -29,6 +35,454 @@ func TestOpenAIRequestView_DecodeKeepsFullMapBehavior(t *testing.T) {
require.IsType(t, []any{}, reqBody["input"])
}
+func TestOpenAIRequestView_ApplyPatches(t *testing.T) {
+ view := newOpenAIRequestView([]byte(`{"model":"gpt-5","previous_response_id":"resp_1","reasoning":{"effort":"minimal"},"input":[{"type":"message","content":"hi"}]}`))
+ view.MarkPatchSet("model", "gpt-5.1")
+ view.MarkPatchDelete("previous_response_id")
+ view.MarkPatchSet("reasoning.effort", "none")
+
+ patched, err := view.ApplyPatches()
+ require.NoError(t, err)
+ require.JSONEq(t, `{"model":"gpt-5.1","reasoning":{"effort":"none"},"input":[{"type":"message","content":"hi"}]}`, string(patched))
+}
+
+func TestOpenAIRequestView_ApplyPatchesDisabled(t *testing.T) {
+ view := newOpenAIRequestView([]byte(`{"model":"gpt-5"}`))
+ view.MarkPatchSet("model", "gpt-5.1")
+ view.DisablePatches()
+
+ _, err := view.ApplyPatches()
+ require.Error(t, err)
+}
+
+func TestOpenAIRequestView_HasPatches(t *testing.T) {
+ view := newOpenAIRequestView([]byte(`{"model":"gpt-5"}`))
+ require.False(t, view.HasPatches())
+
+ view.MarkPatchSet("model", "gpt-5.1")
+ require.True(t, view.HasPatches())
+
+ view.DisablePatches()
+ require.False(t, view.HasPatches())
+}
+
+func TestOpenAIGatewayService_Forward_HTTPPatchPathKeepsLargeInputRaw(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(
+ `{"usage":{"input_tokens":1,"output_tokens":2,"input_tokens_details":{"cached_tokens":0}}}`,
+ )),
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 1,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"gpt-5","stream":false,"reasoning":{"effort":"minimal"},"input":[{"type":"message","content":[{"type":"input_text","text":"hi","nonce":9007199254740993}]}]}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.NotNil(t, upstream.lastReq)
+ require.JSONEq(t, `{"model":"gpt-5","stream":false,"reasoning":{"effort":"none"},"instructions":"You are a helpful coding assistant.","input":[{"type":"message","content":[{"type":"input_text","text":"hi","nonce":9007199254740993}]}]}`, string(upstream.lastBody))
+ require.Equal(t, "9007199254740993", gjson.GetBytes(upstream.lastBody, "input.0.content.0.nonce").Raw)
+}
+
+func TestOpenAIGatewayService_Forward_DecodedMutationKeepsLaterFieldDeletes(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)),
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 2,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"gpt-5.4","stream":false,"max_completion_tokens":12,"tools":[{"type":"image_generation","format":"png"}],"input":[{"type":"message","content":"draw"}]}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.False(t, gjson.GetBytes(upstream.lastBody, "max_completion_tokens").Exists())
+ require.False(t, gjson.GetBytes(upstream.lastBody, "tools.0.format").Exists())
+ require.Equal(t, "png", gjson.GetBytes(upstream.lastBody, "tools.0.output_format").String())
+}
+
+func TestOpenAIGatewayService_Forward_MappedImageModelUsesImageGate(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)),
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 3,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ "model_mapping": map[string]any{"draw-alias": "gpt-image-2"},
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ c.Set("api_key", &APIKey{Group: &Group{AllowImageGeneration: false}})
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"draw-alias","stream":false,"input":"draw"}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.Error(t, err)
+ require.Nil(t, result)
+ require.Nil(t, upstream.lastReq)
+ require.Equal(t, http.StatusForbidden, rec.Code)
+}
+
+func TestOpenAIGatewayService_Forward_TextDataImageDoesNotForceMapMarshal(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)),
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 4,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"gpt-5","stream":false,"input":[{"type":"message","content":[{"type":"input_text","text":"literal data:image/png;base64, only","nonce":1e1000000}]}]}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.Equal(t, "1e1000000", gjson.GetBytes(upstream.lastBody, "input.0.content.0.nonce").Raw)
+}
+
+func TestOpenAIGatewayService_Forward_ImageToolBillingDoesNotForceFullDecode(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(
+ `{"output":[{"id":"ig_1","type":"image_generation_call","result":"final-image"}],"usage":{"input_tokens":1,"output_tokens":2}}`,
+ )),
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 9,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"gpt-5","stream":false,"tools":[{"type":"image_generation","model":"gpt-image-2","size":"2048x1152"}],"input":[{"type":"message","content":[{"type":"input_text","text":"draw","nonce":1e1000000}]}]}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.Equal(t, "1e1000000", gjson.GetBytes(upstream.lastBody, "input.0.content.0.nonce").Raw)
+ require.Equal(t, 1, result.ImageCount)
+ require.Equal(t, "2K", result.ImageSize)
+ require.Equal(t, "gpt-image-2", result.BillingModel)
+}
+
+func TestOpenAIGatewayService_Forward_HTTPRetryRecoveryDoesNotDecodeBeforeError(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ responses: []*http.Response{
+ {
+ StatusCode: http.StatusBadRequest,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(`{"error":{"code":"invalid_encrypted_content","type":"invalid_request_error","message":"bad encrypted content"}}`)),
+ },
+ {
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)),
+ },
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 10,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"gpt-5","stream":false,"input":[{"type":"reasoning","encrypted_content":"gAAA","summary":[{"type":"summary_text","text":"keep me"}]},{"type":"message","content":[{"type":"input_text","text":"hi","nonce":9007199254740993}]}]}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.Len(t, upstream.bodies, 2)
+ require.Equal(t, "gAAA", gjson.GetBytes(upstream.bodies[0], "input.0.encrypted_content").String())
+ require.Equal(t, "9007199254740993", gjson.GetBytes(upstream.bodies[0], "input.1.content.0.nonce").Raw)
+ require.False(t, gjson.GetBytes(upstream.bodies[1], "input.0.encrypted_content").Exists())
+ require.Equal(t, "summary_text", gjson.GetBytes(upstream.bodies[1], "input.0.summary.0.type").String())
+}
+
+func TestOpenAIGatewayService_Forward_CodexSparkRejectsEscapedInputImage(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)),
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 5,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"gpt-5.3-codex-spark","stream":false,"input":[{"type":"input_` + "\\u0069" + `mage","file_id":"file_1"}]}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.Error(t, err)
+ require.Nil(t, result)
+ require.Nil(t, upstream.lastReq)
+ require.Equal(t, http.StatusBadRequest, rec.Code)
+}
+
+func TestOpenAIGatewayService_Forward_CodexBridgeInjectionSetsImageBilling(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(
+ `{"output":[{"id":"ig_1","type":"image_generation_call","result":"final-image","size":"1024x1024"}],"usage":{"input_tokens":1,"output_tokens":2}}`,
+ )),
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ cfg.Gateway.ForceCodexCLI = true
+ cfg.Gateway.CodexImageGenerationBridgeEnabled = true
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 7,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ c.Set("api_key", &APIKey{Group: &Group{AllowImageGeneration: true}})
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"gpt-5","stream":false,"input":"draw if needed"}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.Equal(t, 1, result.ImageCount)
+ require.Equal(t, "2K", result.ImageSize)
+ require.Equal(t, "gpt-image-2", result.BillingModel)
+}
+
+func TestOpenAIGatewayService_Forward_HTTPDeletesPreviousResponseIDWhenPresent(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ account := &Account{
+ ID: 8,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+
+ for _, body := range [][]byte{
+ []byte(`{"model":"gpt-5","stream":false,"previous_response_id":"","input":"hi"}`),
+ []byte(`{"model":"gpt-5","stream":false,"previous_response_id":null,"input":"hi"}`),
+ } {
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)),
+ },
+ }
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.False(t, gjson.GetBytes(upstream.lastBody, "previous_response_id").Exists())
+ }
+}
+
+func TestOpenAIRequestBodyMayContainEmptyBase64InputImageSeesEscapedJSON(t *testing.T) {
+ body := []byte(`{"input":[{"type":"message","content":[{"type":"input_image","image_` + "\\u0075" + `rl":"data:image/png;base64` + "\\u002c" + ` "}]}]}`)
+
+ require.True(t, openAIRequestBodyMayContainEmptyBase64InputImage(body))
+}
+
+func TestOpenAIRequestBodyMayContainEmptyBase64InputImageSeesEscapedImageType(t *testing.T) {
+ body := []byte(`{"input":[{"type":"message","content":[{"type":"input_` + "\\u0069" + `mage","image_url":"data:image/png;base64, "}]}]}`)
+
+ require.True(t, openAIRequestBodyMayContainEmptyBase64InputImage(body))
+}
+
+func TestOpenAIRequestBodyMayContainEmptyBase64InputImageSeesEscapedInputPrefix(t *testing.T) {
+ body := []byte(`{"input":[{"type":"message","content":[{"type":"inp` + "\\u0075" + `t_image","image_url":"data:image/png;base64, "}]}]}`)
+
+ require.True(t, openAIRequestBodyMayContainEmptyBase64InputImage(body))
+}
+
+func TestOpenAIGatewayService_Forward_ImageOnlyModelKeepsSupportedVerbosity(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)),
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 6,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"gpt-image-2","stream":false,"text":{"verbosity":"low"},"input":"draw"}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.Equal(t, "low", gjson.GetBytes(upstream.lastBody, "text.verbosity").String())
+ require.Equal(t, openAIImagesResponsesMainModel, gjson.GetBytes(upstream.lastBody, "model").String())
+}
+
func TestExtractOpenAIRequestMetaFromBody(t *testing.T) {
tests := []struct {
name string
diff --git a/backend/internal/service/openai_ws_forwarder_success_test.go b/backend/internal/service/openai_ws_forwarder_success_test.go
index e949560f..bd262207 100644
--- a/backend/internal/service/openai_ws_forwarder_success_test.go
+++ b/backend/internal/service/openai_ws_forwarder_success_test.go
@@ -171,6 +171,91 @@ func TestOpenAIGatewayService_Forward_WSv2_SuccessAndBindSticky(t *testing.T) {
require.Equal(t, "resp_new_1", gjson.GetBytes(responseBody, "id").String())
}
+func TestOpenAIGatewayService_Forward_WSv2_UsesPatchedBodyAfterValidationDecode(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ type receivedPayload struct {
+ MaxCompletionTokensExists bool
+ }
+ receivedCh := make(chan receivedPayload, 1)
+
+ upgrader := websocket.Upgrader{CheckOrigin: func(r *http.Request) bool { return true }}
+ wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ conn, err := upgrader.Upgrade(w, r, nil)
+ if err != nil {
+ t.Errorf("upgrade websocket failed: %v", err)
+ return
+ }
+ defer func() { _ = conn.Close() }()
+
+ var request map[string]any
+ if err := conn.ReadJSON(&request); err != nil {
+ t.Errorf("read ws request failed: %v", err)
+ return
+ }
+ requestJSON := requestToJSONString(request)
+ receivedCh <- receivedPayload{MaxCompletionTokensExists: gjson.Get(requestJSON, "max_completion_tokens").Exists()}
+
+ if err := conn.WriteJSON(map[string]any{
+ "type": "response.completed",
+ "response": map[string]any{
+ "id": "resp_patched_ws_1",
+ "model": "gpt-5.3-codex-spark",
+ "usage": map[string]any{"input_tokens": 1, "output_tokens": 1},
+ },
+ }); err != nil {
+ t.Errorf("write response.completed failed: %v", err)
+ return
+ }
+ }))
+ defer wsServer.Close()
+
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ c.Request.Header.Set("User-Agent", "unit-test-agent/1.0")
+
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ cfg.Security.URLAllowlist.AllowInsecureHTTP = true
+ cfg.Gateway.OpenAIWS.Enabled = true
+ cfg.Gateway.OpenAIWS.APIKeyEnabled = true
+ cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
+ cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 30
+ cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 10
+
+ svc := &OpenAIGatewayService{
+ cfg: cfg,
+ openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
+ toolCorrector: NewCodexToolCorrector(),
+ }
+
+ account := &Account{
+ ID: 10,
+ Name: "openai-ws",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": wsServer.URL,
+ },
+ Extra: map[string]any{"responses_websockets_v2_enabled": true},
+ }
+
+ body := []byte(`{"model":"gpt-5.4","stream":false,"max_completion_tokens":12,"tools":[{"type":"image_generation"}],"input":[{"type":"input_text","text":"hello"}]}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.True(t, result.OpenAIWSMode)
+
+ received := <-receivedCh
+ require.False(t, received.MaxCompletionTokensExists)
+}
+
func TestOpenAIGatewayService_Forward_WSv2_ImageGenerationCountsOutputs(t *testing.T) {
gin.SetMode(gin.TestMode)
From f10bca81559b18fa003c3b1bcc4fedf148ddbd7f Mon Sep 17 00:00:00 2001
From: visa2
Date: Sun, 31 May 2026 16:07:52 +0800
Subject: [PATCH 60/79] =?UTF-8?q?refactor(apicompat):=20redesign=20the=20C?=
=?UTF-8?q?odex=20Responses=20=E2=86=94=20Chat=20Completions=20bridge?=
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit
Codex CLI speaks the OpenAI Responses protocol (streaming, store:false), while
many upstreams (e.g. DeepSeek in thinking mode) only expose Chat Completions.
The bridge that translates between the two had grown field by field and leaned
on Go's serialization defaults, which both the Responses client (Codex) and the
Chat upstream reject in ways the official OpenAI endpoints tolerate.
Problems this fixes (all observed running Codex CLI against a DeepSeek upstream):
- Streaming reasoning was never shown in the Codex TUI (the answer appeared
with no visible thinking): reasoning deltas were emitted before the reasoning
item was opened, so the strict client discarded them.
- A tool-using turn could wedge the session into a "no response" state: the
function_call stream was never closed (no function_call_arguments.done /
output_item.done), so Codex never saw the tool call complete.
- Parallel tool calls were rejected upstream (400/502): each function_call
became its own assistant message, producing consecutive assistant messages
with mismatched tool replies.
- A tool turn was rejected with "reasoning_content in the thinking mode must be
passed back": the reasoning that produced the tool call was dropped instead
of being returned on the assistant message.
- Items with no Chat equivalent (web_search_call, ...) and Codex's
command-approval notice landed between an assistant tool_calls message and
its tool reply, triggering "An assistant message with 'tool_calls' must be
followed by tool messages responding to each 'tool_call_id'".
- Interrupt/reconnect left an unanswered or dangling tool_call in the history,
triggering the same 400.
The shared root cause is reliance on serialization defaults — omitempty dropping
protocol-required zero values, and unrecognized item types falling through a
generic path — rather than deliberately reproducing the target protocol. The
bridge is reworked into two explicit layers.
Request direction (Responses input -> Chat messages): a parse -> build ->
normalize pipeline.
- reasoning_content is carried back on the assistant message that produced a
tool call (DeepSeek thinking mode requires it to continue the same thought)
- consecutive function_call items (parallel tool calls) are merged into a
single assistant message's tool_calls array
- item types with no Chat equivalent are skipped instead of leaking through a
generic path
- normalizeChatMessages is the single invariant gate: it guarantees every
assistant tool_calls message is immediately followed by one tool reply per
tool_call_id — reordering any intervening message (such as a command-approval
notice) to after the replies, dropping unanswered tool_calls and orphan tool
replies, and preserving bare passthrough tool messages.
Response direction (Chat SSE -> Responses SSE): ResponsesStreamEvent.MarshalJSON
constructs each streamed event explicitly so protocol-required fields are always
present (output_index/content_index/summary_index at 0, message content:[],
reasoning summary:[], function_call call_id/name/arguments, output_text part
text/annotations/logprobs). This is a single source of truth that removes any
post-hoc JSON patching. Reasoning is emitted as its own output item, opened
before its deltas, and tool calls are fully closed
(function_call_arguments.done + output_item.done with complete arguments).
Tests cover request-direction message invariants against golden Codex request
shapes (parallel calls, unknown items, intervening messages, partial/dangling
calls), per-event wire completeness, and streaming lifecycle ordering.
Co-Authored-By: Claude Opus 4.8
---
.../chatcompletions_responses_bridge.go | 476 +++++++++++++++---
...tions_responses_request_invariants_test.go | 187 +++++++
...letions_responses_stream_lifecycle_test.go | 103 ++++
.../apicompat/responses_stream_event_wire.go | 199 ++++++++
.../responses_stream_event_wire_test.go | 109 ++++
backend/internal/pkg/apicompat/types.go | 4 +
6 files changed, 1019 insertions(+), 59 deletions(-)
create mode 100644 backend/internal/pkg/apicompat/chatcompletions_responses_request_invariants_test.go
create mode 100644 backend/internal/pkg/apicompat/chatcompletions_responses_stream_lifecycle_test.go
create mode 100644 backend/internal/pkg/apicompat/responses_stream_event_wire.go
create mode 100644 backend/internal/pkg/apicompat/responses_stream_event_wire_test.go
diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go b/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go
index 09b680c7..cc51cba2 100644
--- a/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go
+++ b/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go
@@ -42,14 +42,24 @@ func ResponsesToChatCompletionsRequest(req *ResponsesRequest) (*ChatCompletionsR
return out, nil
}
+// responsesInputToChatMessages converts a Responses request's instructions +
+// input[] into Chat Completions messages. It is a three-stage pipeline:
+//
+// parse — instructions become a system message; input[] is split into items
+// build — buildChatMessagesFromItems walks items, attaching reasoning to the
+// assistant message that produced a tool call, merging parallel tool
+// calls into one assistant message, and skipping item types that have
+// no Chat equivalent
+// normalize — normalizeChatMessages enforces the invariants DeepSeek requires
+//
+// The build + normalize split keeps every protocol rule in one place rather than
+// scattered across per-item cases, and makes unknown future codex item types
+// fail safe instead of leaking into the upstream request.
func responsesInputToChatMessages(instructions string, inputRaw json.RawMessage) ([]ChatMessage, error) {
var messages []ChatMessage
if strings.TrimSpace(instructions) != "" {
content, _ := json.Marshal(instructions)
- messages = append(messages, ChatMessage{
- Role: "system",
- Content: content,
- })
+ messages = append(messages, ChatMessage{Role: "system", Content: content})
}
inputRaw = bytesTrimSpace(inputRaw)
@@ -57,13 +67,11 @@ func responsesInputToChatMessages(instructions string, inputRaw json.RawMessage)
return messages, nil
}
+ // Bare string input is a single user turn.
var inputText string
if err := json.Unmarshal(inputRaw, &inputText); err == nil {
content, _ := json.Marshal(inputText)
- messages = append(messages, ChatMessage{
- Role: "user",
- Content: content,
- })
+ messages = append(messages, ChatMessage{Role: "user", Content: content})
return messages, nil
}
@@ -72,6 +80,24 @@ func responsesInputToChatMessages(instructions string, inputRaw json.RawMessage)
return nil, fmt.Errorf("parse responses input: %w", err)
}
+ built, err := buildChatMessagesFromItems(messages, rawItems)
+ if err != nil {
+ return nil, err
+ }
+ return normalizeChatMessages(built), nil
+}
+
+// buildChatMessagesFromItems walks the Responses input items and appends the
+// corresponding Chat messages.
+func buildChatMessagesFromItems(messages []ChatMessage, rawItems []json.RawMessage) ([]ChatMessage, error) {
+ // pendingReasoning holds the reasoning text from a reasoning item until the
+ // assistant message it belongs to is emitted. DeepSeek's thinking mode
+ // requires the reasoning_content that produced a tool call to be passed back
+ // on that assistant message; dropping it yields a 400. It only survives
+ // across an assistant message (so a following tool call in the same turn
+ // still receives it); any other role ends the thinking span.
+ var pendingReasoning string
+
for _, raw := range rawItems {
raw = bytesTrimSpace(raw)
if len(raw) == 0 || string(raw) == "null" {
@@ -84,6 +110,7 @@ func responsesInputToChatMessages(instructions string, inputRaw json.RawMessage)
if textErr := json.Unmarshal(raw, &text); textErr == nil {
content, _ := json.Marshal(text)
messages = append(messages, ChatMessage{Role: "user", Content: content})
+ pendingReasoning = ""
continue
}
return nil, fmt.Errorf("parse responses input item: %w", err)
@@ -92,22 +119,40 @@ func responsesInputToChatMessages(instructions string, inputRaw json.RawMessage)
role := chatCompletionsBridgeRole(rawString(item["role"]))
itemType := rawString(item["type"])
switch itemType {
+ case "reasoning":
+ if txt := extractResponsesReasoningText(item); txt != "" {
+ pendingReasoning = txt
+ }
+ continue
case "function_call":
arguments := rawString(item["arguments"])
if strings.TrimSpace(arguments) == "" {
arguments = "{}"
}
- messages = append(messages, ChatMessage{
- Role: "assistant",
- ToolCalls: []ChatToolCall{{
- ID: rawString(item["call_id"]),
- Type: "function",
- Function: ChatFunctionCall{
- Name: rawString(item["name"]),
- Arguments: arguments,
- },
- }},
- })
+ toolCall := ChatToolCall{
+ ID: rawString(item["call_id"]),
+ Type: "function",
+ Function: ChatFunctionCall{
+ Name: rawString(item["name"]),
+ Arguments: arguments,
+ },
+ }
+ // Parallel tool calls arrive as consecutive function_call items and
+ // must share one assistant message; the matching tool replies then
+ // follow it. Merge into the immediately preceding assistant message.
+ if n := len(messages); n > 0 && messages[n-1].Role == "assistant" {
+ messages[n-1].ToolCalls = append(messages[n-1].ToolCalls, toolCall)
+ if messages[n-1].ReasoningContent == "" {
+ messages[n-1].ReasoningContent = pendingReasoning
+ }
+ } else {
+ messages = append(messages, ChatMessage{
+ Role: "assistant",
+ ToolCalls: []ChatToolCall{toolCall},
+ ReasoningContent: pendingReasoning,
+ })
+ }
+ pendingReasoning = ""
continue
case "function_call_output":
content, _ := json.Marshal(rawString(item["output"]))
@@ -116,10 +161,12 @@ func responsesInputToChatMessages(instructions string, inputRaw json.RawMessage)
ToolCallID: rawString(item["call_id"]),
Content: content,
})
+ pendingReasoning = ""
continue
case "input_text", "text":
content, _ := json.Marshal(rawString(item["text"]))
messages = append(messages, ChatMessage{Role: "user", Content: content})
+ pendingReasoning = ""
continue
case "input_image":
content, err := chatContentFromSingleResponsesPart(itemType, item)
@@ -127,6 +174,18 @@ func responsesInputToChatMessages(instructions string, inputRaw json.RawMessage)
return nil, err
}
messages = append(messages, ChatMessage{Role: "user", Content: content})
+ pendingReasoning = ""
+ continue
+ }
+
+ // Only genuine message items become chat messages. Codex emits other
+ // Responses item types with no Chat equivalent (web_search_call,
+ // local_shell_call, custom tool calls, file_search_call, ...). Converting
+ // them via the generic path would insert a spurious message between an
+ // assistant tool_calls message and its tool reply, which DeepSeek rejects
+ // ("insufficient tool messages following tool_calls message"). Skip them.
+ if itemType != "" && itemType != "message" {
+ pendingReasoning = ""
continue
}
@@ -140,15 +199,128 @@ func responsesInputToChatMessages(instructions string, inputRaw json.RawMessage)
if err != nil {
return nil, err
}
- messages = append(messages, ChatMessage{
- Role: role,
- Content: chatContent,
- })
+ messages = append(messages, ChatMessage{Role: role, Content: chatContent})
+ // Reasoning only survives across an assistant text message.
+ if role != "assistant" {
+ pendingReasoning = ""
+ }
}
return messages, nil
}
+// normalizeChatMessages is the single place that enforces the tool-call
+// invariant the DeepSeek / OpenAI Chat Completions schema requires: an assistant
+// message with tool_calls must be immediately followed by one tool message per
+// tool_call_id, in order, with nothing in between.
+//
+// Codex histories violate this in several ways that the builder alone can't fix:
+// - a non-tool message lands between an assistant tool_calls message and its
+// tool replies (e.g. an "Approved command prefix saved" system notice codex
+// injects mid tool-execution);
+// - a parallel tool_call's sibling output never arrives, or a call is left
+// dangling by a mid-execution reconnect (unanswered tool_call);
+// - a tool reply has no announcing assistant tool_call (orphan).
+//
+// It rebuilds the sequence so each assistant's answered tool_calls are followed
+// directly by their replies (in call order); unanswered tool_calls are dropped
+// (and an assistant left with neither tool_calls nor content is dropped); orphan
+// tool replies and intervening messages are emitted in their natural position
+// but never between an assistant tool_calls message and its replies.
+func normalizeChatMessages(messages []ChatMessage) []ChatMessage {
+ // Index every tool reply by its tool_call_id (last wins on duplicates).
+ replies := make(map[string]ChatMessage)
+ for _, m := range messages {
+ if m.Role == "tool" && m.ToolCallID != "" {
+ replies[m.ToolCallID] = m
+ }
+ }
+
+ out := make([]ChatMessage, 0, len(messages))
+ for _, m := range messages {
+ switch {
+ case m.Role == "tool":
+ // A bare tool message with no tool_call_id is a direct Chat
+ // Completions passthrough; keep it in place. A tool reply whose id is
+ // announced by an assistant is emitted right after that assistant
+ // (skip the standalone occurrence). Any other tool reply is an orphan
+ // and is dropped.
+ if m.ToolCallID == "" {
+ out = append(out, m)
+ }
+ continue
+ case len(m.ToolCalls) > 0:
+ kept := make([]ChatToolCall, 0, len(m.ToolCalls))
+ for _, tc := range m.ToolCalls {
+ if tc.ID == "" {
+ continue
+ }
+ if _, ok := replies[tc.ID]; ok {
+ kept = append(kept, tc)
+ }
+ }
+ if len(kept) == 0 {
+ // No answered tool_calls left: keep as a plain message if it has
+ // content, otherwise drop it entirely.
+ if isBlankChatContent(m.Content) {
+ continue
+ }
+ m.ToolCalls = nil
+ out = append(out, m)
+ continue
+ }
+ m.ToolCalls = kept
+ out = append(out, m)
+ for _, tc := range kept {
+ out = append(out, replies[tc.ID])
+ }
+ default:
+ out = append(out, m)
+ }
+ }
+ return out
+}
+
+// isBlankChatContent reports whether a chat message content holds no usable text.
+func isBlankChatContent(raw json.RawMessage) bool {
+ raw = bytesTrimSpace(raw)
+ if len(raw) == 0 || string(raw) == "null" || string(raw) == `""` {
+ return true
+ }
+ return chatMessageContentText(raw) == ""
+}
+
+// extractResponsesReasoningText pulls the reasoning text out of a Responses
+// reasoning item. The Chat→Responses bridge writes the upstream reasoning_content
+// verbatim into the summary_text parts (see closeChatReasoningItem), so codex
+// round-trips it there; prefer summary[].text and fall back to content.
+func extractResponsesReasoningText(item map[string]json.RawMessage) string {
+ var parts []string
+ collect := func(raw json.RawMessage) {
+ raw = bytesTrimSpace(raw)
+ if len(raw) == 0 || string(raw) == "null" {
+ return
+ }
+ var arr []map[string]json.RawMessage
+ if err := json.Unmarshal(raw, &arr); err == nil {
+ for _, p := range arr {
+ if t := rawString(p["text"]); t != "" {
+ parts = append(parts, t)
+ }
+ }
+ return
+ }
+ if t := rawString(raw); t != "" {
+ parts = append(parts, t)
+ }
+ }
+ collect(item["summary"])
+ if len(parts) == 0 {
+ collect(item["content"])
+ }
+ return strings.Join(parts, "\n")
+}
+
func chatCompletionsBridgeRole(role string) string {
trimmed := strings.TrimSpace(role)
if trimmed == "" {
@@ -448,10 +620,32 @@ type ChatCompletionsToResponsesStreamState struct {
CreatedSent bool
CompletedSent bool
+ // nextOutputIndex assigns sequential output_index values to items as they
+ // are opened (reasoning, message, tool calls), so the streamed indices match
+ // the order of items in the final response.output array.
+ nextOutputIndex int
+
+ // Reasoning item lifecycle. DeepSeek-style upstreams stream all
+ // reasoning_content before any content, so reasoning is modeled as its own
+ // "reasoning" output item that must be opened (output_item.added) before any
+ // reasoning delta and closed before the message/tool items open.
+ ReasoningItemID string
+ ReasoningIndex int
+ ReasoningOpen bool
+ ReasoningDone bool
+
+ // Message item + output_text content-part lifecycle.
MessageItemID string
- Text strings.Builder
- Reasoning strings.Builder
- ToolCalls map[int]*ChatToolCall
+ MessageIndex int
+ TextPartOpen bool
+
+ Text strings.Builder
+ Reasoning strings.Builder
+
+ // Tool-call lifecycle, keyed by the upstream tool_call index.
+ ToolCalls map[int]*ChatToolCall
+ ToolItemIDs map[int]string
+ ToolOutputIndex map[int]int
FinishReason string
Usage *ResponsesUsage
@@ -460,13 +654,21 @@ type ChatCompletionsToResponsesStreamState struct {
// NewChatCompletionsToResponsesStreamState returns an initialized stream state.
func NewChatCompletionsToResponsesStreamState(model string) *ChatCompletionsToResponsesStreamState {
return &ChatCompletionsToResponsesStreamState{
- ResponseID: generateResponsesID(),
- Model: model,
- Created: time.Now().Unix(),
- ToolCalls: make(map[int]*ChatToolCall),
+ ResponseID: generateResponsesID(),
+ Model: model,
+ Created: time.Now().Unix(),
+ ToolCalls: make(map[int]*ChatToolCall),
+ ToolItemIDs: make(map[int]string),
+ ToolOutputIndex: make(map[int]int),
}
}
+func (state *ChatCompletionsToResponsesStreamState) allocOutputIndex() int {
+ idx := state.nextOutputIndex
+ state.nextOutputIndex++
+ return idx
+}
+
// ChatCompletionsChunkToResponsesEvents converts one Chat Completions stream
// chunk into zero or more Responses stream events.
func ChatCompletionsChunkToResponsesEvents(
@@ -490,24 +692,34 @@ func ChatCompletionsChunkToResponsesEvents(
events = append(events, ensureChatToResponsesCreated(state)...)
for _, choice := range chunk.Choices {
- if choice.Delta.Content != nil {
+ // Reasoning is emitted as its own output item and must be opened
+ // (output_item.added + reasoning_summary_part.added) before the first
+ // delta, otherwise a strict client discards the delta. The leading
+ // empty-string reasoning delta upstreams send is filtered out.
+ if choice.Delta.ReasoningContent != nil && *choice.Delta.ReasoningContent != "" {
+ events = append(events, ensureChatReasoningItem(state)...)
+ _, _ = state.Reasoning.WriteString(*choice.Delta.ReasoningContent)
+ events = append(events, chatToResponsesEvent(state, "response.reasoning_summary_text.delta", &ResponsesStreamEvent{
+ OutputIndex: state.ReasoningIndex,
+ SummaryIndex: 0,
+ Delta: *choice.Delta.ReasoningContent,
+ ItemID: state.ReasoningItemID,
+ }))
+ }
+ if choice.Delta.Content != nil && *choice.Delta.Content != "" {
+ // First real content closes the reasoning item, then opens the
+ // message item and its output_text content part.
+ events = append(events, closeChatReasoningItem(state)...)
events = append(events, ensureChatToResponsesMessageItem(state)...)
+ events = append(events, ensureChatToResponsesTextPart(state)...)
_, _ = state.Text.WriteString(*choice.Delta.Content)
events = append(events, chatToResponsesEvent(state, "response.output_text.delta", &ResponsesStreamEvent{
- OutputIndex: 0,
+ OutputIndex: state.MessageIndex,
ContentIndex: 0,
Delta: *choice.Delta.Content,
ItemID: state.MessageItemID,
}))
}
- if choice.Delta.ReasoningContent != nil {
- _, _ = state.Reasoning.WriteString(*choice.Delta.ReasoningContent)
- events = append(events, chatToResponsesEvent(state, "response.reasoning_summary_text.delta", &ResponsesStreamEvent{
- OutputIndex: 0,
- SummaryIndex: 0,
- Delta: *choice.Delta.ReasoningContent,
- }))
- }
for _, toolCall := range choice.Delta.ToolCalls {
idx := 0
if toolCall.Index != nil {
@@ -515,6 +727,8 @@ func ChatCompletionsChunkToResponsesEvents(
}
stored, ok := state.ToolCalls[idx]
if !ok {
+ // A tool call closes any open reasoning item first.
+ events = append(events, closeChatReasoningItem(state)...)
copyCall := toolCall
if copyCall.ID == "" {
copyCall.ID = generateItemID()
@@ -522,11 +736,14 @@ func ChatCompletionsChunkToResponsesEvents(
copyCall.Type = "function"
state.ToolCalls[idx] = ©Call
stored = ©Call
+ itemID := generateItemID()
+ state.ToolItemIDs[idx] = itemID
+ state.ToolOutputIndex[idx] = state.allocOutputIndex()
events = append(events, chatToResponsesEvent(state, "response.output_item.added", &ResponsesStreamEvent{
- OutputIndex: idx + 1,
+ OutputIndex: state.ToolOutputIndex[idx],
Item: &ResponsesOutput{
Type: "function_call",
- ID: generateItemID(),
+ ID: itemID,
CallID: stored.ID,
Name: stored.Function.Name,
Status: "in_progress",
@@ -543,7 +760,8 @@ func ChatCompletionsChunkToResponsesEvents(
if toolCall.Function.Arguments != "" {
stored.Function.Arguments += toolCall.Function.Arguments
events = append(events, chatToResponsesEvent(state, "response.function_call_arguments.delta", &ResponsesStreamEvent{
- OutputIndex: idx + 1,
+ OutputIndex: state.ToolOutputIndex[idx],
+ ItemID: state.ToolItemIDs[idx],
Delta: toolCall.Function.Arguments,
CallID: stored.ID,
Name: stored.Function.Name,
@@ -565,24 +783,44 @@ func FinalizeChatCompletionsResponsesStream(state *ChatCompletionsToResponsesStr
}
var events []ResponsesStreamEvent
events = append(events, ensureChatToResponsesCreated(state)...)
+
+ // Close a reasoning item that never transitioned to content (reasoning-only
+ // or empty completion).
+ events = append(events, closeChatReasoningItem(state)...)
+
if state.MessageItemID != "" {
- events = append(events, chatToResponsesEvent(state, "response.output_text.done", &ResponsesStreamEvent{
- OutputIndex: 0,
- ContentIndex: 0,
- Text: state.Text.String(),
- ItemID: state.MessageItemID,
- }))
+ if state.TextPartOpen {
+ events = append(events, chatToResponsesEvent(state, "response.output_text.done", &ResponsesStreamEvent{
+ OutputIndex: state.MessageIndex,
+ ContentIndex: 0,
+ Text: state.Text.String(),
+ ItemID: state.MessageItemID,
+ }))
+ events = append(events, chatToResponsesEvent(state, "response.content_part.done", &ResponsesStreamEvent{
+ OutputIndex: state.MessageIndex,
+ ContentIndex: 0,
+ ItemID: state.MessageItemID,
+ Part: &ResponsesContentPart{Type: "output_text", Text: state.Text.String()},
+ }))
+ }
events = append(events, chatToResponsesEvent(state, "response.output_item.done", &ResponsesStreamEvent{
- OutputIndex: 0,
+ OutputIndex: state.MessageIndex,
Item: &ResponsesOutput{
- Type: "message",
- ID: state.MessageItemID,
- Role: "assistant",
- Status: "completed",
+ Type: "message",
+ ID: state.MessageItemID,
+ Role: "assistant",
+ Content: []ResponsesContentPart{{Type: "output_text", Text: state.Text.String()}},
+ Status: "completed",
},
}))
}
+ // Close every function_call item opened during the stream. Codex finalizes a
+ // tool call only after function_call_arguments.done + output_item.done for
+ // that item; without them the call never completes and the session wedges.
+ // Mirrors cc-switch's finalize_tools.
+ events = append(events, closeChatToolItems(state)...)
+
status := "completed"
var incompleteDetails *ResponsesIncompleteDetails
if state.FinishReason == "length" {
@@ -621,22 +859,142 @@ func ensureChatToResponsesCreated(state *ChatCompletionsToResponsesStreamState)
})}
}
+// ensureChatReasoningItem opens the reasoning output item (output_item.added +
+// reasoning_summary_part.added) before the first reasoning delta. Codex renders
+// streaming reasoning only when this summary-part lifecycle is present.
+func ensureChatReasoningItem(state *ChatCompletionsToResponsesStreamState) []ResponsesStreamEvent {
+ if state.ReasoningOpen || state.ReasoningDone {
+ return nil
+ }
+ state.ReasoningOpen = true
+ state.ReasoningItemID = generateItemID()
+ state.ReasoningIndex = state.allocOutputIndex()
+ return []ResponsesStreamEvent{
+ chatToResponsesEvent(state, "response.output_item.added", &ResponsesStreamEvent{
+ OutputIndex: state.ReasoningIndex,
+ Item: &ResponsesOutput{Type: "reasoning", ID: state.ReasoningItemID, Status: "in_progress"},
+ }),
+ chatToResponsesEvent(state, "response.reasoning_summary_part.added", &ResponsesStreamEvent{
+ OutputIndex: state.ReasoningIndex,
+ SummaryIndex: 0,
+ ItemID: state.ReasoningItemID,
+ Part: &ResponsesContentPart{Type: "summary_text"},
+ }),
+ }
+}
+
+// closeChatReasoningItem emits the reasoning item's terminal events
+// (reasoning_summary_text.done + reasoning_summary_part.done + output_item.done).
+func closeChatReasoningItem(state *ChatCompletionsToResponsesStreamState) []ResponsesStreamEvent {
+ if !state.ReasoningOpen {
+ return nil
+ }
+ state.ReasoningOpen = false
+ state.ReasoningDone = true
+ reasoning := state.Reasoning.String()
+ return []ResponsesStreamEvent{
+ chatToResponsesEvent(state, "response.reasoning_summary_text.done", &ResponsesStreamEvent{
+ OutputIndex: state.ReasoningIndex,
+ SummaryIndex: 0,
+ Text: reasoning,
+ ItemID: state.ReasoningItemID,
+ }),
+ chatToResponsesEvent(state, "response.reasoning_summary_part.done", &ResponsesStreamEvent{
+ OutputIndex: state.ReasoningIndex,
+ SummaryIndex: 0,
+ ItemID: state.ReasoningItemID,
+ Part: &ResponsesContentPart{Type: "summary_text", Text: reasoning},
+ }),
+ chatToResponsesEvent(state, "response.output_item.done", &ResponsesStreamEvent{
+ OutputIndex: state.ReasoningIndex,
+ Item: &ResponsesOutput{
+ Type: "reasoning",
+ ID: state.ReasoningItemID,
+ Status: "completed",
+ Summary: []ResponsesSummary{{Type: "summary_text", Text: reasoning}},
+ },
+ }),
+ }
+}
+
func ensureChatToResponsesMessageItem(state *ChatCompletionsToResponsesStreamState) []ResponsesStreamEvent {
if state.MessageItemID != "" {
return nil
}
state.MessageItemID = generateItemID()
+ state.MessageIndex = state.allocOutputIndex()
return []ResponsesStreamEvent{chatToResponsesEvent(state, "response.output_item.added", &ResponsesStreamEvent{
- OutputIndex: 0,
+ OutputIndex: state.MessageIndex,
Item: &ResponsesOutput{
- Type: "message",
- ID: state.MessageItemID,
- Role: "assistant",
- Status: "in_progress",
+ Type: "message",
+ ID: state.MessageItemID,
+ Role: "assistant",
+ Status: "in_progress",
+ Content: []ResponsesContentPart{{Type: "output_text"}},
},
})}
}
+func ensureChatToResponsesTextPart(state *ChatCompletionsToResponsesStreamState) []ResponsesStreamEvent {
+ if state.TextPartOpen {
+ return nil
+ }
+ state.TextPartOpen = true
+ return []ResponsesStreamEvent{chatToResponsesEvent(state, "response.content_part.added", &ResponsesStreamEvent{
+ OutputIndex: state.MessageIndex,
+ ContentIndex: 0,
+ ItemID: state.MessageItemID,
+ Part: &ResponsesContentPart{Type: "output_text", Text: ""},
+ })}
+}
+
+// closeChatToolItems emits function_call_arguments.done + output_item.done for
+// every tool call opened during the stream, carrying the full call_id/name/
+// arguments so codex can deserialize and execute the call. Mirrors cc-switch's
+// finalize_tools.
+func closeChatToolItems(state *ChatCompletionsToResponsesStreamState) []ResponsesStreamEvent {
+ if len(state.ToolCalls) == 0 {
+ return nil
+ }
+ var events []ResponsesStreamEvent
+ for i := 0; i < len(state.ToolCalls); i++ {
+ toolCall, ok := state.ToolCalls[i]
+ if !ok || toolCall == nil {
+ continue
+ }
+ itemID, opened := state.ToolItemIDs[i]
+ if !opened {
+ continue
+ }
+ arguments := toolCall.Function.Arguments
+ if strings.TrimSpace(arguments) == "" {
+ arguments = "{}"
+ }
+ outputIndex := state.ToolOutputIndex[i]
+ events = append(events,
+ chatToResponsesEvent(state, "response.function_call_arguments.done", &ResponsesStreamEvent{
+ OutputIndex: outputIndex,
+ ItemID: itemID,
+ CallID: toolCall.ID,
+ Name: toolCall.Function.Name,
+ Arguments: arguments,
+ }),
+ chatToResponsesEvent(state, "response.output_item.done", &ResponsesStreamEvent{
+ OutputIndex: outputIndex,
+ Item: &ResponsesOutput{
+ Type: "function_call",
+ ID: itemID,
+ CallID: toolCall.ID,
+ Name: toolCall.Function.Name,
+ Arguments: arguments,
+ Status: "completed",
+ },
+ }),
+ )
+ }
+ return events
+}
+
func (state *ChatCompletionsToResponsesStreamState) chatOutput() []ResponsesOutput {
var outputs []ResponsesOutput
if state.Reasoning.Len() > 0 {
diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_request_invariants_test.go b/backend/internal/pkg/apicompat/chatcompletions_responses_request_invariants_test.go
new file mode 100644
index 00000000..e54a4532
--- /dev/null
+++ b/backend/internal/pkg/apicompat/chatcompletions_responses_request_invariants_test.go
@@ -0,0 +1,187 @@
+package apicompat
+
+import (
+ "encoding/json"
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+// assertChatInvariants enforces the DeepSeek / OpenAI Chat Completions message
+// invariants that, when violated, surface as upstream 400s. Used to validate the
+// request-direction converter against golden codex request shapes.
+func assertChatInvariants(t *testing.T, messages []ChatMessage) {
+ t.Helper()
+ for i, m := range messages {
+ // Every assistant tool_calls message must be immediately followed by one
+ // tool message per tool_call_id, in order.
+ if len(m.ToolCalls) > 0 {
+ for j, tc := range m.ToolCalls {
+ k := i + 1 + j
+ require.Lessf(t, k, len(messages), "tool_call %s has no following tool message", tc.ID)
+ require.Equalf(t, "tool", messages[k].Role, "tool_call %s not followed by a tool message", tc.ID)
+ require.Equalf(t, tc.ID, messages[k].ToolCallID, "tool reply order mismatch for %s", tc.ID)
+ }
+ }
+ // No two consecutive assistant messages.
+ if i > 0 && m.Role == "assistant" && messages[i-1].Role == "assistant" {
+ t.Fatalf("consecutive assistant messages at %d", i)
+ }
+ // No orphan tool replies.
+ if m.Role == "tool" {
+ require.NotEmptyf(t, m.ToolCallID, "tool message without tool_call_id at %d", i)
+ }
+ }
+}
+
+func convertGolden(t *testing.T, input string) []ChatMessage {
+ t.Helper()
+ msgs, err := responsesInputToChatMessages("You are a helpful assistant.", json.RawMessage(input))
+ require.NoError(t, err)
+ return msgs
+}
+
+// Golden sample: a single tool-call turn (codex runs one shell/curl command),
+// the shape that produced the original "no response" / 400.
+func TestGolden_SingleToolCall(t *testing.T) {
+ msgs := convertGolden(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"latest sha?"}]},
+ {"type":"reasoning","summary":[{"type":"summary_text","text":"need to run curl"}]},
+ {"type":"function_call","call_id":"call_a","name":"exec_command","arguments":"{\"cmd\":\"curl x\"}"},
+ {"type":"function_call_output","call_id":"call_a","output":"deadbeef"}
+ ]`)
+ assertChatInvariants(t, msgs)
+ // reasoning_content must ride on the assistant tool-call message.
+ var asst *ChatMessage
+ for i := range msgs {
+ if len(msgs[i].ToolCalls) > 0 {
+ asst = &msgs[i]
+ }
+ }
+ require.NotNil(t, asst)
+ require.Equal(t, "need to run curl", asst.ReasoningContent)
+}
+
+// Golden sample: parallel tool calls (codex runs git log + git tag at once).
+func TestGolden_ParallelToolCalls(t *testing.T) {
+ msgs := convertGolden(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"features?"}]},
+ {"type":"reasoning","summary":[{"type":"summary_text","text":"inspect repo"}]},
+ {"type":"function_call","call_id":"c0","name":"exec_command","arguments":"{\"cmd\":\"git log\"}"},
+ {"type":"function_call","call_id":"c1","name":"exec_command","arguments":"{\"cmd\":\"git tag\"}"},
+ {"type":"function_call_output","call_id":"c0","output":"log"},
+ {"type":"function_call_output","call_id":"c1","output":"tags"}
+ ]`)
+ assertChatInvariants(t, msgs)
+ // Both parallel calls share ONE assistant message.
+ var toolMsgs int
+ for _, m := range msgs {
+ if len(m.ToolCalls) == 2 {
+ require.Equal(t, "c0", m.ToolCalls[0].ID)
+ require.Equal(t, "c1", m.ToolCalls[1].ID)
+ }
+ if m.Role == "tool" {
+ toolMsgs++
+ }
+ }
+ require.Equal(t, 2, toolMsgs)
+}
+
+// Golden sample: an unknown item type (web_search_call from a 联网查询) sitting
+// between a function_call and its output must not break tool↔reply adjacency.
+func TestGolden_UnknownItemBetweenToolCallAndOutput(t *testing.T) {
+ msgs := convertGolden(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"search"}]},
+ {"type":"reasoning","summary":[{"type":"summary_text","text":"let me search"}]},
+ {"type":"function_call","call_id":"c0","name":"exec_command","arguments":"{}"},
+ {"type":"web_search_call","id":"ws_1","status":"completed","action":{"type":"search","query":"x"}},
+ {"type":"function_call_output","call_id":"c0","output":"result"}
+ ]`)
+ assertChatInvariants(t, msgs)
+}
+
+// Sequential tool calls (a tool reply between two calls) must stay in distinct
+// assistant messages.
+func TestRequest_SequentialToolCallsStaySeparate(t *testing.T) {
+ msgs := convertGolden(t, `[
+ {"type":"function_call","call_id":"c1","name":"exec","arguments":"{}"},
+ {"type":"function_call_output","call_id":"c1","output":"r1"},
+ {"type":"function_call","call_id":"c2","name":"exec","arguments":"{}"},
+ {"type":"function_call_output","call_id":"c2","output":"r2"}
+ ]`)
+ assertChatInvariants(t, msgs)
+ assistants := 0
+ for _, m := range msgs {
+ if len(m.ToolCalls) == 1 {
+ assistants++
+ }
+ }
+ require.Equal(t, 2, assistants)
+}
+
+// Golden sample: codex injects a message (e.g. an "Approved command prefix
+// saved" notice) between a function_call and its output. The intervening message
+// must be moved after the tool reply so the assistant tool_calls is immediately
+// followed by its reply.
+func TestGolden_MessageBetweenToolCallAndOutput(t *testing.T) {
+ msgs := convertGolden(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"do it"}]},
+ {"type":"reasoning","summary":[{"type":"summary_text","text":"run cmd"}]},
+ {"type":"function_call","call_id":"A","name":"exec","arguments":"{}"},
+ {"type":"message","role":"developer","content":[{"type":"input_text","text":"Approved command prefix saved"}]},
+ {"type":"function_call_output","call_id":"A","output":"ok"}
+ ]`)
+ assertChatInvariants(t, msgs)
+ // The assistant tool_calls message is immediately followed by its tool reply.
+ for i, m := range msgs {
+ if len(m.ToolCalls) > 0 {
+ require.Equal(t, "tool", msgs[i+1].Role)
+ require.Equal(t, "A", msgs[i+1].ToolCallID)
+ }
+ }
+}
+
+// Golden sample: a parallel tool call where one sibling's output is missing
+// (codex interrupted/reconnected mid-execution). The unanswered tool_call must
+// be dropped so the remaining assistant tool_calls are all answered.
+func TestGolden_PartialParallelDropsUnansweredCall(t *testing.T) {
+ msgs := convertGolden(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"q"}]},
+ {"type":"reasoning","summary":[{"type":"summary_text","text":"r"}]},
+ {"type":"function_call","call_id":"A","name":"exec","arguments":"{}"},
+ {"type":"function_call","call_id":"B","name":"exec","arguments":"{}"},
+ {"type":"function_call_output","call_id":"A","output":"oa"}
+ ]`)
+ assertChatInvariants(t, msgs)
+ for _, m := range msgs {
+ for _, tc := range m.ToolCalls {
+ require.NotEqual(t, "B", tc.ID, "unanswered tool_call B should have been dropped")
+ }
+ }
+}
+
+// Golden sample: a dangling tool_call at the end of the history (no output yet).
+// The assistant message holding only that call must be dropped entirely.
+func TestGolden_DanglingToolCallDropped(t *testing.T) {
+ msgs := convertGolden(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"q"}]},
+ {"type":"reasoning","summary":[{"type":"summary_text","text":"r"}]},
+ {"type":"function_call","call_id":"A","name":"exec","arguments":"{}"}
+ ]`)
+ assertChatInvariants(t, msgs)
+ for _, m := range msgs {
+ require.Empty(t, m.ToolCalls, "dangling unanswered tool_call should have been dropped")
+ }
+}
+
+// normalizeChatMessages drops an orphan tool reply whose tool_call was never
+// announced.
+func TestNormalize_DropsOrphanToolReply(t *testing.T) {
+ msgs := convertGolden(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"q"}]},
+ {"type":"function_call_output","call_id":"ghost","output":"orphan"}
+ ]`)
+ for _, m := range msgs {
+ require.NotEqualf(t, "tool", m.Role, "orphan tool reply should have been dropped")
+ }
+}
diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_stream_lifecycle_test.go b/backend/internal/pkg/apicompat/chatcompletions_responses_stream_lifecycle_test.go
new file mode 100644
index 00000000..beb47303
--- /dev/null
+++ b/backend/internal/pkg/apicompat/chatcompletions_responses_stream_lifecycle_test.go
@@ -0,0 +1,103 @@
+package apicompat
+
+import (
+ "encoding/json"
+ "strings"
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+func collectStreamEvents(t *testing.T, chunks []string) []ResponsesStreamEvent {
+ t.Helper()
+ state := NewChatCompletionsToResponsesStreamState("deepseek-v4-pro")
+ var events []ResponsesStreamEvent
+ for _, payload := range chunks {
+ var chunk ChatCompletionsChunk
+ require.NoError(t, json.Unmarshal([]byte(payload), &chunk))
+ events = append(events, ChatCompletionsChunkToResponsesEvents(&chunk, state)...)
+ }
+ events = append(events, FinalizeChatCompletionsResponsesStream(state)...)
+ return events
+}
+
+// TestStream_ReasoningOpensItemBeforeDelta guards the bug where a strict client
+// (Codex) drops reasoning deltas that reference an item not yet opened.
+func TestStream_ReasoningOpensItemBeforeDelta(t *testing.T) {
+ events := collectStreamEvents(t, []string{
+ `{"choices":[{"index":0,"delta":{"role":"assistant","content":null,"reasoning_content":""}}]}`,
+ `{"choices":[{"index":0,"delta":{"reasoning_content":"think"}}]}`,
+ `{"choices":[{"index":0,"delta":{"content":"hello"}}]}`,
+ `{"choices":[{"index":0,"delta":{"content":""},"finish_reason":"stop"}],"usage":{"prompt_tokens":1,"completion_tokens":2,"total_tokens":3}}`,
+ })
+
+ open := map[int]string{} // output_index -> item type
+ for _, e := range events {
+ switch e.Type {
+ case "response.output_item.added":
+ require.NotNil(t, e.Item)
+ open[e.OutputIndex] = e.Item.Type
+ case "response.reasoning_summary_text.delta":
+ require.Equalf(t, "reasoning", open[e.OutputIndex], "reasoning delta before its item was opened")
+ case "response.output_text.delta":
+ require.Equalf(t, "message", open[e.OutputIndex], "text delta before its item was opened")
+ }
+ }
+}
+
+// TestStream_ToolCallLifecycleComplete guards that a tool call is fully closed
+// (function_call_arguments.done + output_item.done with full arguments), which
+// codex needs to execute the call.
+func TestStream_ToolCallLifecycleComplete(t *testing.T) {
+ events := collectStreamEvents(t, []string{
+ `{"choices":[{"index":0,"delta":{"role":"assistant","reasoning_content":"plan"}}]}`,
+ `{"choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"id":"call_a","type":"function","function":{"name":"exec","arguments":""}}]}}]}`,
+ `{"choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"arguments":"{\"cmd\":\"ls\"}"}}]}}]}`,
+ `{"choices":[{"index":0,"delta":{},"finish_reason":"tool_calls"}],"usage":{"prompt_tokens":1,"completion_tokens":2,"total_tokens":3}}`,
+ })
+
+ var sawAdded, sawArgsDone, sawItemDone bool
+ for _, e := range events {
+ switch e.Type {
+ case "response.output_item.added":
+ if e.Item != nil && e.Item.Type == "function_call" {
+ sawAdded = true
+ }
+ case "response.function_call_arguments.done":
+ sawArgsDone = true
+ require.Equal(t, `{"cmd":"ls"}`, e.Arguments)
+ case "response.output_item.done":
+ if e.Item != nil && e.Item.Type == "function_call" {
+ sawItemDone = true
+ require.Equal(t, `{"cmd":"ls"}`, e.Item.Arguments)
+ require.Equal(t, "call_a", e.Item.CallID)
+ }
+ }
+ }
+ require.True(t, sawAdded, "function_call output_item.added missing")
+ require.True(t, sawArgsDone, "function_call_arguments.done missing")
+ require.True(t, sawItemDone, "function_call output_item.done missing")
+}
+
+// TestStream_SSEWireComplete drives the full stream through SSE encoding and
+// asserts the function_call events carry complete fields on the wire.
+func TestStream_SSEWireComplete(t *testing.T) {
+ events := collectStreamEvents(t, []string{
+ `{"choices":[{"index":0,"delta":{"role":"assistant","reasoning_content":"plan"}}]}`,
+ `{"choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"id":"call_a","type":"function","function":{"name":"exec","arguments":"{}"}}]}}]}`,
+ `{"choices":[{"index":0,"delta":{},"finish_reason":"tool_calls"}]}`,
+ })
+
+ var addedLine string
+ for _, e := range events {
+ sse, err := ResponsesEventToSSE(e)
+ require.NoError(t, err)
+ if e.Type == "response.output_item.added" && e.Item != nil && e.Item.Type == "function_call" {
+ addedLine = sse
+ }
+ }
+ require.NotEmpty(t, addedLine)
+ // The function_call added event must carry arguments:"" on the wire.
+ require.True(t, strings.Contains(addedLine, `"arguments":""`), "added line missing arguments: %s", addedLine)
+ require.Contains(t, addedLine, `"call_id":"call_a"`)
+}
diff --git a/backend/internal/pkg/apicompat/responses_stream_event_wire.go b/backend/internal/pkg/apicompat/responses_stream_event_wire.go
new file mode 100644
index 00000000..df7a82e3
--- /dev/null
+++ b/backend/internal/pkg/apicompat/responses_stream_event_wire.go
@@ -0,0 +1,199 @@
+package apicompat
+
+import "encoding/json"
+
+// MarshalJSON renders a ResponsesStreamEvent into its wire form.
+//
+// The OpenAI Responses streaming protocol requires several fields to be present
+// even when they hold a zero value: output_index/content_index/summary_index are
+// meaningful at 0, a function_call item must always carry call_id/name/arguments
+// (arguments may be ""), a message item must carry content:[] and an output_text
+// part must carry text/annotations/logprobs. Go's `omitempty` drops exactly those
+// zero values, and strict clients (Codex CLI) reject items/deltas whose required
+// fields are missing.
+//
+// Rather than marshalling with omitempty and patching the JSON afterwards, every
+// streamed event type is constructed explicitly here — the Go analogue of the
+// reference gateways' (cc-switch, CCX) per-event object construction. This is the
+// single source of truth for Responses SSE field presence and applies uniformly
+// to every emitter (Chat→Responses bridge and Anthropic→Responses converter).
+//
+// Event types not listed fall back to the default struct marshalling, which
+// bounds the blast radius of this method to the streamed item/part/text/tool
+// events.
+func (e ResponsesStreamEvent) MarshalJSON() ([]byte, error) {
+ switch e.Type {
+ case "response.output_text.delta", "response.output_text.done":
+ m := e.wireBase()
+ e.putItemID(m)
+ m["output_index"] = e.OutputIndex
+ m["content_index"] = e.ContentIndex
+ if e.Type == "response.output_text.done" {
+ m["text"] = e.Text
+ } else {
+ m["delta"] = e.Delta
+ }
+ return json.Marshal(m)
+
+ case "response.content_part.added", "response.content_part.done":
+ m := e.wireBase()
+ e.putItemID(m)
+ m["output_index"] = e.OutputIndex
+ m["content_index"] = e.ContentIndex
+ m["part"] = outputTextPartWire(e.Part)
+ return json.Marshal(m)
+
+ case "response.reasoning_summary_text.delta", "response.reasoning_summary_text.done":
+ m := e.wireBase()
+ e.putItemID(m)
+ m["output_index"] = e.OutputIndex
+ m["summary_index"] = e.SummaryIndex
+ if e.Type == "response.reasoning_summary_text.done" {
+ m["text"] = e.Text
+ } else {
+ m["delta"] = e.Delta
+ }
+ return json.Marshal(m)
+
+ case "response.reasoning_summary_part.added", "response.reasoning_summary_part.done":
+ m := e.wireBase()
+ e.putItemID(m)
+ m["output_index"] = e.OutputIndex
+ m["summary_index"] = e.SummaryIndex
+ m["part"] = summaryTextPartWire(e.Part)
+ return json.Marshal(m)
+
+ case "response.output_item.added", "response.output_item.done":
+ m := e.wireBase()
+ m["output_index"] = e.OutputIndex
+ m["item"] = responsesItemWire(e.Item)
+ return json.Marshal(m)
+
+ case "response.function_call_arguments.delta", "response.function_call_arguments.done":
+ m := e.wireBase()
+ e.putItemID(m)
+ m["output_index"] = e.OutputIndex
+ if e.CallID != "" {
+ m["call_id"] = e.CallID
+ }
+ if e.Name != "" {
+ m["name"] = e.Name
+ }
+ if e.Type == "response.function_call_arguments.done" {
+ m["arguments"] = e.Arguments
+ } else {
+ m["delta"] = e.Delta
+ }
+ return json.Marshal(m)
+
+ default:
+ // response.created / completed / done / failed / incomplete and any
+ // event type not shaped above keep the default struct marshalling.
+ type alias ResponsesStreamEvent
+ return json.Marshal(alias(e))
+ }
+}
+
+func (e ResponsesStreamEvent) wireBase() map[string]any {
+ m := map[string]any{
+ "type": e.Type,
+ "sequence_number": e.SequenceNumber,
+ }
+ return m
+}
+
+func (e ResponsesStreamEvent) putItemID(m map[string]any) {
+ if e.ItemID != "" {
+ m["item_id"] = e.ItemID
+ }
+}
+
+// outputTextPartWire renders a content part for a message's output_text, always
+// carrying text/annotations/logprobs (matching cc-switch's push_text_delta).
+func outputTextPartWire(part *ResponsesContentPart) map[string]any {
+ text := ""
+ if part != nil {
+ text = part.Text
+ }
+ return map[string]any{
+ "type": "output_text",
+ "text": text,
+ "annotations": []any{},
+ "logprobs": []any{},
+ }
+}
+
+// summaryTextPartWire renders a reasoning summary part.
+func summaryTextPartWire(part *ResponsesContentPart) map[string]any {
+ text := ""
+ if part != nil {
+ text = part.Text
+ }
+ return map[string]any{
+ "type": "summary_text",
+ "text": text,
+ }
+}
+
+// responsesItemWire renders an output_item with every field the item's type
+// requires to be present, including the empty arrays/strings that omitempty
+// would otherwise drop. Mirrors cc-switch's response_function_call_item and the
+// message/reasoning item shapes codex expects.
+func responsesItemWire(item *ResponsesOutput) map[string]any {
+ if item == nil {
+ return map[string]any{}
+ }
+ m := map[string]any{
+ "type": item.Type,
+ "id": item.ID,
+ }
+ if item.Status != "" {
+ m["status"] = item.Status
+ }
+ switch item.Type {
+ case "message":
+ role := item.Role
+ if role == "" {
+ role = "assistant"
+ }
+ m["role"] = role
+ m["content"] = messageContentWire(item.Content)
+ case "reasoning":
+ m["summary"] = reasoningSummaryWire(item.Summary)
+ if item.EncryptedContent != "" {
+ m["encrypted_content"] = item.EncryptedContent
+ }
+ case "function_call":
+ m["call_id"] = item.CallID
+ m["name"] = item.Name
+ m["arguments"] = item.Arguments
+ }
+ return m
+}
+
+// messageContentWire renders a message item's content array; always an array
+// (never null), with each output_text part carrying its text.
+func messageContentWire(parts []ResponsesContentPart) []map[string]any {
+ out := make([]map[string]any, 0, len(parts))
+ for _, p := range parts {
+ typ := p.Type
+ if typ == "" {
+ typ = "output_text"
+ }
+ out = append(out, map[string]any{"type": typ, "text": p.Text})
+ }
+ return out
+}
+
+// reasoningSummaryWire renders a reasoning item's summary array; always an array.
+func reasoningSummaryWire(summary []ResponsesSummary) []map[string]any {
+ out := make([]map[string]any, 0, len(summary))
+ for _, s := range summary {
+ typ := s.Type
+ if typ == "" {
+ typ = "summary_text"
+ }
+ out = append(out, map[string]any{"type": typ, "text": s.Text})
+ }
+ return out
+}
diff --git a/backend/internal/pkg/apicompat/responses_stream_event_wire_test.go b/backend/internal/pkg/apicompat/responses_stream_event_wire_test.go
new file mode 100644
index 00000000..fbef45af
--- /dev/null
+++ b/backend/internal/pkg/apicompat/responses_stream_event_wire_test.go
@@ -0,0 +1,109 @@
+package apicompat
+
+import (
+ "encoding/json"
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+// marshalEvent marshals through the custom MarshalJSON and returns the decoded
+// object plus the set of top-level keys.
+func marshalEvent(t *testing.T, e ResponsesStreamEvent) map[string]any {
+ t.Helper()
+ b, err := json.Marshal(e)
+ require.NoError(t, err)
+ var m map[string]any
+ require.NoError(t, json.Unmarshal(b, &m))
+ return m
+}
+
+// TestWire_IndexFieldsPresentAtZero guards the omitempty trap: output_index/
+// content_index/summary_index must serialize even when 0.
+func TestWire_IndexFieldsPresentAtZero(t *testing.T) {
+ m := marshalEvent(t, ResponsesStreamEvent{
+ Type: "response.output_text.delta", OutputIndex: 0, ContentIndex: 0, ItemID: "msg_1", Delta: "hi",
+ })
+ require.Contains(t, m, "output_index")
+ require.Contains(t, m, "content_index")
+ require.EqualValues(t, 0, m["output_index"])
+
+ r := marshalEvent(t, ResponsesStreamEvent{
+ Type: "response.reasoning_summary_text.delta", OutputIndex: 0, SummaryIndex: 0, ItemID: "rs_1", Delta: "think",
+ })
+ require.Contains(t, r, "output_index")
+ require.Contains(t, r, "summary_index")
+}
+
+// TestWire_FunctionCallItemAlwaysComplete guards that a function_call item
+// always carries call_id/name/arguments, including arguments:"" on .added.
+func TestWire_FunctionCallItemAlwaysComplete(t *testing.T) {
+ added := marshalEvent(t, ResponsesStreamEvent{
+ Type: "response.output_item.added",
+ OutputIndex: 1,
+ Item: &ResponsesOutput{Type: "function_call", ID: "fc_1", CallID: "call_a", Name: "exec", Status: "in_progress"},
+ })
+ item := added["item"].(map[string]any)
+ for _, k := range []string{"call_id", "name", "arguments"} {
+ require.Containsf(t, item, k, "function_call item missing %q", k)
+ }
+ require.Equal(t, "", item["arguments"])
+}
+
+// TestWire_MessageItemContentAlwaysArray guards content:[] presence.
+func TestWire_MessageItemContentAlwaysArray(t *testing.T) {
+ m := marshalEvent(t, ResponsesStreamEvent{
+ Type: "response.output_item.added",
+ OutputIndex: 0,
+ Item: &ResponsesOutput{Type: "message", ID: "msg_1", Role: "assistant", Status: "in_progress"},
+ })
+ item := m["item"].(map[string]any)
+ require.Contains(t, item, "content")
+ _, ok := item["content"].([]any)
+ require.True(t, ok, "content must be an array")
+}
+
+// TestWire_ReasoningItemSummaryAlwaysArray guards summary:[] presence.
+func TestWire_ReasoningItemSummaryAlwaysArray(t *testing.T) {
+ m := marshalEvent(t, ResponsesStreamEvent{
+ Type: "response.output_item.added",
+ OutputIndex: 0,
+ Item: &ResponsesOutput{Type: "reasoning", ID: "rs_1", Status: "in_progress"},
+ })
+ item := m["item"].(map[string]any)
+ require.Contains(t, item, "summary")
+ _, ok := item["summary"].([]any)
+ require.True(t, ok, "summary must be an array")
+}
+
+// TestWire_ContentPartCarriesAnnotationsLogprobs guards the output_text part shape.
+func TestWire_ContentPartCarriesAnnotationsLogprobs(t *testing.T) {
+ m := marshalEvent(t, ResponsesStreamEvent{
+ Type: "response.content_part.added", OutputIndex: 0, ContentIndex: 0, ItemID: "msg_1",
+ Part: &ResponsesContentPart{Type: "output_text", Text: ""},
+ })
+ part := m["part"].(map[string]any)
+ require.Equal(t, "output_text", part["type"])
+ require.Contains(t, part, "text")
+ require.Contains(t, part, "annotations")
+ require.Contains(t, part, "logprobs")
+}
+
+// TestWire_ArgumentsDonePresentEvenEmpty guards arguments presence on done.
+func TestWire_ArgumentsDonePresentEvenEmpty(t *testing.T) {
+ m := marshalEvent(t, ResponsesStreamEvent{
+ Type: "response.function_call_arguments.done", OutputIndex: 1, ItemID: "fc_1", CallID: "call_a", Name: "exec", Arguments: "",
+ })
+ require.Contains(t, m, "arguments")
+ require.Equal(t, "", m["arguments"])
+}
+
+// TestWire_UnknownEventFallsBackToDefault ensures non-streamed event types keep
+// default marshalling (the response object is preserved).
+func TestWire_UnknownEventFallsBackToDefault(t *testing.T) {
+ m := marshalEvent(t, ResponsesStreamEvent{
+ Type: "response.completed",
+ Response: &ResponsesResponse{ID: "resp_1", Object: "response", Status: "completed"},
+ })
+ require.Contains(t, m, "response")
+}
diff --git a/backend/internal/pkg/apicompat/types.go b/backend/internal/pkg/apicompat/types.go
index b4451f23..d2937802 100644
--- a/backend/internal/pkg/apicompat/types.go
+++ b/backend/internal/pkg/apicompat/types.go
@@ -406,6 +406,10 @@ type ResponsesStreamEvent struct {
// Reuses Text/Delta fields above, SummaryIndex identifies which summary part
SummaryIndex int `json:"summary_index,omitempty"`
+ // response.content_part.added / done and
+ // response.reasoning_summary_part.added / done
+ Part *ResponsesContentPart `json:"part,omitempty"`
+
// error event fields
Code string `json:"code,omitempty"`
Param string `json:"param,omitempty"`
From 1afae0a4cded8d4ca81bc7ea70bf55bbbaa5c72c Mon Sep 17 00:00:00 2001
From: name <136912576+is7Qin@users.noreply.github.com>
Date: Sun, 31 May 2026 16:38:42 +0800
Subject: [PATCH 61/79] fix(gateway): address OpenAI OOM review lint
---
.../service/gateway_service_benchmark_test.go | 94 +++++++++----------
.../service/openai_gateway_service.go | 19 +++-
.../openai_gateway_service_hotpath_test.go | 45 +++++++++
.../service/openai_tool_continuation.go | 2 +-
4 files changed, 109 insertions(+), 51 deletions(-)
diff --git a/backend/internal/service/gateway_service_benchmark_test.go b/backend/internal/service/gateway_service_benchmark_test.go
index f6cd1404..42a711db 100644
--- a/backend/internal/service/gateway_service_benchmark_test.go
+++ b/backend/internal/service/gateway_service_benchmark_test.go
@@ -229,16 +229,16 @@ func benchmarkBodySizes() []struct {
func buildSystemCacheableRequest(parts int) *ParsedRequest {
var builder strings.Builder
- builder.WriteString(`{"system":[`)
+ _, _ = builder.WriteString(`{"system":[`)
for i := 0; i < parts; i++ {
if i > 0 {
- builder.WriteByte(',')
+ _ = builder.WriteByte(',')
}
- builder.WriteString(`{"text":"system_part_`)
- builder.WriteString(strconv.Itoa(i))
- builder.WriteString(`","cache_control":{"type":"ephemeral"}}`)
+ _, _ = builder.WriteString(`{"text":"system_part_`)
+ _, _ = builder.WriteString(strconv.Itoa(i))
+ _, _ = builder.WriteString(`","cache_control":{"type":"ephemeral"}}`)
}
- builder.WriteString(`]}`)
+ _, _ = builder.WriteString(`]}`)
parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(builder.String())), "")
if err != nil {
panic(err)
@@ -249,93 +249,93 @@ func buildSystemCacheableRequest(parts int) *ParsedRequest {
func buildLargeAnthropicMessagesBody(targetBytes int, includeCacheControl bool) []byte {
var builder strings.Builder
builder.Grow(targetBytes + 1024)
- builder.WriteString(`{"model":"claude-sonnet-4-5","stream":true,"system":[{"type":"text","text":"system seed"}],"messages":[`)
+ _, _ = builder.WriteString(`{"model":"claude-sonnet-4-5","stream":true,"system":[{"type":"text","text":"system seed"}],"messages":[`)
for i := 0; builder.Len() < targetBytes; i++ {
if i > 0 {
- builder.WriteByte(',')
+ _ = builder.WriteByte(',')
}
- builder.WriteString(`{"role":"user","content":[{"type":"text","text":"`)
- builder.WriteString(strings.Repeat("anthropic payload ", 64))
- builder.WriteString(strconv.Itoa(i))
- builder.WriteByte('"')
+ _, _ = builder.WriteString(`{"role":"user","content":[{"type":"text","text":"`)
+ _, _ = builder.WriteString(strings.Repeat("anthropic payload ", 64))
+ _, _ = builder.WriteString(strconv.Itoa(i))
+ _ = builder.WriteByte('"')
if includeCacheControl && i%32 == 0 {
- builder.WriteString(`,"cache_control":{"type":"ephemeral"}`)
+ _, _ = builder.WriteString(`,"cache_control":{"type":"ephemeral"}`)
}
- builder.WriteString(`}]}`)
+ _, _ = builder.WriteString(`}]}`)
}
- builder.WriteString(`]}`)
+ _, _ = builder.WriteString(`]}`)
return []byte(builder.String())
}
func buildLargeGeminiContentsBody(targetBytes int) []byte {
var builder strings.Builder
builder.Grow(targetBytes + 1024)
- builder.WriteString(`{"model":"gemini-2.5-pro","systemInstruction":{"parts":[{"text":"system seed"}]},"contents":[`)
+ _, _ = builder.WriteString(`{"model":"gemini-2.5-pro","systemInstruction":{"parts":[{"text":"system seed"}]},"contents":[`)
for i := 0; builder.Len() < targetBytes; i++ {
if i > 0 {
- builder.WriteByte(',')
+ _ = builder.WriteByte(',')
}
- builder.WriteString(`{"role":"user","parts":[{"text":"`)
- builder.WriteString(strings.Repeat("gemini payload ", 64))
- builder.WriteString(strconv.Itoa(i))
- builder.WriteString(`"}]}`)
+ _, _ = builder.WriteString(`{"role":"user","parts":[{"text":"`)
+ _, _ = builder.WriteString(strings.Repeat("gemini payload ", 64))
+ _, _ = builder.WriteString(strconv.Itoa(i))
+ _, _ = builder.WriteString(`"}]}`)
}
- builder.WriteString(`]}`)
+ _, _ = builder.WriteString(`]}`)
return []byte(builder.String())
}
func buildLargeOpenAIResponsesBody(targetBytes int) []byte {
var builder strings.Builder
builder.Grow(targetBytes + 1024)
- builder.WriteString(`{"model":"gpt-5.4","stream":true,"prompt_cache_key":"session-benchmark","input":[`)
+ _, _ = builder.WriteString(`{"model":"gpt-5.4","stream":true,"prompt_cache_key":"session-benchmark","input":[`)
for i := 0; builder.Len() < targetBytes; i++ {
if i > 0 {
- builder.WriteByte(',')
+ _ = builder.WriteByte(',')
}
- builder.WriteString(`{"type":"message","role":"user","content":[{"type":"input_text","text":"`)
- builder.WriteString(strings.Repeat("openai responses payload ", 48))
- builder.WriteString(strconv.Itoa(i))
- builder.WriteString(`"}]}`)
+ _, _ = builder.WriteString(`{"type":"message","role":"user","content":[{"type":"input_text","text":"`)
+ _, _ = builder.WriteString(strings.Repeat("openai responses payload ", 48))
+ _, _ = builder.WriteString(strconv.Itoa(i))
+ _, _ = builder.WriteString(`"}]}`)
}
- builder.WriteString(`],"tools":[{"type":"function","name":"lookup","parameters":{"type":"object","properties":{"query":{"type":"string"}}}}]}`)
+ _, _ = builder.WriteString(`],"tools":[{"type":"function","name":"lookup","parameters":{"type":"object","properties":{"query":{"type":"string"}}}}]}`)
return []byte(builder.String())
}
func buildLargeOpenAIResponsesToolContinuationBody(targetBytes int) []byte {
var builder strings.Builder
builder.Grow(targetBytes + 1024)
- builder.WriteString(`{"model":"gpt-5.4","stream":true,"previous_response_id":"resp_benchmark","input":[`)
+ _, _ = builder.WriteString(`{"model":"gpt-5.4","stream":true,"previous_response_id":"resp_benchmark","input":[`)
for i := 0; builder.Len() < targetBytes; i++ {
if i > 0 {
- builder.WriteByte(',')
+ _ = builder.WriteByte(',')
}
callID := "call_" + strconv.Itoa(i)
- builder.WriteString(`{"type":"item_reference","id":"`)
- builder.WriteString(callID)
- builder.WriteString(`"},{"type":"function_call_output","call_id":"`)
- builder.WriteString(callID)
- builder.WriteString(`","output":"`)
- builder.WriteString(strings.Repeat("tool output payload ", 48))
- builder.WriteString(strconv.Itoa(i))
- builder.WriteString(`"}`)
+ _, _ = builder.WriteString(`{"type":"item_reference","id":"`)
+ _, _ = builder.WriteString(callID)
+ _, _ = builder.WriteString(`"},{"type":"function_call_output","call_id":"`)
+ _, _ = builder.WriteString(callID)
+ _, _ = builder.WriteString(`","output":"`)
+ _, _ = builder.WriteString(strings.Repeat("tool output payload ", 48))
+ _, _ = builder.WriteString(strconv.Itoa(i))
+ _, _ = builder.WriteString(`"}`)
}
- builder.WriteString(`]}`)
+ _, _ = builder.WriteString(`]}`)
return []byte(builder.String())
}
func buildLargeOpenAIResponsesImageToolBody(targetBytes int) []byte {
var builder strings.Builder
builder.Grow(targetBytes + 1024)
- builder.WriteString(`{"model":"gpt-5.4","stream":false,"tools":[{"type":"image_generation","model":"gpt-image-2","size":"2048x1152"}],"input":[`)
+ _, _ = builder.WriteString(`{"model":"gpt-5.4","stream":false,"tools":[{"type":"image_generation","model":"gpt-image-2","size":"2048x1152"}],"input":[`)
for i := 0; builder.Len() < targetBytes; i++ {
if i > 0 {
- builder.WriteByte(',')
+ _ = builder.WriteByte(',')
}
- builder.WriteString(`{"type":"message","role":"user","content":[{"type":"input_text","text":"`)
- builder.WriteString(strings.Repeat("openai image billing payload ", 48))
- builder.WriteString(strconv.Itoa(i))
- builder.WriteString(`"}]}`)
+ _, _ = builder.WriteString(`{"type":"message","role":"user","content":[{"type":"input_text","text":"`)
+ _, _ = builder.WriteString(strings.Repeat("openai image billing payload ", 48))
+ _, _ = builder.WriteString(strconv.Itoa(i))
+ _, _ = builder.WriteString(`"}]}`)
}
- builder.WriteString(`]}`)
+ _, _ = builder.WriteString(`]}`)
return []byte(builder.String())
}
diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go
index d17b85d0..8fd4bc60 100644
--- a/backend/internal/service/openai_gateway_service.go
+++ b/backend/internal/service/openai_gateway_service.go
@@ -2544,6 +2544,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Added Codex image_generation bridge instructions")
}
} else if imageGenerationAllowed && imageIntent && openAIRequestBodyHasImageGenerationTool(body) {
+ // 完整 image_generation tool 只做 raw 计费读取,校验/桥接/旧字段迁移命中时才展开大 input map。
logger.LegacyPrintf("service.openai_gateway", "[OpenAI] /responses image_generation request inbound_model=%s mapped_model=%s account_type=%s", requestView.Model, upstreamModel, account.Type)
}
@@ -6215,7 +6216,7 @@ func (v *openAIRequestView) MarkPatchSet(path string, value any) {
return
}
path = strings.TrimSpace(path)
- if path == "" {
+ if !isSimpleOpenAIRequestPatchPath(path) {
v.DisablePatches()
return
}
@@ -6227,13 +6228,25 @@ func (v *openAIRequestView) MarkPatchDelete(path string) {
return
}
path = strings.TrimSpace(path)
- if path == "" {
+ if !isSimpleOpenAIRequestPatchPath(path) {
v.DisablePatches()
return
}
v.patches = append(v.patches, openAIRequestPatch{path: path, delete: true})
}
+func isSimpleOpenAIRequestPatchPath(path string) bool {
+ if path == "" || strings.ContainsRune(path, '\\') {
+ return false
+ }
+ for _, part := range strings.Split(path, ".") {
+ if strings.TrimSpace(part) == "" {
+ return false
+ }
+ }
+ return true
+}
+
func (v *openAIRequestView) DisablePatches() {
if v == nil {
return
@@ -6772,7 +6785,7 @@ func openAIRequestBodyMayContainImageInput(body []byte) bool {
return false
}
input := gjson.GetBytes(body, "input")
- messages := gjson.GetBytes(body, "messages")
+ messages := gjson.GetBytes(body, "messages.#-1")
return openAIJSONValueMayContainImageInput(input) || openAIJSONValueMayContainImageInput(messages)
}
diff --git a/backend/internal/service/openai_gateway_service_hotpath_test.go b/backend/internal/service/openai_gateway_service_hotpath_test.go
index d240cec8..92a0d1ac 100644
--- a/backend/internal/service/openai_gateway_service_hotpath_test.go
+++ b/backend/internal/service/openai_gateway_service_hotpath_test.go
@@ -46,6 +46,15 @@ func TestOpenAIRequestView_ApplyPatches(t *testing.T) {
require.JSONEq(t, `{"model":"gpt-5.1","reasoning":{"effort":"none"},"input":[{"type":"message","content":"hi"}]}`, string(patched))
}
+func TestOpenAIRequestView_RejectsEscapedPatchPath(t *testing.T) {
+ view := newOpenAIRequestView([]byte(`{"metadata":{"user.id":"old"}}`))
+ view.MarkPatchSet(`metadata.user\.id`, "new")
+
+ require.False(t, view.HasPatches())
+ _, err := view.ApplyPatches()
+ require.Error(t, err)
+}
+
func TestOpenAIRequestView_ApplyPatchesDisabled(t *testing.T) {
view := newOpenAIRequestView([]byte(`{"model":"gpt-5"}`))
view.MarkPatchSet("model", "gpt-5.1")
@@ -260,6 +269,42 @@ func TestOpenAIGatewayService_Forward_ImageToolBillingDoesNotForceFullDecode(t *
require.Equal(t, "gpt-image-2", result.BillingModel)
}
+func TestOpenAIGatewayService_Forward_ImageToolWithImageOnlyModelIsNormalized(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)),
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 11,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"gpt-image-2","stream":false,"tools":[{"type":"image_generation","model":"gpt-image-2"}],"input":"draw"}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.Equal(t, openAIImagesResponsesMainModel, gjson.GetBytes(upstream.lastBody, "model").String())
+}
+
func TestOpenAIGatewayService_Forward_HTTPRetryRecoveryDoesNotDecodeBeforeError(t *testing.T) {
gin.SetMode(gin.TestMode)
upstream := &httpUpstreamRecorder{
diff --git a/backend/internal/service/openai_tool_continuation.go b/backend/internal/service/openai_tool_continuation.go
index 92bc2613..6515c0c4 100644
--- a/backend/internal/service/openai_tool_continuation.go
+++ b/backend/internal/service/openai_tool_continuation.go
@@ -199,7 +199,7 @@ func ValidateFunctionCallOutputContextBytes(body []byte) FunctionCallOutputValid
}
referenceIDs[idValue] = struct{}{}
}
- return !(result.HasFunctionCallOutput && result.HasToolCallContext)
+ return !result.HasFunctionCallOutput || !result.HasToolCallContext
})
if !result.HasFunctionCallOutput || result.HasToolCallContext || len(callIDs) == 0 || len(referenceIDs) == 0 {
return result
From a01686c637c81c29d19d7139f154c7ca9d4a1add Mon Sep 17 00:00:00 2001
From: moonagic
Date: Sun, 31 May 2026 21:44:40 +0800
Subject: [PATCH 62/79] fix antigravity gemini rate limit and account
scheduling
Squash of 4 commits:
- Fix Gemini rate limit scheduling
- fix antigravity gemini rate limit scheduling
- fix antigravity gemini limited account scheduling
- fix antigravity test stubs for default lint
---
backend/internal/handler/gateway_handler.go | 12 +-
.../internal/handler/gemini_v1beta_handler.go | 13 +-
.../internal/repository/scheduler_cache.go | 1 +
.../repository/scheduler_cache_unit_test.go | 26 ++++
.../antigravity_default_test_stubs_test.go | 61 +++++++++
.../service/antigravity_gateway_service.go | 126 +++++++++++++-----
.../antigravity_gateway_service_test.go | 80 +++++++++++
.../service/antigravity_quota_scope.go | 12 +-
.../service/antigravity_rate_limit_test.go | 95 +++++++++++++
.../antigravity_single_account_retry_test.go | 8 +-
.../service/antigravity_smart_retry_test.go | 28 ++--
backend/internal/service/error_policy_test.go | 54 ++++++++
.../service/gateway_multiplatform_test.go | 100 ++++++++++++++
backend/internal/service/model_rate_limit.go | 32 ++++-
.../internal/service/model_rate_limit_test.go | 62 +++++++++
.../scheduler_snapshot_hydration_test.go | 88 ++++++++++++
16 files changed, 747 insertions(+), 51 deletions(-)
create mode 100644 backend/internal/service/antigravity_default_test_stubs_test.go
diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go
index eb5c4a42..50c9fe34 100644
--- a/backend/internal/handler/gateway_handler.go
+++ b/backend/internal/handler/gateway_handler.go
@@ -440,7 +440,17 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
// 记录 Forward 前已写入字节数,Forward 后若增加则说明 SSE 内容已发,禁止 failover
writerSizeBeforeForward := c.Writer.Size()
if account.Platform == service.PlatformAntigravity {
- result, err = h.antigravityGatewayService.ForwardGemini(requestCtx, c, account, reqModel, "generateContent", reqStream, body, hasBoundSession)
+ result, err = h.antigravityGatewayService.ForwardGemini(
+ requestCtx,
+ c,
+ account,
+ reqModel,
+ "generateContent",
+ reqStream,
+ body,
+ hasBoundSession,
+ service.WithForwardGeminiSession(derefGroupID(apiKey.GroupID), sessionKey),
+ )
} else {
result, err = h.geminiCompatService.Forward(requestCtx, c, account, body)
}
diff --git a/backend/internal/handler/gemini_v1beta_handler.go b/backend/internal/handler/gemini_v1beta_handler.go
index 0b33ca3e..dfe0b2a8 100644
--- a/backend/internal/handler/gemini_v1beta_handler.go
+++ b/backend/internal/handler/gemini_v1beta_handler.go
@@ -477,8 +477,19 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
if fs.SwitchCount > 0 {
requestCtx = service.WithAccountSwitchCount(requestCtx, fs.SwitchCount, h.metadataBridgeEnabled())
}
+ sessionGroupID := derefGroupID(apiKey.GroupID)
if account.Platform == service.PlatformAntigravity && account.Type != service.AccountTypeAPIKey {
- result, err = h.antigravityGatewayService.ForwardGemini(requestCtx, c, account, modelName, action, stream, body, hasBoundSession)
+ result, err = h.antigravityGatewayService.ForwardGemini(
+ requestCtx,
+ c,
+ account,
+ modelName,
+ action,
+ stream,
+ body,
+ hasBoundSession,
+ service.WithForwardGeminiSession(sessionGroupID, sessionKey),
+ )
} else {
result, err = h.geminiCompatService.ForwardNative(requestCtx, c, account, modelName, action, stream, body)
}
diff --git a/backend/internal/repository/scheduler_cache.go b/backend/internal/repository/scheduler_cache.go
index cf19deda..921aa081 100644
--- a/backend/internal/repository/scheduler_cache.go
+++ b/backend/internal/repository/scheduler_cache.go
@@ -559,6 +559,7 @@ func filterSchedulerExtra(extra map[string]any) map[string]any {
"auto_pause_7d_threshold",
"auto_pause_5h_disabled",
"auto_pause_7d_disabled",
+ "model_rate_limits",
}
filtered := make(map[string]any)
for _, key := range keys {
diff --git a/backend/internal/repository/scheduler_cache_unit_test.go b/backend/internal/repository/scheduler_cache_unit_test.go
index a4667591..c14721cd 100644
--- a/backend/internal/repository/scheduler_cache_unit_test.go
+++ b/backend/internal/repository/scheduler_cache_unit_test.go
@@ -108,3 +108,29 @@ func TestBuildSchedulerMetadataAccount_KeepsQuotaAutoPauseFields(t *testing.T) {
require.Equal(t, true, got.Extra["auto_pause_5h_disabled"])
require.Equal(t, false, got.Extra["auto_pause_7d_disabled"])
}
+
+func TestBuildSchedulerMetadataAccount_KeepsModelRateLimits(t *testing.T) {
+ account := service.Account{
+ ID: 90,
+ Platform: service.PlatformAntigravity,
+ Extra: map[string]any{
+ "model_rate_limits": map[string]any{
+ "gemini-3-flash": map[string]any{
+ "rate_limit_reset_at": "2026-05-30T10:10:00Z",
+ },
+ "antigravity:gemini": map[string]any{
+ "rate_limit_reset_at": "2026-05-30T10:10:00Z",
+ },
+ },
+ "unused_large_field": "drop-me",
+ },
+ }
+
+ got := buildSchedulerMetadataAccount(account)
+
+ limits, ok := got.Extra["model_rate_limits"].(map[string]any)
+ require.True(t, ok)
+ require.Contains(t, limits, "gemini-3-flash")
+ require.Contains(t, limits, "antigravity:gemini")
+ require.Nil(t, got.Extra["unused_large_field"])
+}
diff --git a/backend/internal/service/antigravity_default_test_stubs_test.go b/backend/internal/service/antigravity_default_test_stubs_test.go
new file mode 100644
index 00000000..d3c2c57a
--- /dev/null
+++ b/backend/internal/service/antigravity_default_test_stubs_test.go
@@ -0,0 +1,61 @@
+//go:build !unit
+
+package service
+
+import (
+ "context"
+ "time"
+)
+
+type defaultRateLimitCall struct {
+ accountID int64
+ resetAt time.Time
+}
+
+type defaultModelRateLimitCall struct {
+ accountID int64
+ modelKey string
+ resetAt time.Time
+}
+
+type defaultExtraUpdateCall struct {
+ accountID int64
+ updates map[string]any
+}
+
+type stubAntigravityAccountRepo struct {
+ AccountRepository
+ rateCalls []defaultRateLimitCall
+ modelRateLimitCalls []defaultModelRateLimitCall
+ extraUpdateCalls []defaultExtraUpdateCall
+}
+
+func (s *stubAntigravityAccountRepo) SetRateLimited(_ context.Context, id int64, resetAt time.Time) error {
+ s.rateCalls = append(s.rateCalls, defaultRateLimitCall{accountID: id, resetAt: resetAt})
+ return nil
+}
+
+func (s *stubAntigravityAccountRepo) SetModelRateLimit(_ context.Context, id int64, modelKey string, resetAt time.Time, _ ...string) error {
+ s.modelRateLimitCalls = append(s.modelRateLimitCalls, defaultModelRateLimitCall{accountID: id, modelKey: modelKey, resetAt: resetAt})
+ return nil
+}
+
+func (s *stubAntigravityAccountRepo) UpdateExtra(_ context.Context, id int64, updates map[string]any) error {
+ s.extraUpdateCalls = append(s.extraUpdateCalls, defaultExtraUpdateCall{accountID: id, updates: updates})
+ return nil
+}
+
+type defaultDeleteSessionCall struct {
+ groupID int64
+ sessionHash string
+}
+
+type stubSmartRetryCache struct {
+ GatewayCache
+ deleteCalls []defaultDeleteSessionCall
+}
+
+func (c *stubSmartRetryCache) DeleteSessionAccountID(_ context.Context, groupID int64, sessionHash string) error {
+ c.deleteCalls = append(c.deleteCalls, defaultDeleteSessionCall{groupID: groupID, sessionHash: sessionHash})
+ return nil
+}
diff --git a/backend/internal/service/antigravity_gateway_service.go b/backend/internal/service/antigravity_gateway_service.go
index 2b849bdd..8aca2e13 100644
--- a/backend/internal/service/antigravity_gateway_service.go
+++ b/backend/internal/service/antigravity_gateway_service.go
@@ -228,12 +228,11 @@ func (s *AntigravityGatewayService) handleSmartRetry(p antigravityRetryLoopParam
p.prefix, resp.StatusCode, modelName, p.account.ID, rateLimitDuration, truncateForLog(respBody, 200))
resetAt := time.Now().Add(rateLimitDuration)
- if !setModelRateLimitByModelName(p.ctx, p.accountRepo, p.account.ID, modelName, p.prefix, resp.StatusCode, resetAt, false) {
+ if !s.setAntigravityModelRateLimits(p.ctx, p.accountRepo, p.account, modelName, p.prefix, resp.StatusCode, resetAt, false) {
p.handleError(p.ctx, p.prefix, p.account, resp.StatusCode, resp.Header, respBody, p.requestedModel, p.groupID, p.sessionHash, p.isStickySession)
logger.LegacyPrintf("service.antigravity_gateway", "%s status=%d rate_limited account=%d (no model mapping)", p.prefix, resp.StatusCode, p.account.ID)
- } else {
- s.updateAccountModelRateLimitInCache(p.ctx, p.account, modelName, resetAt)
}
+ s.clearStickySession(p.ctx, p.groupID, p.sessionHash)
// 返回账号切换信号,让上层切换账号重试
return &smartRetryResult{
@@ -392,20 +391,10 @@ func (s *AntigravityGatewayService) handleSmartRetry(p antigravityRetryLoopParam
p.prefix, resp.StatusCode, maxAttempts, modelName, p.account.ID, rateLimitDuration, truncateForLog(retryBody, 200))
resetAt := time.Now().Add(rateLimitDuration)
- if p.accountRepo != nil && modelName != "" {
- if err := p.accountRepo.SetModelRateLimit(p.ctx, p.account.ID, modelName, resetAt); err != nil {
- logger.LegacyPrintf("service.antigravity_gateway", "%s status=%d model_rate_limit_failed model=%s error=%v", p.prefix, resp.StatusCode, modelName, err)
- } else {
- logger.LegacyPrintf("service.antigravity_gateway", "%s status=%d model_rate_limited_after_smart_retry model=%s account=%d reset_in=%v",
- p.prefix, resp.StatusCode, modelName, p.account.ID, rateLimitDuration)
- s.updateAccountModelRateLimitInCache(p.ctx, p.account, modelName, resetAt)
- }
- }
+ s.setAntigravityModelRateLimits(p.ctx, p.accountRepo, p.account, modelName, p.prefix, resp.StatusCode, resetAt, true)
// 清除粘性会话绑定,避免下次请求仍命中限流账号
- if s.cache != nil && p.sessionHash != "" {
- _ = s.cache.DeleteSessionAccountID(p.ctx, p.groupID, p.sessionHash)
- }
+ s.clearStickySession(p.ctx, p.groupID, p.sessionHash)
// 返回账号切换信号,让上层切换账号重试
return &smartRetryResult{
@@ -938,8 +927,14 @@ func (s *AntigravityGatewayService) checkErrorPolicy(ctx context.Context, accoun
func (s *AntigravityGatewayService) applyErrorPolicy(p antigravityRetryLoopParams, statusCode int, headers http.Header, respBody []byte) (handled bool, outStatus int, retErr error) {
switch s.checkErrorPolicy(p.ctx, p.account, statusCode, respBody) {
case ErrorPolicySkipped:
+ if s.handleAntigravityModelRateLimitBeforePolicy(p, statusCode, headers, respBody) {
+ return true, statusCode, nil
+ }
return true, http.StatusInternalServerError, nil
case ErrorPolicyMatched:
+ if s.handleAntigravityModelRateLimitBeforePolicy(p, statusCode, headers, respBody) {
+ return true, statusCode, nil
+ }
_ = p.handleError(p.ctx, p.prefix, p.account, statusCode, headers, respBody,
p.requestedModel, p.groupID, p.sessionHash, p.isStickySession)
return true, statusCode, nil
@@ -951,6 +946,31 @@ func (s *AntigravityGatewayService) applyErrorPolicy(p antigravityRetryLoopParam
return false, statusCode, nil
}
+func (s *AntigravityGatewayService) handleAntigravityModelRateLimitBeforePolicy(p antigravityRetryLoopParams, statusCode int, headers http.Header, respBody []byte) bool {
+ if statusCode != http.StatusTooManyRequests && statusCode != http.StatusServiceUnavailable {
+ return false
+ }
+ if p.account == nil || p.account.Platform != PlatformAntigravity {
+ return false
+ }
+ _, shouldRateLimitModel, waitDuration, modelName, isModelCapacityExhausted := shouldTriggerAntigravitySmartRetry(p.account, respBody)
+ if isModelCapacityExhausted || !shouldRateLimitModel || strings.TrimSpace(modelName) == "" {
+ return false
+ }
+ rateLimitDuration := waitDuration
+ if rateLimitDuration <= 0 {
+ rateLimitDuration = antigravityDefaultRateLimitDuration
+ }
+ resetAt := time.Now().Add(rateLimitDuration)
+ if !s.setAntigravityModelRateLimits(p.ctx, p.accountRepo, p.account, modelName, p.prefix, statusCode, resetAt, false) {
+ return false
+ }
+ s.clearStickySession(p.ctx, p.groupID, p.sessionHash)
+ logger.LegacyPrintf("service.antigravity_gateway", "%s status=%d model_rate_limited_before_error_policy model=%s account=%d reset_in=%v",
+ p.prefix, statusCode, modelName, p.account.ID, rateLimitDuration)
+ return true
+}
+
// mapAntigravityModel 获取映射后的模型名
// 完全依赖映射配置:账户映射(通配符)→ 默认映射兜底(DefaultAntigravityModelMapping)
// 注意:返回空字符串表示模型不被支持,调度时会过滤掉该账号
@@ -958,6 +978,7 @@ func mapAntigravityModel(account *Account, requestedModel string) string {
if account == nil {
return ""
}
+ requestedModel = strings.TrimPrefix(requestedModel, "models/")
// 获取映射表(未配置时自动使用 DefaultAntigravityModelMapping)
mapping := account.GetModelMapping()
@@ -2057,8 +2078,28 @@ func stripSignatureSensitiveBlocksFromClaudeRequest(req *antigravity.ClaudeReque
// └─ retryDelay < 7s → 等待后重试 1 次
// ├─ 成功 → 正常返回
// └─ 失败 → 设置模型限流 + 清除粘性绑定 → 切换账号
-func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Context, account *Account, originalModel string, action string, stream bool, body []byte, isStickySession bool) (*ForwardResult, error) {
+type ForwardGeminiOption func(*forwardGeminiOptions)
+
+type forwardGeminiOptions struct {
+ groupID int64
+ sessionHash string
+}
+
+func WithForwardGeminiSession(groupID int64, sessionHash string) ForwardGeminiOption {
+ return func(opts *forwardGeminiOptions) {
+ opts.groupID = groupID
+ opts.sessionHash = sessionHash
+ }
+}
+
+func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Context, account *Account, originalModel string, action string, stream bool, body []byte, isStickySession bool, options ...ForwardGeminiOption) (*ForwardResult, error) {
startTime := time.Now()
+ forwardOpts := forwardGeminiOptions{}
+ for _, apply := range options {
+ if apply != nil {
+ apply(&forwardOpts)
+ }
+ }
sessionID := getSessionID(c)
prefix := logPrefix(sessionID, account.Name)
@@ -2163,8 +2204,8 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co
handleError: s.handleUpstreamError,
requestedModel: originalModel,
isStickySession: isStickySession, // ForwardGemini 由上层判断粘性会话
- groupID: 0, // ForwardGemini 方法没有 groupID,由上层处理粘性会话清除
- sessionHash: "", // ForwardGemini 方法没有 sessionHash,由上层处理粘性会话清除
+ groupID: forwardOpts.groupID,
+ sessionHash: forwardOpts.sessionHash,
})
if err != nil {
// 检查是否是账号切换信号,转换为 UpstreamFailoverError 让 Handler 切换账号
@@ -2262,8 +2303,8 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co
handleError: s.handleUpstreamError,
requestedModel: originalModel,
isStickySession: isStickySession,
- groupID: 0,
- sessionHash: "",
+ groupID: forwardOpts.groupID,
+ sessionHash: forwardOpts.sessionHash,
})
if retryErr == nil {
retryResp := retryResult.resp
@@ -2339,7 +2380,7 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co
if unwrapErr != nil || len(unwrappedForOps) == 0 {
unwrappedForOps = respBody
}
- s.handleUpstreamError(ctx, prefix, account, resp.StatusCode, resp.Header, respBody, originalModel, 0, "", isStickySession)
+ s.handleUpstreamError(ctx, prefix, account, resp.StatusCode, resp.Header, respBody, originalModel, forwardOpts.groupID, forwardOpts.sessionHash, isStickySession)
upstreamMsg := strings.TrimSpace(extractAntigravityErrorMessage(unwrappedForOps))
upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg)
upstreamDetail := s.getUpstreamErrorDetail(unwrappedForOps)
@@ -2550,6 +2591,34 @@ func setModelRateLimitByModelName(ctx context.Context, repo AccountRepository, a
return true
}
+func (s *AntigravityGatewayService) setAntigravityModelRateLimits(ctx context.Context, repo AccountRepository, account *Account, modelName, prefix string, statusCode int, resetAt time.Time, afterSmartRetry bool) bool {
+ if account == nil || repo == nil {
+ return false
+ }
+ keys := antigravityModelRateLimitKeys(modelName)
+ if len(keys) == 0 {
+ return false
+ }
+
+ success := false
+ for _, key := range keys {
+ if setModelRateLimitByModelName(ctx, repo, account.ID, key, prefix, statusCode, resetAt, afterSmartRetry) {
+ s.updateAccountModelRateLimitInCache(ctx, account, key, resetAt)
+ success = true
+ }
+ }
+ return success
+}
+
+func (s *AntigravityGatewayService) clearStickySession(ctx context.Context, groupID int64, sessionHash string) {
+ if s == nil || s.cache == nil || strings.TrimSpace(sessionHash) == "" {
+ return
+ }
+ if err := s.cache.DeleteSessionAccountID(ctx, groupID, sessionHash); err != nil {
+ logger.LegacyPrintf("service.antigravity_gateway", "[antigravity-Forward] sticky_session_clear_failed group_id=%d session=%s err=%v", groupID, shortSessionHash(sessionHash), err)
+ }
+}
+
func antigravityFallbackCooldownSeconds() (time.Duration, bool) {
raw := strings.TrimSpace(os.Getenv(antigravityFallbackSecondsEnv))
if raw == "" {
@@ -2628,7 +2697,7 @@ func parseAntigravitySmartRetryInfo(body []byte) *antigravitySmartRetryInfo {
if atType == googleRPCTypeErrorInfo {
if meta, ok := dm["metadata"].(map[string]any); ok {
if model, ok := meta["model"].(string); ok {
- modelName = model
+ modelName = normalizeAntigravityModelName(model)
}
}
// 检查 reason
@@ -2802,13 +2871,7 @@ func (s *AntigravityGatewayService) setModelRateLimitAndClearSession(p *handleMo
logger.LegacyPrintf("service.antigravity_gateway", "%s status=%d model_rate_limited model=%s account=%d reset_in=%v",
p.prefix, p.statusCode, info.ModelName, p.account.ID, info.RetryDelay)
- // 设置模型限流状态(数据库)
- if err := s.accountRepo.SetModelRateLimit(p.ctx, p.account.ID, info.ModelName, resetAt); err != nil {
- logger.LegacyPrintf("service.antigravity_gateway", "%s model_rate_limit_failed model=%s error=%v", p.prefix, info.ModelName, err)
- }
-
- // 立即更新 Redis 快照中账号的限流状态,避免并发请求重复选中
- s.updateAccountModelRateLimitInCache(p.ctx, p.account, info.ModelName, resetAt)
+ s.setAntigravityModelRateLimits(p.ctx, s.accountRepo, p.account, info.ModelName, p.prefix, p.statusCode, resetAt, false)
// 清除粘性会话绑定
if p.cache != nil && p.sessionHash != "" {
@@ -2898,12 +2961,11 @@ func (s *AntigravityGatewayService) handleUpstreamError(
}
if modelKey != "" {
ra := s.resolveResetTime(resetAt, defaultDur)
- if err := s.accountRepo.SetModelRateLimit(ctx, account.ID, modelKey, ra); err != nil {
- logger.LegacyPrintf("service.antigravity_gateway", "%s status=429 model_rate_limit_set_failed model=%s error=%v", prefix, modelKey, err)
+ if !s.setAntigravityModelRateLimits(ctx, s.accountRepo, account, modelKey, prefix, statusCode, ra, false) {
+ logger.LegacyPrintf("service.antigravity_gateway", "%s status=429 model_rate_limit_set_failed model=%s", prefix, modelKey)
} else {
logger.LegacyPrintf("service.antigravity_gateway", "%s status=429 model_rate_limited model=%s account=%d reset_at=%v reset_in=%v",
prefix, modelKey, account.ID, ra.Format("15:04:05"), time.Until(ra).Truncate(time.Second))
- s.updateAccountModelRateLimitInCache(ctx, account, modelKey, ra)
}
return nil
}
diff --git a/backend/internal/service/antigravity_gateway_service_test.go b/backend/internal/service/antigravity_gateway_service_test.go
index 22124374..1aa1ac9d 100644
--- a/backend/internal/service/antigravity_gateway_service_test.go
+++ b/backend/internal/service/antigravity_gateway_service_test.go
@@ -491,6 +491,86 @@ func TestAntigravityGatewayService_ForwardGemini_StickySessionForceCacheBilling(
require.True(t, failoverErr.ForceCacheBilling, "ForceCacheBilling should be true for sticky session switch")
}
+func TestAntigravityGatewayService_ForwardGemini_ClearsStickySessionOnGeminiRateLimit(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ writer := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(writer)
+
+ body, err := json.Marshal(map[string]any{
+ "contents": []map[string]any{
+ {"role": "user", "parts": []map[string]any{{"text": "hi"}}},
+ },
+ })
+ require.NoError(t, err)
+
+ req := httptest.NewRequest(http.MethodPost, "/v1beta/models/gemini-3-flash-preview:generateContent", bytes.NewReader(body))
+ c.Request = req
+
+ respBody := []byte(`{
+ "error": {
+ "status": "RESOURCE_EXHAUSTED",
+ "details": [
+ {"@type": "type.googleapis.com/google.rpc.ErrorInfo", "metadata": {"model": "gemini-3-flash"}, "reason": "RATE_LIMIT_EXCEEDED"},
+ {"@type": "type.googleapis.com/google.rpc.RetryInfo", "retryDelay": "15s"}
+ ]
+ }
+ }`)
+ upstream := &httpUpstreamStub{resp: &http.Response{
+ StatusCode: http.StatusTooManyRequests,
+ Header: http.Header{},
+ Body: io.NopCloser(bytes.NewReader(respBody)),
+ }}
+ repo := &stubAntigravityAccountRepo{}
+ cache := &stubSmartRetryCache{}
+ svc := &AntigravityGatewayService{
+ tokenProvider: &AntigravityTokenProvider{},
+ httpUpstream: upstream,
+ accountRepo: repo,
+ cache: cache,
+ }
+
+ account := &Account{
+ ID: 44,
+ Name: "acc-gemini-runtime-rate-limited",
+ Platform: PlatformAntigravity,
+ Type: AccountTypeOAuth,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "access_token": "token",
+ "expires_at": time.Now().Add(time.Hour).Format(time.RFC3339),
+ "project_id": "proj",
+ },
+ Extra: map[string]any{
+ "mixed_scheduling": true,
+ },
+ }
+
+ result, err := svc.ForwardGemini(
+ context.Background(),
+ c,
+ account,
+ "gemini-3-flash-preview",
+ "generateContent",
+ false,
+ body,
+ true,
+ WithForwardGeminiSession(77, "gemini:sticky-runtime"),
+ )
+
+ require.Nil(t, result)
+ var failoverErr *UpstreamFailoverError
+ require.ErrorAs(t, err, &failoverErr)
+ require.Equal(t, http.StatusServiceUnavailable, failoverErr.StatusCode)
+ require.Len(t, repo.modelRateLimitCalls, 2)
+ require.Equal(t, "gemini-3-flash", repo.modelRateLimitCalls[0].modelKey)
+ require.Equal(t, antigravityGeminiModelRateLimitKey, repo.modelRateLimitCalls[1].modelKey)
+ require.Len(t, cache.deleteCalls, 1)
+ require.Equal(t, int64(77), cache.deleteCalls[0].groupID)
+ require.Equal(t, "gemini:sticky-runtime", cache.deleteCalls[0].sessionHash)
+}
+
// TestAntigravityGatewayService_Forward_BillsWithMappedModel
// 验证:Antigravity Claude 转发返回的计费模型使用映射后的模型
func TestAntigravityGatewayService_Forward_BillsWithMappedModel(t *testing.T) {
diff --git a/backend/internal/service/antigravity_quota_scope.go b/backend/internal/service/antigravity_quota_scope.go
index b536d16c..75862633 100644
--- a/backend/internal/service/antigravity_quota_scope.go
+++ b/backend/internal/service/antigravity_quota_scope.go
@@ -8,7 +8,17 @@ import (
func normalizeAntigravityModelName(model string) string {
normalized := strings.ToLower(strings.TrimSpace(model))
- normalized = strings.TrimPrefix(normalized, "models/")
+ if idx := strings.LastIndex(normalized, "/publishers/google/models/"); idx != -1 {
+ normalized = normalized[idx+len("/publishers/google/models/"):]
+ } else if idx := strings.LastIndex(normalized, "/publishers/anthropic/models/"); idx != -1 {
+ normalized = normalized[idx+len("/publishers/anthropic/models/"):]
+ } else if idx := strings.LastIndex(normalized, "/models/"); idx != -1 {
+ normalized = normalized[idx+len("/models/"):]
+ } else {
+ normalized = strings.TrimPrefix(normalized, "publishers/google/models/")
+ normalized = strings.TrimPrefix(normalized, "publishers/anthropic/models/")
+ normalized = strings.TrimPrefix(normalized, "models/")
+ }
return normalized
}
diff --git a/backend/internal/service/antigravity_rate_limit_test.go b/backend/internal/service/antigravity_rate_limit_test.go
index c3e49458..374b29f6 100644
--- a/backend/internal/service/antigravity_rate_limit_test.go
+++ b/backend/internal/service/antigravity_rate_limit_test.go
@@ -821,6 +821,51 @@ func TestSetModelRateLimitByModelName_NotConvertToScope(t *testing.T) {
require.NotEqual(t, "claude_sonnet", call.modelKey, "should NOT be scope")
}
+func TestSetAntigravityModelRateLimits_GeminiWritesFamilyScope(t *testing.T) {
+ repo := &stubAntigravityAccountRepo{}
+ svc := &AntigravityGatewayService{}
+ account := &Account{ID: 789, Platform: PlatformAntigravity}
+ resetAt := time.Now().Add(30 * time.Second)
+
+ success := svc.setAntigravityModelRateLimits(
+ context.Background(),
+ repo,
+ account,
+ "gemini-3-pro",
+ "[test]",
+ 429,
+ resetAt,
+ false,
+ )
+
+ require.True(t, success)
+ require.Len(t, repo.modelRateLimitCalls, 2)
+ require.Equal(t, "gemini-3-pro", repo.modelRateLimitCalls[0].modelKey)
+ require.Equal(t, antigravityGeminiModelRateLimitKey, repo.modelRateLimitCalls[1].modelKey)
+}
+
+func TestSetAntigravityModelRateLimits_ClaudeDoesNotWriteGeminiScope(t *testing.T) {
+ repo := &stubAntigravityAccountRepo{}
+ svc := &AntigravityGatewayService{}
+ account := &Account{ID: 790, Platform: PlatformAntigravity}
+ resetAt := time.Now().Add(30 * time.Second)
+
+ success := svc.setAntigravityModelRateLimits(
+ context.Background(),
+ repo,
+ account,
+ "claude-sonnet-4-5",
+ "[test]",
+ 429,
+ resetAt,
+ false,
+ )
+
+ require.True(t, success)
+ require.Len(t, repo.modelRateLimitCalls, 1)
+ require.Equal(t, "claude-sonnet-4-5", repo.modelRateLimitCalls[0].modelKey)
+}
+
func TestAntigravityRetryLoop_PreCheck_SwitchesWhenRateLimited(t *testing.T) {
upstream := &recordingOKUpstream{}
account := &Account{
@@ -1124,3 +1169,53 @@ func TestSchedulerSnapshotService_UpdateAccountInCache(t *testing.T) {
require.ErrorIs(t, err, expectedErr)
})
}
+func TestNormalizeAntigravityModelName(t *testing.T) {
+ tests := []struct {
+ name string
+ model string
+ expected string
+ }{
+ {
+ name: "plain model name",
+ model: "gemini-1.5-pro",
+ expected: "gemini-1.5-pro",
+ },
+ {
+ name: "models/ prefix",
+ model: "models/gemini-1.5-pro",
+ expected: "gemini-1.5-pro",
+ },
+ {
+ name: "publishers/google/models/ prefix",
+ model: "publishers/google/models/gemini-1.5-pro",
+ expected: "gemini-1.5-pro",
+ },
+ {
+ name: "projects/.../publishers/google/models/ path",
+ model: "projects/my-proj/locations/us-central1/publishers/google/models/gemini-2.5-flash",
+ expected: "gemini-2.5-flash",
+ },
+ {
+ name: "publishers/anthropic/models/ prefix",
+ model: "publishers/anthropic/models/claude-sonnet-4-5",
+ expected: "claude-sonnet-4-5",
+ },
+ {
+ name: "projects/.../publishers/anthropic/models/ path",
+ model: "projects/my-proj/locations/global/publishers/anthropic/models/claude-sonnet-4-5",
+ expected: "claude-sonnet-4-5",
+ },
+ {
+ name: "mixed case and spaces",
+ model: " Models/Gemini-1.5-Pro ",
+ expected: "gemini-1.5-pro",
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ actual := normalizeAntigravityModelName(tt.model)
+ require.Equal(t, tt.expected, actual)
+ })
+ }
+}
diff --git a/backend/internal/service/antigravity_single_account_retry_test.go b/backend/internal/service/antigravity_single_account_retry_test.go
index 675e9c0c..6e58ab75 100644
--- a/backend/internal/service/antigravity_single_account_retry_test.go
+++ b/backend/internal/service/antigravity_single_account_retry_test.go
@@ -196,8 +196,10 @@ func TestHandleSmartRetry_503_LongDelay_NoSingleAccountRetry_StillSwitches(t *te
require.Nil(t, result.resp, "should not return resp when switchError is set")
// 对照:多账号模式应设模型限流
- require.Len(t, repo.modelRateLimitCalls, 1,
+ require.Len(t, repo.modelRateLimitCalls, 2,
"multi-account mode SHOULD set model rate limit")
+ require.Equal(t, "gemini-3-pro-high", repo.modelRateLimitCalls[0].modelKey)
+ require.Equal(t, antigravityGeminiModelRateLimitKey, repo.modelRateLimitCalls[1].modelKey)
}
// TestHandleSmartRetry_429_LongDelay_SingleAccountRetry_StillSwitches
@@ -412,8 +414,10 @@ func TestHandleSmartRetry_503_ShortDelay_NoSingleAccountRetry_SetsRateLimit(t *t
// 对照:多账号模式应返回 switchError
require.NotNil(t, result.switchError, "multi-account mode should return switchError for 503")
// 对照:多账号模式应设模型限流
- require.Len(t, repo.modelRateLimitCalls, 1,
+ require.Len(t, repo.modelRateLimitCalls, 2,
"multi-account mode should set model rate limit")
+ require.Equal(t, "gemini-3-flash", repo.modelRateLimitCalls[0].modelKey)
+ require.Equal(t, antigravityGeminiModelRateLimitKey, repo.modelRateLimitCalls[1].modelKey)
}
// ---------------------------------------------------------------------------
diff --git a/backend/internal/service/antigravity_smart_retry_test.go b/backend/internal/service/antigravity_smart_retry_test.go
index e3b60a27..9f06b13b 100644
--- a/backend/internal/service/antigravity_smart_retry_test.go
+++ b/backend/internal/service/antigravity_smart_retry_test.go
@@ -328,9 +328,10 @@ func TestHandleSmartRetry_ShortDelay_SmartRetryFailed_ReturnsSwitchError(t *test
require.Equal(t, "gemini-3-flash", result.switchError.RateLimitedModel)
require.False(t, result.switchError.IsStickySession)
- // 验证模型限流已设置
- require.Len(t, repo.modelRateLimitCalls, 1)
+ // 验证模型限流已设置:Gemini 同时写入精确模型和家族级 scope
+ require.Len(t, repo.modelRateLimitCalls, 2)
require.Equal(t, "gemini-3-flash", repo.modelRateLimitCalls[0].modelKey)
+ require.Equal(t, antigravityGeminiModelRateLimitKey, repo.modelRateLimitCalls[1].modelKey)
require.Len(t, upstream.calls, 1, "should have made one retry call (max attempts)")
}
@@ -1104,10 +1105,9 @@ func TestHandleSmartRetry_ShortDelay_StickySession_SuccessRetry_NoDeleteSession(
require.Len(t, cache.deleteCalls, 0, "should NOT call DeleteSessionAccountID on successful retry")
}
-// TestHandleSmartRetry_LongDelay_StickySession_NoDeleteInHandleSmartRetry
-// 长延迟路径(情况1)在 handleSmartRetry 中不直接调用 DeleteSessionAccountID
-// (清除由 handler 层的 shouldClearStickySession 在下次请求时处理)
-func TestHandleSmartRetry_LongDelay_StickySession_NoDeleteInHandleSmartRetry(t *testing.T) {
+// TestHandleSmartRetry_LongDelay_StickySession_ClearsSession
+// 长延迟路径(情况1)应立即清除 sticky 绑定,避免下一次请求继续命中已限流账号。
+func TestHandleSmartRetry_LongDelay_StickySession_ClearsSession(t *testing.T) {
repo := &stubAntigravityAccountRepo{}
cache := &stubSmartRetryCache{}
account := &Account{
@@ -1159,10 +1159,9 @@ func TestHandleSmartRetry_LongDelay_StickySession_NoDeleteInHandleSmartRetry(t *
require.NotNil(t, result.switchError)
require.True(t, result.switchError.IsStickySession)
- // 长延迟路径不在 handleSmartRetry 中调用 DeleteSessionAccountID
- // (由上游 handler 的 shouldClearStickySession 处理)
- require.Len(t, cache.deleteCalls, 0,
- "long delay path should NOT call DeleteSessionAccountID in handleSmartRetry (handled by handler layer)")
+ require.Len(t, cache.deleteCalls, 1, "long delay path should clear sticky session in handleSmartRetry")
+ require.Equal(t, int64(42), cache.deleteCalls[0].groupID)
+ require.Equal(t, "sticky-hash-long-delay", cache.deleteCalls[0].sessionHash)
}
// TestHandleSmartRetry_ShortDelay_NetworkError_StickySession_ClearsSession
@@ -1227,6 +1226,10 @@ func TestHandleSmartRetry_ShortDelay_NetworkError_StickySession_ClearsSession(t
require.Len(t, cache.deleteCalls, 1, "should call DeleteSessionAccountID after network error exhausts retry")
require.Equal(t, int64(99), cache.deleteCalls[0].groupID)
require.Equal(t, "sticky-net-error", cache.deleteCalls[0].sessionHash)
+
+ require.Len(t, repo.modelRateLimitCalls, 2)
+ require.Equal(t, "gemini-3-flash", repo.modelRateLimitCalls[0].modelKey)
+ require.Equal(t, antigravityGeminiModelRateLimitKey, repo.modelRateLimitCalls[1].modelKey)
}
// TestHandleSmartRetry_ShortDelay_503_StickySession_FailedRetry_ClearsSession
@@ -1308,9 +1311,10 @@ func TestHandleSmartRetry_ShortDelay_503_StickySession_FailedRetry_ClearsSession
require.Equal(t, int64(77), cache.deleteCalls[0].groupID)
require.Equal(t, "sticky-503-short", cache.deleteCalls[0].sessionHash)
- // 验证模型限流已设置
- require.Len(t, repo.modelRateLimitCalls, 1)
+ // 验证模型限流已设置:Gemini 同时写入精确模型和家族级 scope
+ require.Len(t, repo.modelRateLimitCalls, 2)
require.Equal(t, "gemini-3-pro", repo.modelRateLimitCalls[0].modelKey)
+ require.Equal(t, antigravityGeminiModelRateLimitKey, repo.modelRateLimitCalls[1].modelKey)
}
// TestAntigravityRetryLoop_SmartRetryFailed_StickySession_SwitchErrorPropagates
diff --git a/backend/internal/service/error_policy_test.go b/backend/internal/service/error_policy_test.go
index 297a954c..2aa7a421 100644
--- a/backend/internal/service/error_policy_test.go
+++ b/backend/internal/service/error_policy_test.go
@@ -389,6 +389,60 @@ func TestApplyErrorPolicy(t *testing.T) {
}
}
+func TestApplyErrorPolicy_GeminiRateLimitBypassesCustomSkip(t *testing.T) {
+ repo := &stubAntigravityAccountRepo{}
+ cache := &stubSmartRetryCache{}
+ rlSvc := NewRateLimitService(repo, nil, &config.Config{}, nil, nil)
+ svc := &AntigravityGatewayService{
+ rateLimitService: rlSvc,
+ accountRepo: repo,
+ cache: cache,
+ }
+
+ account := &Account{
+ ID: 31,
+ Type: AccountTypeAPIKey,
+ Platform: PlatformAntigravity,
+ Credentials: map[string]any{
+ "custom_error_codes_enabled": true,
+ "custom_error_codes": []any{float64(500)},
+ },
+ }
+ body := []byte(`{
+ "error": {
+ "status": "RESOURCE_EXHAUSTED",
+ "details": [
+ {"@type": "type.googleapis.com/google.rpc.ErrorInfo", "metadata": {"model": "gemini-3-flash"}, "reason": "RATE_LIMIT_EXCEEDED"},
+ {"@type": "type.googleapis.com/google.rpc.RetryInfo", "retryDelay": "15s"}
+ ]
+ }
+ }`)
+ p := antigravityRetryLoopParams{
+ ctx: context.Background(),
+ prefix: "[test]",
+ account: account,
+ accountRepo: repo,
+ groupID: 42,
+ sessionHash: "gemini:sticky",
+ handleError: func(context.Context, string, *Account, int, http.Header, []byte, string, int64, string, bool) *handleModelRateLimitResult {
+ t.Fatal("model rate limit should be handled before custom error fallback")
+ return nil
+ },
+ }
+
+ handled, outStatus, retErr := svc.applyErrorPolicy(p, http.StatusTooManyRequests, http.Header{}, body)
+
+ require.True(t, handled)
+ require.Equal(t, http.StatusTooManyRequests, outStatus)
+ require.NoError(t, retErr)
+ require.Len(t, repo.modelRateLimitCalls, 2)
+ require.Equal(t, "gemini-3-flash", repo.modelRateLimitCalls[0].modelKey)
+ require.Equal(t, antigravityGeminiModelRateLimitKey, repo.modelRateLimitCalls[1].modelKey)
+ require.Len(t, cache.deleteCalls, 1)
+ require.Equal(t, int64(42), cache.deleteCalls[0].groupID)
+ require.Equal(t, "gemini:sticky", cache.deleteCalls[0].sessionHash)
+}
+
// ---------------------------------------------------------------------------
// errorPolicyRepoStub — minimal AccountRepository stub for error policy tests
// ---------------------------------------------------------------------------
diff --git a/backend/internal/service/gateway_multiplatform_test.go b/backend/internal/service/gateway_multiplatform_test.go
index a00b7dc1..7a6acaac 100644
--- a/backend/internal/service/gateway_multiplatform_test.go
+++ b/backend/internal/service/gateway_multiplatform_test.go
@@ -1229,6 +1229,106 @@ func TestGatewayService_selectAccountWithMixedScheduling(t *testing.T) {
require.Equal(t, int64(2), acc.ID, "应选择优先级最高的账户(包含启用混合调度的antigravity)")
})
+ t.Run("混合调度-Gemini家族限流后跳过Antigravity账户", func(t *testing.T) {
+ resetAt := time.Now().Add(10 * time.Minute).Format(time.RFC3339)
+ repo := &mockAccountRepoForPlatform{
+ accounts: []Account{
+ {
+ ID: 1,
+ Platform: PlatformAntigravity,
+ Priority: 1,
+ Status: StatusActive,
+ Schedulable: true,
+ Extra: map[string]any{
+ "mixed_scheduling": true,
+ modelRateLimitsKey: map[string]any{
+ antigravityGeminiModelRateLimitKey: map[string]any{
+ "rate_limit_reset_at": resetAt,
+ },
+ },
+ },
+ },
+ {
+ ID: 2,
+ Platform: PlatformAntigravity,
+ Priority: 1,
+ Status: StatusActive,
+ Schedulable: true,
+ Extra: map[string]any{
+ "mixed_scheduling": true,
+ modelRateLimitsKey: map[string]any{
+ antigravityGeminiModelRateLimitKey: map[string]any{
+ "rate_limit_reset_at": resetAt,
+ },
+ },
+ },
+ },
+ {
+ ID: 3,
+ Platform: PlatformAntigravity,
+ Priority: 2,
+ Status: StatusActive,
+ Schedulable: true,
+ Extra: map[string]any{"mixed_scheduling": true},
+ },
+ },
+ accountsByID: map[int64]*Account{},
+ }
+ for i := range repo.accounts {
+ repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i]
+ }
+
+ svc := &GatewayService{
+ accountRepo: repo,
+ cache: &mockGatewayCacheForPlatform{},
+ cfg: testConfig(),
+ }
+
+ acc, err := svc.selectAccountWithMixedScheduling(ctx, nil, "", "gemini-3-pro-preview", nil, PlatformGemini)
+ require.NoError(t, err)
+ require.NotNil(t, acc)
+ require.Equal(t, int64(3), acc.ID)
+ })
+
+ t.Run("混合调度-Gemini家族限流不影响Claude调度", func(t *testing.T) {
+ resetAt := time.Now().Add(10 * time.Minute).Format(time.RFC3339)
+ repo := &mockAccountRepoForPlatform{
+ accounts: []Account{
+ {
+ ID: 1,
+ Platform: PlatformAntigravity,
+ Priority: 1,
+ Status: StatusActive,
+ Schedulable: true,
+ Extra: map[string]any{
+ "mixed_scheduling": true,
+ modelRateLimitsKey: map[string]any{
+ antigravityGeminiModelRateLimitKey: map[string]any{
+ "rate_limit_reset_at": resetAt,
+ },
+ },
+ },
+ },
+ {ID: 2, Platform: PlatformAnthropic, Priority: 2, Status: StatusActive, Schedulable: true},
+ },
+ accountsByID: map[int64]*Account{},
+ }
+ for i := range repo.accounts {
+ repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i]
+ }
+
+ svc := &GatewayService{
+ accountRepo: repo,
+ cache: &mockGatewayCacheForPlatform{},
+ cfg: testConfig(),
+ }
+
+ acc, err := svc.selectAccountWithMixedScheduling(ctx, nil, "", "claude-sonnet-4-5", nil, PlatformAnthropic)
+ require.NoError(t, err)
+ require.NotNil(t, acc)
+ require.Equal(t, int64(1), acc.ID)
+ })
+
t.Run("混合调度-路由优先选择路由账号", func(t *testing.T) {
groupID := int64(30)
requestedModel := "claude-sonnet-4-5"
diff --git a/backend/internal/service/model_rate_limit.go b/backend/internal/service/model_rate_limit.go
index c45615cc..420f1446 100644
--- a/backend/internal/service/model_rate_limit.go
+++ b/backend/internal/service/model_rate_limit.go
@@ -6,7 +6,10 @@ import (
"time"
)
-const modelRateLimitsKey = "model_rate_limits"
+const (
+ modelRateLimitsKey = "model_rate_limits"
+ antigravityGeminiModelRateLimitKey = "antigravity:gemini"
+)
// isRateLimitActiveForKey 检查指定 key 的限流是否生效
func (a *Account) isRateLimitActiveForKey(key string) bool {
@@ -35,6 +38,9 @@ func (a *Account) isModelRateLimitedWithContext(ctx context.Context, requestedMo
modelKey := a.GetMappedModel(requestedModel)
if a.Platform == PlatformAntigravity {
modelKey = resolveFinalAntigravityModelKey(ctx, a, requestedModel)
+ if isAntigravityGeminiModel(modelKey) && a.isRateLimitActiveForKey(antigravityGeminiModelRateLimitKey) {
+ return true
+ }
}
modelKey = strings.TrimSpace(modelKey)
if modelKey == "" {
@@ -62,7 +68,13 @@ func (a *Account) GetModelRateLimitRemainingTimeWithContext(ctx context.Context,
if modelKey == "" {
return 0
}
- return a.getRateLimitRemainingForKey(modelKey)
+ remaining := a.getRateLimitRemainingForKey(modelKey)
+ if a.Platform == PlatformAntigravity && isAntigravityGeminiModel(modelKey) {
+ if familyRemaining := a.getRateLimitRemainingForKey(antigravityGeminiModelRateLimitKey); familyRemaining > remaining {
+ return familyRemaining
+ }
+ }
+ return remaining
}
func resolveFinalAntigravityModelKey(ctx context.Context, account *Account, requestedModel string) string {
@@ -77,6 +89,22 @@ func resolveFinalAntigravityModelKey(ctx context.Context, account *Account, requ
return modelKey
}
+func isAntigravityGeminiModel(model string) bool {
+ return strings.HasPrefix(normalizeAntigravityModelName(model), "gemini-")
+}
+
+func antigravityModelRateLimitKeys(model string) []string {
+ model = strings.TrimSpace(model)
+ if model == "" {
+ return nil
+ }
+ keys := []string{model}
+ if isAntigravityGeminiModel(model) && model != antigravityGeminiModelRateLimitKey {
+ keys = append(keys, antigravityGeminiModelRateLimitKey)
+ }
+ return keys
+}
+
func (a *Account) modelRateLimitResetAt(scope string) *time.Time {
if a == nil || a.Extra == nil || scope == "" {
return nil
diff --git a/backend/internal/service/model_rate_limit_test.go b/backend/internal/service/model_rate_limit_test.go
index b79b9688..3cce6459 100644
--- a/backend/internal/service/model_rate_limit_test.go
+++ b/backend/internal/service/model_rate_limit_test.go
@@ -121,6 +121,36 @@ func TestIsModelRateLimited(t *testing.T) {
requestedModel: "gemini-3-pro-preview",
expected: true,
},
+ {
+ name: "antigravity platform - gemini family rate limit blocks mapped preview",
+ account: &Account{
+ Platform: PlatformAntigravity,
+ Extra: map[string]any{
+ modelRateLimitsKey: map[string]any{
+ antigravityGeminiModelRateLimitKey: map[string]any{
+ "rate_limit_reset_at": future,
+ },
+ },
+ },
+ },
+ requestedModel: "gemini-3-pro-preview",
+ expected: true,
+ },
+ {
+ name: "antigravity platform - gemini family rate limit does not block claude",
+ account: &Account{
+ Platform: PlatformAntigravity,
+ Extra: map[string]any{
+ modelRateLimitsKey: map[string]any{
+ antigravityGeminiModelRateLimitKey: map[string]any{
+ "rate_limit_reset_at": future,
+ },
+ },
+ },
+ },
+ requestedModel: "claude-sonnet-4-5",
+ expected: false,
+ },
{
name: "non-antigravity platform - gemini-3-pro-preview NOT mapped",
account: &Account{
@@ -306,6 +336,38 @@ func TestGetModelRateLimitRemainingTime(t *testing.T) {
minExpected: 4 * time.Minute,
maxExpected: 6 * time.Minute,
},
+ {
+ name: "antigravity platform - gemini family rate limit remaining",
+ account: &Account{
+ Platform: PlatformAntigravity,
+ Extra: map[string]any{
+ modelRateLimitsKey: map[string]any{
+ antigravityGeminiModelRateLimitKey: map[string]any{
+ "rate_limit_reset_at": future10m,
+ },
+ },
+ },
+ },
+ requestedModel: "gemini-3-pro-preview",
+ minExpected: 9 * time.Minute,
+ maxExpected: 11 * time.Minute,
+ },
+ {
+ name: "antigravity platform - gemini family remaining ignored for claude",
+ account: &Account{
+ Platform: PlatformAntigravity,
+ Extra: map[string]any{
+ modelRateLimitsKey: map[string]any{
+ antigravityGeminiModelRateLimitKey: map[string]any{
+ "rate_limit_reset_at": future10m,
+ },
+ },
+ },
+ },
+ requestedModel: "claude-sonnet-4-5",
+ minExpected: 0,
+ maxExpected: 0,
+ },
}
for _, tt := range tests {
diff --git a/backend/internal/service/scheduler_snapshot_hydration_test.go b/backend/internal/service/scheduler_snapshot_hydration_test.go
index 778cab23..0a1d0a0a 100644
--- a/backend/internal/service/scheduler_snapshot_hydration_test.go
+++ b/backend/internal/service/scheduler_snapshot_hydration_test.go
@@ -6,6 +6,8 @@ import (
"context"
"testing"
"time"
+
+ "github.com/Wei-Shaw/sub2api/internal/config"
)
type snapshotHydrationCache struct {
@@ -186,3 +188,89 @@ func TestGatewaySelectAccountWithLoadAwareness_HydratesSelectedAccountFromSchedu
t.Fatalf("expected hydrated api key, got %q", got)
}
}
+
+func TestGatewaySelectAccountWithLoadAwareness_SkipsAntigravityGeminiFamilyRateLimitedSnapshot(t *testing.T) {
+ resetAt := time.Now().Add(10 * time.Minute).Format(time.RFC3339)
+ cache := &snapshotHydrationCache{
+ snapshot: []*Account{
+ {
+ ID: 1,
+ Platform: PlatformAntigravity,
+ Type: AccountTypeOAuth,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 1,
+ AccountGroups: []AccountGroup{
+ {AccountID: 1, GroupID: 22},
+ },
+ GroupIDs: []int64{22},
+ Extra: map[string]any{
+ "mixed_scheduling": true,
+ modelRateLimitsKey: map[string]any{
+ antigravityGeminiModelRateLimitKey: map[string]any{
+ "rate_limit_reset_at": resetAt,
+ },
+ },
+ },
+ },
+ {
+ ID: 2,
+ Platform: PlatformAntigravity,
+ Type: AccountTypeOAuth,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 2,
+ AccountGroups: []AccountGroup{
+ {AccountID: 2, GroupID: 22},
+ },
+ GroupIDs: []int64{22},
+ Extra: map[string]any{
+ "mixed_scheduling": true,
+ },
+ },
+ },
+ accounts: map[int64]*Account{
+ 1: {ID: 1, Platform: PlatformAntigravity, Type: AccountTypeOAuth},
+ 2: {ID: 2, Platform: PlatformAntigravity, Type: AccountTypeOAuth},
+ },
+ }
+ groupID := int64(22)
+ svc := &GatewayService{
+ schedulerSnapshot: NewSchedulerSnapshotService(cache, nil, nil, nil, nil),
+ groupRepo: &mockGroupRepoForGateway{
+ groups: map[int64]*Group{
+ groupID: {
+ ID: groupID,
+ Platform: PlatformGemini,
+ Status: StatusActive,
+ Hydrated: true,
+ },
+ },
+ },
+ concurrencyService: NewConcurrencyService(&mockConcurrencyCache{}),
+ cfg: &config.Config{
+ Gateway: config.GatewayConfig{
+ Scheduling: config.GatewaySchedulingConfig{
+ LoadBatchEnabled: true,
+ StickySessionMaxWaiting: 3,
+ StickySessionWaitTimeout: time.Second,
+ FallbackWaitTimeout: time.Second,
+ FallbackMaxWaiting: 10,
+ },
+ },
+ },
+ }
+
+ result, err := svc.SelectAccountWithLoadAwareness(context.Background(), &groupID, "", "gemini-3-flash-preview", nil, "", 0)
+ if err != nil {
+ t.Fatalf("SelectAccountWithLoadAwareness error: %v", err)
+ }
+ if result == nil || result.Account == nil {
+ t.Fatalf("expected selected account")
+ }
+ if result.Account.ID != 2 {
+ t.Fatalf("expected scheduler to skip Gemini-family limited antigravity account 1, got %d", result.Account.ID)
+ }
+}
From 0560340bd4c748c436053febd9009674301ecfb2 Mon Sep 17 00:00:00 2001
From: fatelei
Date: Mon, 1 Jun 2026 08:41:35 +0800
Subject: [PATCH 63/79] fix: change balance to pointer type
---
.../internal/handler/admin/user_handler.go | 16 +++---
backend/internal/service/admin_service.go | 11 +++-
.../service/admin_service_create_user_test.go | 55 ++++++++++++++++++-
frontend/src/api/admin/users.ts | 3 +
.../components/admin/user/UserCreateModal.vue | 14 +++--
5 files changed, 83 insertions(+), 16 deletions(-)
diff --git a/backend/internal/handler/admin/user_handler.go b/backend/internal/handler/admin/user_handler.go
index db35472e..061bc1b2 100644
--- a/backend/internal/handler/admin/user_handler.go
+++ b/backend/internal/handler/admin/user_handler.go
@@ -34,14 +34,14 @@ func NewUserHandler(adminService service.AdminService, concurrencyService *servi
// CreateUserRequest represents admin create user request
type CreateUserRequest struct {
- Email string `json:"email" binding:"required,email"`
- Password string `json:"password" binding:"required,min=6"`
- Username string `json:"username"`
- Notes string `json:"notes"`
- Balance float64 `json:"balance"`
- Concurrency int `json:"concurrency"`
- RPMLimit int `json:"rpm_limit"`
- AllowedGroups []int64 `json:"allowed_groups"`
+ Email string `json:"email" binding:"required,email"`
+ Password string `json:"password" binding:"required,min=6"`
+ Username string `json:"username"`
+ Notes string `json:"notes"`
+ Balance *float64 `json:"balance"`
+ Concurrency int `json:"concurrency"`
+ RPMLimit int `json:"rpm_limit"`
+ AllowedGroups []int64 `json:"allowed_groups"`
}
// UpdateUserRequest represents admin update user request
diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go
index eb5994d5..09aeef86 100644
--- a/backend/internal/service/admin_service.go
+++ b/backend/internal/service/admin_service.go
@@ -120,7 +120,7 @@ type CreateUserInput struct {
Password string
Username string
Notes string
- Balance float64
+ Balance *float64
Concurrency int
RPMLimit int
AllowedGroups []int64
@@ -661,12 +661,19 @@ func (s *adminServiceImpl) GetUser(ctx context.Context, id int64) (*User, error)
}
func (s *adminServiceImpl) CreateUser(ctx context.Context, input *CreateUserInput) (*User, error) {
+ balance := 0.0
+ if input.Balance != nil {
+ balance = *input.Balance
+ } else if s.settingService != nil {
+ balance = s.settingService.GetDefaultBalance(ctx)
+ }
+
user := &User{
Email: input.Email,
Username: input.Username,
Notes: input.Notes,
Role: RoleUser, // Always create as regular user, never admin
- Balance: input.Balance,
+ Balance: balance,
Concurrency: input.Concurrency,
RPMLimit: input.RPMLimit,
Status: StatusActive,
diff --git a/backend/internal/service/admin_service_create_user_test.go b/backend/internal/service/admin_service_create_user_test.go
index c5b1e38d..5e9578ab 100644
--- a/backend/internal/service/admin_service_create_user_test.go
+++ b/backend/internal/service/admin_service_create_user_test.go
@@ -14,13 +14,14 @@ import (
func TestAdminService_CreateUser_Success(t *testing.T) {
repo := &userRepoStub{nextID: 10}
svc := &adminServiceImpl{userRepo: repo}
+ balance := 12.5
input := &CreateUserInput{
Email: "user@test.com",
Password: "strong-pass",
Username: "tester",
Notes: "note",
- Balance: 12.5,
+ Balance: &balance,
Concurrency: 7,
AllowedGroups: []int64{3, 5},
}
@@ -32,7 +33,7 @@ func TestAdminService_CreateUser_Success(t *testing.T) {
require.Equal(t, input.Email, user.Email)
require.Equal(t, input.Username, user.Username)
require.Equal(t, input.Notes, user.Notes)
- require.Equal(t, input.Balance, user.Balance)
+ require.Equal(t, balance, user.Balance)
require.Equal(t, input.Concurrency, user.Concurrency)
require.Equal(t, input.AllowedGroups, user.AllowedGroups)
require.Equal(t, RoleUser, user.Role)
@@ -42,6 +43,56 @@ func TestAdminService_CreateUser_Success(t *testing.T) {
require.Equal(t, user, repo.created[0])
}
+func TestAdminService_CreateUser_UsesDefaultBalanceWhenBalanceOmitted(t *testing.T) {
+ repo := &userRepoStub{nextID: 11}
+ cfg := &config.Config{
+ Default: config.DefaultConfig{
+ UserBalance: 0,
+ },
+ }
+ settingService := NewSettingService(&settingRepoStub{values: map[string]string{
+ SettingKeyDefaultBalance: "0.02",
+ }}, cfg)
+ svc := &adminServiceImpl{userRepo: repo, settingService: settingService}
+
+ user, err := svc.CreateUser(context.Background(), &CreateUserInput{
+ Email: "default-balance@test.com",
+ Password: "strong-pass",
+ })
+
+ require.NoError(t, err)
+ require.NotNil(t, user)
+ require.Equal(t, 0.02, user.Balance)
+ require.Len(t, repo.created, 1)
+ require.Equal(t, 0.02, repo.created[0].Balance)
+}
+
+func TestAdminService_CreateUser_ExplicitZeroBalanceOverridesDefault(t *testing.T) {
+ repo := &userRepoStub{nextID: 12}
+ cfg := &config.Config{
+ Default: config.DefaultConfig{
+ UserBalance: 0,
+ },
+ }
+ settingService := NewSettingService(&settingRepoStub{values: map[string]string{
+ SettingKeyDefaultBalance: "0.02",
+ }}, cfg)
+ svc := &adminServiceImpl{userRepo: repo, settingService: settingService}
+ balance := 0.0
+
+ user, err := svc.CreateUser(context.Background(), &CreateUserInput{
+ Email: "zero-balance@test.com",
+ Password: "strong-pass",
+ Balance: &balance,
+ })
+
+ require.NoError(t, err)
+ require.NotNil(t, user)
+ require.Equal(t, 0.0, user.Balance)
+ require.Len(t, repo.created, 1)
+ require.Equal(t, 0.0, repo.created[0].Balance)
+}
+
func TestAdminService_CreateUser_EmailExists(t *testing.T) {
repo := &userRepoStub{createErr: ErrEmailExists}
svc := &adminServiceImpl{userRepo: repo}
diff --git a/frontend/src/api/admin/users.ts b/frontend/src/api/admin/users.ts
index fabc69bc..c33eacee 100644
--- a/frontend/src/api/admin/users.ts
+++ b/frontend/src/api/admin/users.ts
@@ -115,8 +115,11 @@ export async function getById(id: number): Promise {
export async function create(userData: {
email: string
password: string
+ username?: string
+ notes?: string
balance?: number
concurrency?: number
+ rpm_limit?: number
allowed_groups?: number[] | null
}): Promise {
const { data } = await apiClient.post('/admin/users', userData)
diff --git a/frontend/src/components/admin/user/UserCreateModal.vue b/frontend/src/components/admin/user/UserCreateModal.vue
index 2966a23b..a638e79c 100644
--- a/frontend/src/components/admin/user/UserCreateModal.vue
+++ b/frontend/src/components/admin/user/UserCreateModal.vue
@@ -28,7 +28,7 @@
-
+
@@ -69,18 +69,24 @@ import Icon from '@/components/icons/Icon.vue'
const props = defineProps<{ show: boolean }>()
const emit = defineEmits(['close', 'success']); const { t } = useI18n()
-const form = reactive({ email: '', password: '', username: '', notes: '', balance: 0, concurrency: 1, rpm_limit: 0 })
+const form = reactive({ email: '', password: '', username: '', notes: '', balance: '', concurrency: 1, rpm_limit: 0 })
const { loading, submit } = useForm({
form,
submitFn: async (data) => {
- await adminAPI.users.create(data)
+ const { balance: rawBalance, ...rest } = data
+ const balance = String(rawBalance).trim()
+ const payload: typeof rest & { balance?: number } = { ...rest }
+ if (balance !== '') {
+ payload.balance = Number(balance)
+ }
+ await adminAPI.users.create(payload)
emit('success'); emit('close')
},
successMsg: t('admin.users.userCreated')
})
-watch(() => props.show, (v) => { if(v) Object.assign(form, { email: '', password: '', username: '', notes: '', balance: 0, concurrency: 1, rpm_limit: 0 }) })
+watch(() => props.show, (v) => { if(v) Object.assign(form, { email: '', password: '', username: '', notes: '', balance: '', concurrency: 1, rpm_limit: 0 }) })
const generateRandomPassword = () => {
const chars = 'ABCDEFGHJKLMNPQRSTUVWXYZabcdefghjkmnpqrstuvwxyz23456789!@#$%^&*'
From c8cd91e3ce1b588c374e310ddcb80bd975e7af0d Mon Sep 17 00:00:00 2001
From: wucm667
Date: Sun, 31 May 2026 08:47:13 +0800
Subject: [PATCH 64/79] =?UTF-8?q?test(openai):=20=E8=A6=86=E7=9B=96=20fail?=
=?UTF-8?q?over=20=E8=AF=B7=E6=B1=82=E4=BD=93=E9=87=8D=E6=98=A0=E5=B0=84?=
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit
---
.../openai_failover_cached_body_test.go | 134 ++++++++++++++++++
1 file changed, 134 insertions(+)
create mode 100644 backend/internal/service/openai_failover_cached_body_test.go
diff --git a/backend/internal/service/openai_failover_cached_body_test.go b/backend/internal/service/openai_failover_cached_body_test.go
new file mode 100644
index 00000000..776f1b1a
--- /dev/null
+++ b/backend/internal/service/openai_failover_cached_body_test.go
@@ -0,0 +1,134 @@
+package service
+
+import (
+ "bytes"
+ "context"
+ "errors"
+ "io"
+ "net/http"
+ "net/http/httptest"
+ "strings"
+ "testing"
+
+ "github.com/gin-gonic/gin"
+ "github.com/stretchr/testify/require"
+ "github.com/tidwall/gjson"
+)
+
+func TestOpenAIGatewayService_Forward_FailoverReparsesCachedBodyForNextAccount(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ tests := []struct {
+ name string
+ requestModel string
+ firstMapping map[string]any
+ secondMapping map[string]any
+ wantFirst string
+ wantSecond string
+ }{
+ {
+ name: "both accounts have mapping",
+ firstMapping: map[string]any{"alias-model": "base-model-a"},
+ secondMapping: map[string]any{"alias-model": "base-model-b"},
+ wantFirst: "base-model-a",
+ wantSecond: "base-model-b",
+ },
+ {
+ name: "first account has mapping second account has none",
+ requestModel: "gpt-5.4-high",
+ firstMapping: map[string]any{"gpt-5.4-high": "gpt-5.4"},
+ wantFirst: "gpt-5.4",
+ wantSecond: "gpt-5.4",
+ },
+ {
+ name: "first account has no mapping second account has mapping",
+ secondMapping: map[string]any{"alias-model": "base-model-b"},
+ wantFirst: "alias-model",
+ wantSecond: "base-model-b",
+ },
+ {
+ name: "legacy context cache is ignored when mappings differ",
+ firstMapping: map[string]any{"alias-model": "base-model-a"},
+ secondMapping: map[string]any{"alias-model": "base-model-b"},
+ wantFirst: "base-model-a",
+ wantSecond: "base-model-b",
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ requestModel := tt.requestModel
+ if requestModel == "" {
+ requestModel = "alias-model"
+ }
+ body := []byte(`{"model":"` + requestModel + `","stream":false,"instructions":"cache-test","input":"hello"}`)
+
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
+ c.Request.Header.Set("Content-Type", "application/json")
+
+ upstream := &httpUpstreamRecorder{responses: []*http.Response{
+ {
+ StatusCode: http.StatusTooManyRequests,
+ Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid-failover-a"}},
+ Body: io.NopCloser(strings.NewReader(`{"error":{"type":"rate_limit_error","message":"rate limited"}}`)),
+ },
+ {
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid-ok-b"}},
+ Body: io.NopCloser(strings.NewReader(`{"id":"resp_123","status":"completed","model":"ok","output":[],"usage":{"input_tokens":1,"output_tokens":1}}`)),
+ },
+ }}
+ svc := &OpenAIGatewayService{httpUpstream: upstream}
+
+ firstAccount := openAIFailoverCachedBodyTestAccount(1, "account-a", tt.firstMapping)
+ secondAccount := openAIFailoverCachedBodyTestAccount(2, "account-b", tt.secondMapping)
+
+ _, err := svc.Forward(context.Background(), c, firstAccount, body)
+ require.Error(t, err)
+ var failoverErr *UpstreamFailoverError
+ require.True(t, errors.As(err, &failoverErr))
+ require.Len(t, upstream.bodies, 1)
+ require.Equal(t, tt.wantFirst, gjson.GetBytes(upstream.bodies[0], "model").String())
+
+ c.Set("openai_parsed_request_body", map[string]any{"model": tt.wantFirst, "stream": true})
+ result, err := svc.Forward(context.Background(), c, secondAccount, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.Len(t, upstream.bodies, 2)
+ require.Equal(t, tt.wantSecond, gjson.GetBytes(upstream.bodies[1], "model").String())
+ })
+ }
+}
+
+func TestGetOpenAIRequestBodyMap_IgnoresLegacyContextCache(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Set("openai_parsed_request_body", map[string]any{"model": "base-model-a", "stream": true})
+
+ got, err := getOpenAIRequestBodyMap(c, []byte(`{"model":"alias-model","stream":false}`))
+ require.NoError(t, err)
+ require.Equal(t, "alias-model", got["model"])
+ require.Equal(t, false, got["stream"])
+}
+
+func openAIFailoverCachedBodyTestAccount(id int64, name string, mapping map[string]any) *Account {
+ credentials := map[string]any{"access_token": "oauth-token", "chatgpt_account_id": "chatgpt-account"}
+ if mapping != nil {
+ credentials["model_mapping"] = mapping
+ }
+ return &Account{
+ ID: id,
+ Name: name,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Concurrency: 1,
+ Credentials: credentials,
+ Status: StatusActive,
+ Schedulable: true,
+ RateMultiplier: f64p(1),
+ }
+}
From 08e19bb15c2091885dbe214297d168c045c474c6 Mon Sep 17 00:00:00 2001
From: xlx0852 <33200995+xlx0852@users.noreply.github.com>
Date: Mon, 1 Jun 2026 10:32:32 +0800
Subject: [PATCH 65/79] fix(openai): bridge oversized websocket requests
---
backend/internal/config/config.go | 18 +
backend/internal/config/config_test.go | 9 +
.../handler/openai_gateway_handler.go | 2 +-
.../service/openai_gateway_service.go | 3 +
.../internal/service/openai_ws_forwarder.go | 266 ++++++++--
...penai_ws_forwarder_ingress_session_test.go | 155 ++++++
.../internal/service/openai_ws_http_bridge.go | 387 +++++++++++++++
.../service/openai_ws_http_bridge_test.go | 461 ++++++++++++++++++
8 files changed, 1257 insertions(+), 44 deletions(-)
create mode 100644 backend/internal/service/openai_ws_http_bridge.go
create mode 100644 backend/internal/service/openai_ws_http_bridge_test.go
diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go
index df9dcefc..0df7d09e 100644
--- a/backend/internal/config/config.go
+++ b/backend/internal/config/config.go
@@ -885,6 +885,12 @@ type GatewayOpenAIWSConfig struct {
StoreDisabledForceNewConn bool `mapstructure:"store_disabled_force_new_conn"`
// PrewarmGenerateEnabled: 是否启用 WSv2 generate=false 预热(默认 false)
PrewarmGenerateEnabled bool `mapstructure:"prewarm_generate_enabled"`
+ // ClientReadLimitBytes: 入站客户端 WS 单帧读取上限。
+ ClientReadLimitBytes int64 `mapstructure:"client_read_limit_bytes"`
+ // HTTPBridgeEnabled: 首包过大时,保持客户端 WS,改用 HTTP Responses 上游。
+ HTTPBridgeEnabled bool `mapstructure:"http_bridge_enabled"`
+ // HTTPBridgeThresholdBytes: 触发 HTTP bridge 的入站 WS payload 阈值。
+ HTTPBridgeThresholdBytes int64 `mapstructure:"http_bridge_threshold_bytes"`
// Feature 开关:v2 优先于 v1
ResponsesWebsockets bool `mapstructure:"responses_websockets"`
@@ -1806,6 +1812,9 @@ func setDefaults() {
viper.SetDefault("gateway.openai_ws.store_disabled_conn_mode", "strict")
viper.SetDefault("gateway.openai_ws.store_disabled_force_new_conn", true)
viper.SetDefault("gateway.openai_ws.prewarm_generate_enabled", false)
+ viper.SetDefault("gateway.openai_ws.client_read_limit_bytes", 64*1024*1024)
+ viper.SetDefault("gateway.openai_ws.http_bridge_enabled", true)
+ viper.SetDefault("gateway.openai_ws.http_bridge_threshold_bytes", 15*1024*1024)
viper.SetDefault("gateway.openai_ws.responses_websockets", false)
viper.SetDefault("gateway.openai_ws.responses_websockets_v2", true)
viper.SetDefault("gateway.openai_ws.max_conns_per_account", 128)
@@ -2543,6 +2552,15 @@ func (c *Config) Validate() error {
if c.Gateway.OpenAIWS.PrewarmCooldownMS < 0 {
return fmt.Errorf("gateway.openai_ws.prewarm_cooldown_ms must be non-negative")
}
+ if c.Gateway.OpenAIWS.ClientReadLimitBytes <= 0 {
+ return fmt.Errorf("gateway.openai_ws.client_read_limit_bytes must be positive")
+ }
+ if c.Gateway.OpenAIWS.HTTPBridgeThresholdBytes < 0 {
+ return fmt.Errorf("gateway.openai_ws.http_bridge_threshold_bytes must be non-negative")
+ }
+ if c.Gateway.OpenAIWS.HTTPBridgeEnabled && c.Gateway.OpenAIWS.HTTPBridgeThresholdBytes == 0 {
+ return fmt.Errorf("gateway.openai_ws.http_bridge_threshold_bytes must be positive when http_bridge_enabled is true")
+ }
if c.Gateway.OpenAIWS.FallbackCooldownSeconds < 0 {
return fmt.Errorf("gateway.openai_ws.fallback_cooldown_seconds must be non-negative")
}
diff --git a/backend/internal/config/config_test.go b/backend/internal/config/config_test.go
index 1eae5ed9..9478b510 100644
--- a/backend/internal/config/config_test.go
+++ b/backend/internal/config/config_test.go
@@ -134,6 +134,15 @@ func TestLoadDefaultOpenAIWSConfig(t *testing.T) {
if cfg.Gateway.OpenAIWS.PrewarmCooldownMS != 300 {
t.Fatalf("Gateway.OpenAIWS.PrewarmCooldownMS = %d, want 300", cfg.Gateway.OpenAIWS.PrewarmCooldownMS)
}
+ if cfg.Gateway.OpenAIWS.ClientReadLimitBytes != 64*1024*1024 {
+ t.Fatalf("Gateway.OpenAIWS.ClientReadLimitBytes = %d, want %d", cfg.Gateway.OpenAIWS.ClientReadLimitBytes, 64*1024*1024)
+ }
+ if !cfg.Gateway.OpenAIWS.HTTPBridgeEnabled {
+ t.Fatalf("Gateway.OpenAIWS.HTTPBridgeEnabled = false, want true")
+ }
+ if cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes != 15*1024*1024 {
+ t.Fatalf("Gateway.OpenAIWS.HTTPBridgeThresholdBytes = %d, want %d", cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes, 15*1024*1024)
+ }
if cfg.Gateway.OpenAIWS.RetryBackoffInitialMS != 120 {
t.Fatalf("Gateway.OpenAIWS.RetryBackoffInitialMS = %d, want 120", cfg.Gateway.OpenAIWS.RetryBackoffInitialMS)
}
diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go
index 0aa477b0..f3d4caf0 100644
--- a/backend/internal/handler/openai_gateway_handler.go
+++ b/backend/internal/handler/openai_gateway_handler.go
@@ -1167,7 +1167,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
defer func() {
_ = wsConn.CloseNow()
}()
- wsConn.SetReadLimit(16 * 1024 * 1024)
+ wsConn.SetReadLimit(service.ResolveOpenAIWSClientReadLimitBytes(h.cfg))
ctx := c.Request.Context()
readCtx, cancel := context.WithTimeout(ctx, 30*time.Second)
diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go
index 89ddaa7d..b1a95594 100644
--- a/backend/internal/service/openai_gateway_service.go
+++ b/backend/internal/service/openai_gateway_service.go
@@ -256,6 +256,9 @@ type OpenAIForwardResult struct {
ImageOutputSizes []string
ImageSizeSource string
ImageSizeBreakdown map[string]int
+
+ wsReplayInput []json.RawMessage
+ wsReplayInputExists bool
}
type OpenAIWSRetryMetricsSnapshot struct {
diff --git a/backend/internal/service/openai_ws_forwarder.go b/backend/internal/service/openai_ws_forwarder.go
index dd2f45fa..66c66134 100644
--- a/backend/internal/service/openai_ws_forwarder.go
+++ b/backend/internal/service/openai_ws_forwarder.go
@@ -1560,6 +1560,38 @@ func openAIWSRawItemsHasFunctionCallOutput(items []json.RawMessage) bool {
return false
}
+func openAIWSRawItemsHaveToolCallContextForOutputs(items []json.RawMessage) bool {
+ if len(items) == 0 {
+ return false
+ }
+ contextCallIDs := make(map[string]struct{})
+ outputCallIDs := make(map[string]struct{})
+ for _, item := range items {
+ itemType := gjson.GetBytes(item, "type").String()
+ callID := strings.TrimSpace(gjson.GetBytes(item, "call_id").String())
+ switch {
+ case isCodexToolCallContextItemType(itemType):
+ if callID != "" {
+ contextCallIDs[callID] = struct{}{}
+ }
+ case isCodexToolCallOutputItemType(itemType):
+ if callID == "" {
+ return false
+ }
+ outputCallIDs[callID] = struct{}{}
+ }
+ }
+ if len(outputCallIDs) == 0 || len(contextCallIDs) == 0 {
+ return false
+ }
+ for callID := range outputCallIDs {
+ if _, ok := contextCallIDs[callID]; !ok {
+ return false
+ }
+ }
+ return true
+}
+
func openAIWSRawPayloadHasToolCallOutput(payload []byte) bool {
if len(payload) == 0 {
return false
@@ -2664,6 +2696,27 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
}, nil
}
+ writeClientMessage := func(message []byte) error {
+ writeCtx, cancel := context.WithTimeout(ctx, s.openAIWSWriteTimeout())
+ defer cancel()
+ return clientConn.Write(writeCtx, coderws.MessageText, message)
+ }
+
+ readClientMessage := func() ([]byte, error) {
+ msgType, payload, readErr := clientConn.Read(ctx)
+ if readErr != nil {
+ return nil, readErr
+ }
+ if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
+ return nil, NewOpenAIWSClientCloseError(
+ coderws.StatusPolicyViolation,
+ fmt.Sprintf("unsupported websocket client message type: %s", msgType.String()),
+ nil,
+ )
+ }
+ return payload, nil
+ }
+
firstPayload, err := parseClientPayload(firstClientMessage)
if err != nil {
return err
@@ -2672,25 +2725,152 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
turnState := strings.TrimSpace(c.GetHeader(openAIWSTurnStateHeader))
stateStore := s.getOpenAIWSStateStore()
groupID := getOpenAIGroupIDFromContext(c)
- sessionHash := s.GenerateSessionHash(c, firstPayload.rawForHash)
- if turnState == "" && stateStore != nil && sessionHash != "" {
- if savedTurnState, ok := stateStore.GetSessionTurnState(groupID, sessionHash); ok {
- turnState = savedTurnState
- }
- }
-
- preferredConnID := ""
- if stateStore != nil && firstPayload.previousResponseID != "" {
- if connID, ok := stateStore.GetResponseConn(firstPayload.previousResponseID); ok {
- preferredConnID = connID
- }
- }
-
- storeDisabled := s.isOpenAIWSStoreDisabledInRequestRaw(firstPayload.payloadRaw, account)
storeDisabledConnMode := s.openAIWSStoreDisabledConnMode()
- if stateStore != nil && storeDisabled && firstPayload.previousResponseID == "" && sessionHash != "" {
- if connID, ok := stateStore.GetSessionConn(groupID, sessionHash); ok {
- preferredConnID = connID
+ sessionHash := ""
+ preferredConnID := ""
+ storeDisabled := false
+ refreshIngressRouteState := func(payload openAIWSClientPayload) {
+ sessionHash = s.GenerateSessionHash(c, payload.rawForHash)
+ if turnState == "" && stateStore != nil && sessionHash != "" {
+ if savedTurnState, ok := stateStore.GetSessionTurnState(groupID, sessionHash); ok {
+ turnState = savedTurnState
+ }
+ }
+
+ preferredConnID = ""
+ if stateStore != nil && payload.previousResponseID != "" {
+ if connID, ok := stateStore.GetResponseConn(payload.previousResponseID); ok {
+ preferredConnID = connID
+ }
+ }
+
+ storeDisabled = s.isOpenAIWSStoreDisabledInRequestRaw(payload.payloadRaw, account)
+ if stateStore != nil && storeDisabled && payload.previousResponseID == "" && sessionHash != "" {
+ if connID, ok := stateStore.GetSessionConn(groupID, sessionHash); ok {
+ preferredConnID = connID
+ }
+ }
+ }
+ refreshIngressRouteState(firstPayload)
+
+ if s.shouldBridgeOpenAIWSHTTP(firstPayload.payloadBytes, firstPayload.previousResponseID) {
+ logOpenAIWSModeInfo(
+ "ingress_ws_http_bridge_start account_id=%d account_type=%s payload_bytes=%d threshold_bytes=%d has_session_hash=%v store_disabled=%v",
+ account.ID,
+ account.Type,
+ firstPayload.payloadBytes,
+ s.openAIWSHTTPBridgeThresholdBytes(),
+ sessionHash != "",
+ storeDisabled,
+ )
+ currentBridgePayload := firstPayload
+ var bridgeReplayInput []json.RawMessage
+ bridgeReplayInputExists := false
+ for turn := 1; ; turn++ {
+ if turn > 1 && hooks != nil && hooks.BeforeRequest != nil {
+ if err := hooks.BeforeRequest(turn, currentBridgePayload.payloadRaw, currentBridgePayload.originalModel); err != nil {
+ return err
+ }
+ }
+ if hooks != nil && hooks.BeforeTurn != nil {
+ if err := hooks.BeforeTurn(turn); err != nil {
+ return err
+ }
+ }
+ if turnState != "" && c != nil && c.Request != nil {
+ c.Request.Header.Set(openAIWSTurnStateHeader, turnState)
+ }
+ bridgePayloadRaw := currentBridgePayload.payloadRaw
+ bridgePayloadBytes := currentBridgePayload.payloadBytes
+ needsBridgeReplay := currentBridgePayload.previousResponseID != "" || openAIWSRawPayloadHasToolCallOutput(currentBridgePayload.payloadRaw)
+ turnReplayInput, turnReplayInputExists, replayInputErr := buildOpenAIWSReplayInputSequence(
+ bridgeReplayInput,
+ bridgeReplayInputExists,
+ currentBridgePayload.payloadRaw,
+ needsBridgeReplay,
+ )
+ if replayInputErr != nil {
+ return fmt.Errorf("build websocket http bridge replay input: %w", replayInputErr)
+ }
+ if needsBridgeReplay && turnReplayInputExists {
+ updatedPayload, setInputErr := setOpenAIWSPayloadInputSequence(
+ currentBridgePayload.payloadRaw,
+ turnReplayInput,
+ true,
+ )
+ if setInputErr != nil {
+ return fmt.Errorf("set websocket http bridge replay input: %w", setInputErr)
+ }
+ bridgePayloadRaw = updatedPayload
+ bridgePayloadBytes = len(updatedPayload)
+ logOpenAIWSModeInfo(
+ "ingress_ws_http_bridge_replay_input account_id=%d turn=%d input_items=%d previous_response_id_present=%v has_tool_output=%v",
+ account.ID,
+ turn,
+ len(turnReplayInput),
+ currentBridgePayload.previousResponseID != "",
+ openAIWSRawPayloadHasToolCallOutput(currentBridgePayload.payloadRaw),
+ )
+ }
+ result, bridgeErr := s.proxyOpenAIWSHTTPBridgeTurn(
+ ctx,
+ c,
+ account,
+ token,
+ bridgePayloadRaw,
+ bridgePayloadBytes,
+ currentBridgePayload.originalModel,
+ currentBridgePayload.imageBillingModel,
+ currentBridgePayload.imageSizeTier,
+ currentBridgePayload.imageInputSize,
+ turn,
+ writeClientMessage,
+ )
+ if hooks != nil && hooks.AfterTurn != nil {
+ hooks.AfterTurn(turn, result, bridgeErr)
+ }
+ if bridgeErr != nil {
+ return bridgeErr
+ }
+ if result == nil {
+ return errors.New("websocket http bridge turn result is nil")
+ }
+ bridgeReplayInput = cloneOpenAIWSRawMessages(turnReplayInput)
+ bridgeReplayInputExists = turnReplayInputExists
+ if result.wsReplayInputExists {
+ bridgeReplayInput = append(bridgeReplayInput, cloneOpenAIWSRawMessages(result.wsReplayInput)...)
+ bridgeReplayInputExists = true
+ }
+ if bridgeTurnState := strings.TrimSpace(result.ResponseHeaders.Get(openAIWSTurnStateHeader)); bridgeTurnState != "" {
+ turnState = bridgeTurnState
+ if stateStore != nil && sessionHash != "" {
+ stateStore.BindSessionTurnState(groupID, sessionHash, bridgeTurnState, s.openAIWSSessionStickyTTL())
+ }
+ }
+ responseID := strings.TrimSpace(result.RequestID)
+ if responseID != "" && stateStore != nil {
+ ttl := s.openAIWSResponseStickyTTL()
+ logOpenAIWSBindResponseAccountWarn(groupID, account.ID, responseID, stateStore.BindResponseAccount(ctx, groupID, responseID, account.ID, ttl))
+ }
+ nextClientMessage, readErr := readClientMessage()
+ if readErr != nil {
+ if isOpenAIWSClientDisconnectError(readErr) {
+ closeStatus, closeReason := summarizeOpenAIWSReadCloseError(readErr)
+ logOpenAIWSModeInfo(
+ "ingress_ws_http_bridge_client_closed account_id=%d close_status=%s close_reason=%s",
+ account.ID,
+ closeStatus,
+ truncateOpenAIWSLogValue(closeReason, openAIWSHeaderValueMaxLen),
+ )
+ return nil
+ }
+ return fmt.Errorf("read client websocket request: %w", readErr)
+ }
+ nextPayload, parseErr := parseClientPayload(nextClientMessage)
+ if parseErr != nil {
+ return parseErr
+ }
+ currentBridgePayload = nextPayload
}
}
@@ -2844,27 +3024,6 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
return lease, nil
}
- writeClientMessage := func(message []byte) error {
- writeCtx, cancel := context.WithTimeout(ctx, s.openAIWSWriteTimeout())
- defer cancel()
- return clientConn.Write(writeCtx, coderws.MessageText, message)
- }
-
- readClientMessage := func() ([]byte, error) {
- msgType, payload, readErr := clientConn.Read(ctx)
- if readErr != nil {
- return nil, readErr
- }
- if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
- return nil, NewOpenAIWSClientCloseError(
- coderws.StatusPolicyViolation,
- fmt.Sprintf("unsupported websocket client message type: %s", msgType.String()),
- nil,
- )
- }
- return payload, nil
- }
-
sendAndRelay := func(turn int, lease *openAIWSConnLease, payload []byte, payloadBytes int, originalModel string, imageBillingModel string, imageSizeTier string, imageInputSize string) (*OpenAIForwardResult, error) {
if lease == nil {
return nil, errors.New("upstream websocket lease is nil")
@@ -2901,6 +3060,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
eventCount := 0
tokenEventCount := 0
terminalEventCount := 0
+ replayCollector := &openAIWSToolCallReplayCollector{}
firstEventType := ""
lastEventType := ""
needModelReplace := false
@@ -3031,6 +3191,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
upstreamMessage = corrected
}
}
+ replayCollector.AddEvent(eventType, upstreamMessage)
if err := writeClientMessage(upstreamMessage); err != nil {
if isOpenAIWSClientDisconnectError(err) {
clientDisconnected = true
@@ -3094,6 +3255,10 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
Duration: time.Since(turnStart),
FirstTokenMs: firstTokenMs,
}
+ if replayInput := replayCollector.Items(); len(replayInput) > 0 {
+ result.wsReplayInput = replayInput
+ result.wsReplayInputExists = true
+ }
if imageCount > 0 {
result.ImageCount = imageCount
result.ImageSize = imageSizeTier
@@ -3487,9 +3652,12 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
if forcePreferredConn {
// 携带 function_call_output 的请求不能丢弃 previous_response_id:
// 上游 API 需要 response chain 来匹配 tool_result 与之前的 tool_use,
- // 丢弃后会导致 "No tool call found for function call output" 400 错误。
+ // 除非 replay input 已经包含与每个 tool_result 匹配的 tool_use 上下文。
hasFCOutput := hasFunctionCallOutput
- if !turnPrevRecoveryTried && currentPreviousResponseID != "" && !hasFCOutput {
+ hasReplayToolContext := hasFCOutput &&
+ currentTurnReplayInputExists &&
+ openAIWSRawItemsHaveToolCallContextForOutputs(currentTurnReplayInput)
+ if !turnPrevRecoveryTried && currentPreviousResponseID != "" && (!hasFCOutput || hasReplayToolContext) {
updatedPayload, removed, dropErr := dropPreviousResponseIDFromRawPayload(currentPayload)
if dropErr != nil || !removed {
reason := "not_removed"
@@ -3521,11 +3689,13 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
)
} else {
logOpenAIWSModeInfo(
- "ingress_ws_preflight_ping_recovery account_id=%d turn=%d conn_id=%s action=drop_previous_response_id_retry previous_response_id=%s",
+ "ingress_ws_preflight_ping_recovery account_id=%d turn=%d conn_id=%s action=drop_previous_response_id_retry previous_response_id=%s has_function_call_output=%v has_replay_tool_context=%v",
account.ID,
turn,
truncateOpenAIWSLogValue(sessionConnID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(currentPreviousResponseID, openAIWSIDValueMaxLen),
+ hasFCOutput,
+ hasReplayToolContext,
)
turnPrevRecoveryTried = true
currentPayload = updatedWithInput
@@ -3537,12 +3707,18 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
}
}
if hasFCOutput && currentPreviousResponseID != "" {
+ reason := "function_call_output_missing_replay_context"
+ if hasReplayToolContext {
+ reason = "function_call_output_replay_not_applied"
+ }
logOpenAIWSModeInfo(
- "ingress_ws_preflight_ping_recovery_skip account_id=%d turn=%d conn_id=%s reason=function_call_output action=fail_close previous_response_id=%s",
+ "ingress_ws_preflight_ping_recovery_skip account_id=%d turn=%d conn_id=%s reason=%s action=fail_close previous_response_id=%s has_replay_tool_context=%v",
account.ID,
turn,
truncateOpenAIWSLogValue(sessionConnID, openAIWSIDValueMaxLen),
+ reason,
truncateOpenAIWSLogValue(currentPreviousResponseID, openAIWSIDValueMaxLen),
+ hasReplayToolContext,
)
}
resetSessionLease(true)
@@ -3622,6 +3798,10 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
lastTurnPayload = cloneOpenAIWSPayloadBytes(currentPayload)
lastTurnReplayInput = cloneOpenAIWSRawMessages(currentTurnReplayInput)
lastTurnReplayInputExists = currentTurnReplayInputExists
+ if result.wsReplayInputExists {
+ lastTurnReplayInput = append(lastTurnReplayInput, cloneOpenAIWSRawMessages(result.wsReplayInput)...)
+ lastTurnReplayInputExists = true
+ }
nextStrictState, strictStateErr := buildOpenAIWSIngressPreviousTurnStrictState(currentPayload)
if strictStateErr != nil {
lastTurnStrictState = nil
diff --git a/backend/internal/service/openai_ws_forwarder_ingress_session_test.go b/backend/internal/service/openai_ws_forwarder_ingress_session_test.go
index b7f1bc4f..069a9dee 100644
--- a/backend/internal/service/openai_ws_forwarder_ingress_session_test.go
+++ b/backend/internal/service/openai_ws_forwarder_ingress_session_test.go
@@ -2305,6 +2305,161 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_StoreDisabledStr
require.Equal(t, "world", gjson.Get(secondWrite, "input.1.text").String())
}
+func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_StoreDisabledPreflightPingFailReplaysFunctionCallOutputWithContext(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ prevPreflightPingIdle := openAIWSIngressPreflightPingIdle
+ openAIWSIngressPreflightPingIdle = 0
+ defer func() {
+ openAIWSIngressPreflightPingIdle = prevPreflightPingIdle
+ }()
+
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ cfg.Security.URLAllowlist.AllowInsecureHTTP = true
+ cfg.Gateway.OpenAIWS.Enabled = true
+ cfg.Gateway.OpenAIWS.OAuthEnabled = true
+ cfg.Gateway.OpenAIWS.APIKeyEnabled = true
+ cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
+ cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 2
+ cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0
+ cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 2
+ cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8
+ cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
+
+ firstConn := &openAIWSPreflightFailConn{
+ events: [][]byte{
+ []byte(`{"type":"response.completed","response":{"id":"resp_turn_ping_replay_ctx_1","model":"gpt-5.1","output":[{"type":"function_call","id":"fc_replay_1","call_id":"call_replay_1","name":"shell","arguments":"{}"}],"usage":{"input_tokens":1,"output_tokens":1}}}`),
+ },
+ }
+ secondConn := &openAIWSCaptureConn{
+ events: [][]byte{
+ []byte(`{"type":"response.completed","response":{"id":"resp_turn_ping_replay_ctx_2","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`),
+ },
+ }
+ dialer := &openAIWSQueueDialer{
+ conns: []openAIWSClientConn{firstConn, secondConn},
+ }
+ pool := newOpenAIWSConnPool(cfg)
+ pool.setClientDialerForTest(dialer)
+
+ svc := &OpenAIGatewayService{
+ cfg: cfg,
+ httpUpstream: &httpUpstreamRecorder{},
+ cache: &stubGatewayCache{},
+ openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
+ toolCorrector: NewCodexToolCorrector(),
+ openaiWSPool: pool,
+ }
+
+ account := &Account{
+ ID: 128,
+ Name: "openai-ingress-preflight-replay-function-output-with-context",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ },
+ Extra: map[string]any{
+ "responses_websockets_v2_enabled": true,
+ },
+ }
+
+ serverErrCh := make(chan error, 1)
+ wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{
+ CompressionMode: coderws.CompressionContextTakeover,
+ })
+ if err != nil {
+ serverErrCh <- err
+ return
+ }
+ defer func() {
+ _ = conn.CloseNow()
+ }()
+
+ rec := httptest.NewRecorder()
+ ginCtx, _ := gin.CreateTestContext(rec)
+ req := r.Clone(r.Context())
+ req.Header = req.Header.Clone()
+ req.Header.Set("User-Agent", "unit-test-agent/1.0")
+ ginCtx.Request = req
+
+ readCtx, cancel := context.WithTimeout(r.Context(), 3*time.Second)
+ msgType, firstMessage, readErr := conn.Read(readCtx)
+ cancel()
+ if readErr != nil {
+ serverErrCh <- readErr
+ return
+ }
+ if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
+ serverErrCh <- errors.New("unsupported websocket client message type")
+ return
+ }
+
+ serverErrCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "sk-test", firstMessage, nil)
+ }))
+ defer wsServer.Close()
+
+ dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
+ clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
+ cancelDial()
+ require.NoError(t, err)
+ defer func() {
+ _ = clientConn.CloseNow()
+ }()
+
+ writeMessage := func(payload string) {
+ writeCtx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
+ defer cancel()
+ require.NoError(t, clientConn.Write(writeCtx, coderws.MessageText, []byte(payload)))
+ }
+ readMessage := func() []byte {
+ readCtx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
+ defer cancel()
+ msgType, message, readErr := clientConn.Read(readCtx)
+ require.NoError(t, readErr)
+ require.Equal(t, coderws.MessageText, msgType)
+ return message
+ }
+
+ writeMessage(`{"type":"response.create","model":"gpt-5.1","stream":false,"store":false,"input":[{"type":"message","role":"user","content":[{"type":"input_text","text":"call tool"}]}]}`)
+ firstTurn := readMessage()
+ require.Equal(t, "resp_turn_ping_replay_ctx_1", gjson.GetBytes(firstTurn, "response.id").String())
+
+ writeMessage(`{"type":"response.create","model":"gpt-5.1","stream":false,"store":false,"previous_response_id":"resp_turn_ping_replay_ctx_1","input":[{"type":"function_call_output","call_id":"call_replay_1","output":"ok"}]}`)
+ secondTurn := readMessage()
+ require.Equal(t, "resp_turn_ping_replay_ctx_2", gjson.GetBytes(secondTurn, "response.id").String())
+
+ require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done"))
+ select {
+ case serverErr := <-serverErrCh:
+ require.NoError(t, serverErr)
+ case <-time.After(5 * time.Second):
+ t.Fatal("等待 ingress websocket function_call_output 自包含重放后结束超时")
+ }
+
+ require.Equal(t, 2, dialer.DialCount(), "带完整 tool 上下文的 function_call_output 应在 ping 失败后换新连接重放")
+ require.Equal(t, 1, firstConn.WriteCount())
+ require.GreaterOrEqual(t, firstConn.PingCount(), 1)
+ secondConn.mu.Lock()
+ secondWrites := append([]map[string]any(nil), secondConn.writes...)
+ secondConn.mu.Unlock()
+ require.Len(t, secondWrites, 1)
+ secondWrite := requestToJSONString(secondWrites[0])
+ require.False(t, gjson.Get(secondWrite, "previous_response_id").Exists())
+ require.Equal(t, 3, len(gjson.Get(secondWrite, "input").Array()))
+ require.Equal(t, "message", gjson.Get(secondWrite, "input.0.type").String())
+ require.Equal(t, "function_call", gjson.Get(secondWrite, "input.1.type").String())
+ require.Equal(t, "call_replay_1", gjson.Get(secondWrite, "input.1.call_id").String())
+ require.Equal(t, "function_call_output", gjson.Get(secondWrite, "input.2.type").String())
+ require.Equal(t, "call_replay_1", gjson.Get(secondWrite, "input.2.call_id").String())
+}
+
func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_StoreDisabledPreflightPingFailClosesWhenFunctionCallOutputNeedsPreviousResponseID(t *testing.T) {
gin.SetMode(gin.TestMode)
prevPreflightPingIdle := openAIWSIngressPreflightPingIdle
diff --git a/backend/internal/service/openai_ws_http_bridge.go b/backend/internal/service/openai_ws_http_bridge.go
new file mode 100644
index 00000000..1f0f32a0
--- /dev/null
+++ b/backend/internal/service/openai_ws_http_bridge.go
@@ -0,0 +1,387 @@
+package service
+
+import (
+ "bufio"
+ "context"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "io"
+ "net/http"
+ "strings"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/config"
+ "github.com/gin-gonic/gin"
+ "github.com/tidwall/gjson"
+)
+
+const (
+ openAIWSClientReadLimitBytesDefault int64 = 64 * 1024 * 1024
+ openAIWSHTTPBridgeThresholdBytesDefault int64 = 15 * 1024 * 1024
+ openAIWSHTTPBridgeErrorBodyLimitBytes = 64 * 1024
+)
+
+func ResolveOpenAIWSClientReadLimitBytes(cfg *config.Config) int64 {
+ if cfg == nil || cfg.Gateway.OpenAIWS.ClientReadLimitBytes <= 0 {
+ return openAIWSClientReadLimitBytesDefault
+ }
+ return cfg.Gateway.OpenAIWS.ClientReadLimitBytes
+}
+
+func (s *OpenAIGatewayService) openAIWSHTTPBridgeEnabled() bool {
+ return s != nil && s.cfg != nil && s.cfg.Gateway.OpenAIWS.HTTPBridgeEnabled
+}
+
+func (s *OpenAIGatewayService) openAIWSHTTPBridgeThresholdBytes() int64 {
+ if s == nil || s.cfg == nil || s.cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes <= 0 {
+ return openAIWSHTTPBridgeThresholdBytesDefault
+ }
+ return s.cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes
+}
+
+func (s *OpenAIGatewayService) shouldBridgeOpenAIWSHTTP(payloadBytes int, previousResponseID string) bool {
+ if !s.openAIWSHTTPBridgeEnabled() {
+ return false
+ }
+ if strings.TrimSpace(previousResponseID) != "" {
+ return false
+ }
+ threshold := s.openAIWSHTTPBridgeThresholdBytes()
+ return threshold > 0 && int64(payloadBytes) >= threshold
+}
+
+func prepareOpenAIWSHTTPBridgeBody(payload []byte) ([]byte, error) {
+ var body map[string]any
+ if err := json.Unmarshal(payload, &body); err != nil {
+ return nil, err
+ }
+ if body == nil {
+ return nil, errors.New("response.create payload must be a JSON object")
+ }
+ delete(body, "type")
+ delete(body, "generate")
+ delete(body, "previous_response_id")
+ body["stream"] = true
+ return json.Marshal(body)
+}
+
+type openAIWSToolCallReplayCollector struct {
+ items []json.RawMessage
+ seen map[string]struct{}
+}
+
+func (c *openAIWSToolCallReplayCollector) AddEvent(eventType string, message []byte) {
+ switch strings.TrimSpace(eventType) {
+ case "response.output_item.done":
+ c.addItem(gjson.GetBytes(message, "item"))
+ case "response.completed", "response.done":
+ output := gjson.GetBytes(message, "response.output")
+ if !output.IsArray() {
+ return
+ }
+ for _, item := range output.Array() {
+ c.addItem(item)
+ }
+ }
+}
+
+func (c *openAIWSToolCallReplayCollector) Items() []json.RawMessage {
+ return cloneOpenAIWSRawMessages(c.items)
+}
+
+func (c *openAIWSToolCallReplayCollector) addItem(item gjson.Result) {
+ if !item.Exists() || item.Type != gjson.JSON {
+ return
+ }
+ raw := strings.TrimSpace(item.Raw)
+ if raw == "" || !strings.HasPrefix(raw, "{") {
+ return
+ }
+ if !isCodexToolCallContextItemType(item.Get("type").String()) {
+ return
+ }
+ key := strings.TrimSpace(item.Get("id").String())
+ if key == "" {
+ key = strings.TrimSpace(item.Get("call_id").String())
+ }
+ if key == "" {
+ key = raw
+ }
+ if c.seen == nil {
+ c.seen = make(map[string]struct{})
+ }
+ if _, ok := c.seen[key]; ok {
+ return
+ }
+ c.seen[key] = struct{}{}
+ c.items = append(c.items, json.RawMessage(raw))
+}
+
+func buildOpenAIWSHTTPBridgeErrorEvent(statusCode int, message string) []byte {
+ message = strings.TrimSpace(message)
+ if message == "" {
+ message = http.StatusText(statusCode)
+ }
+ if message == "" {
+ message = "upstream request failed"
+ }
+ event := map[string]any{
+ "type": "error",
+ "status": statusCode,
+ "error": map[string]any{
+ "type": "upstream_error",
+ "message": message,
+ },
+ }
+ body, err := json.Marshal(event)
+ if err != nil {
+ return []byte(`{"type":"error","error":{"type":"upstream_error","message":"upstream request failed"}}`)
+ }
+ return body
+}
+
+func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
+ ctx context.Context,
+ c *gin.Context,
+ account *Account,
+ token string,
+ payload []byte,
+ payloadBytes int,
+ originalModel string,
+ imageBillingModel string,
+ imageSizeTier string,
+ imageInputSize string,
+ turn int,
+ writeClientMessage func([]byte) error,
+) (*OpenAIForwardResult, error) {
+ if s == nil {
+ return nil, errors.New("service is nil")
+ }
+ if s.httpUpstream == nil {
+ return nil, errors.New("openai http upstream is nil")
+ }
+ if account == nil {
+ return nil, errors.New("account is nil")
+ }
+ if writeClientMessage == nil {
+ return nil, errors.New("client websocket writer is nil")
+ }
+
+ body, err := prepareOpenAIWSHTTPBridgeBody(payload)
+ if err != nil {
+ return nil, fmt.Errorf("prepare http bridge body: %w", err)
+ }
+
+ upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
+ upstreamReq, err := s.buildUpstreamRequestOpenAIPassthrough(upstreamCtx, c, account, body, token)
+ releaseUpstreamCtx()
+ if err != nil {
+ return nil, err
+ }
+
+ proxyURL := ""
+ if account.ProxyID != nil && account.Proxy != nil {
+ proxyURL = account.Proxy.URL()
+ }
+ if c != nil {
+ c.Set("openai_passthrough", true)
+ c.Set("openai_ws_http_bridge", true)
+ }
+
+ turnStart := time.Now()
+ resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency)
+ if err != nil {
+ safeErr := sanitizeUpstreamErrorMessage(err.Error())
+ _ = writeClientMessage(buildOpenAIWSHTTPBridgeErrorEvent(http.StatusBadGateway, "Upstream request failed"))
+ return nil, fmt.Errorf("upstream http bridge request failed: %s", safeErr)
+ }
+ defer func() { _ = resp.Body.Close() }()
+
+ if resp.StatusCode >= 400 {
+ respBody, _ := io.ReadAll(io.LimitReader(resp.Body, openAIWSHTTPBridgeErrorBodyLimitBytes))
+ upstreamMsg := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(respBody)))
+ if upstreamMsg == "" {
+ upstreamMsg = http.StatusText(resp.StatusCode)
+ }
+ _ = writeClientMessage(buildOpenAIWSHTTPBridgeErrorEvent(resp.StatusCode, upstreamMsg))
+ return nil, fmt.Errorf("upstream http bridge error: status=%d message=%s", resp.StatusCode, upstreamMsg)
+ }
+
+ responseID := ""
+ usage := OpenAIUsage{}
+ imageCounter := newOpenAIImageOutputCounter()
+ var firstTokenMs *int
+ reqStream := openAIWSPayloadBoolFromRaw(body, "stream", true)
+ eventCount := 0
+ tokenEventCount := 0
+ terminalEventCount := 0
+ replayCollector := &openAIWSToolCallReplayCollector{}
+ firstEventType := ""
+ lastEventType := ""
+ sawDone := false
+ wroteDownstream := false
+ clientDisconnected := false
+ mappedModel := ""
+ needModelReplace := false
+ var mappedModelBytes []byte
+ if originalModel != "" {
+ mappedModel = normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel))
+ needModelReplace = mappedModel != "" && mappedModel != originalModel
+ if needModelReplace {
+ mappedModelBytes = []byte(mappedModel)
+ }
+ }
+
+ resultWithUsage := func() *OpenAIForwardResult {
+ imageCount := imageCounter.Count()
+ result := &OpenAIForwardResult{
+ RequestID: responseID,
+ Usage: usage,
+ Model: originalModel,
+ UpstreamModel: mappedModel,
+ ServiceTier: extractOpenAIServiceTierFromBody(body),
+ ReasoningEffort: extractOpenAIReasoningEffortFromBody(body, originalModel),
+ Stream: reqStream,
+ OpenAIWSMode: true,
+ ResponseHeaders: cloneHeader(resp.Header),
+ Duration: time.Since(turnStart),
+ FirstTokenMs: firstTokenMs,
+ }
+ if replayInput := replayCollector.Items(); len(replayInput) > 0 {
+ result.wsReplayInput = replayInput
+ result.wsReplayInputExists = true
+ }
+ if imageCount > 0 {
+ result.ImageCount = imageCount
+ result.ImageSize = imageSizeTier
+ result.ImageInputSize = imageInputSize
+ result.ImageOutputSizes = imageCounter.Sizes()
+ result.BillingModel = imageBillingModel
+ }
+ return result
+ }
+
+ scanner := bufio.NewScanner(resp.Body)
+ maxLineSize := defaultMaxLineSize
+ if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 {
+ maxLineSize = s.cfg.Gateway.MaxLineSize
+ }
+ scanBuf := getSSEScannerBuf64K()
+ scanner.Buffer(scanBuf[:0], maxLineSize)
+ defer putSSEScannerBuf64K(scanBuf)
+
+ for scanner.Scan() {
+ line := scanner.Text()
+ data, ok := extractOpenAISSEDataLine(line)
+ if !ok {
+ continue
+ }
+ trimmedData := strings.TrimSpace(data)
+ if trimmedData == "" {
+ continue
+ }
+ if trimmedData == "[DONE]" {
+ sawDone = true
+ continue
+ }
+
+ upstreamMessage := []byte(trimmedData)
+ eventType, eventResponseID, _ := parseOpenAIWSEventEnvelope(upstreamMessage)
+ if responseID == "" && eventResponseID != "" {
+ responseID = eventResponseID
+ }
+ if eventType != "" {
+ eventCount++
+ if firstEventType == "" {
+ firstEventType = eventType
+ }
+ lastEventType = eventType
+ }
+ if isOpenAIWSTokenEvent(eventType) {
+ tokenEventCount++
+ if firstTokenMs == nil {
+ ms := int(time.Since(turnStart).Milliseconds())
+ firstTokenMs = &ms
+ }
+ }
+ if openAIWSEventShouldParseUsage(eventType) {
+ parseOpenAIWSResponseUsageFromCompletedEvent(upstreamMessage, &usage)
+ }
+ imageCounter.AddSSEData(upstreamMessage)
+
+ if needModelReplace && len(mappedModelBytes) > 0 && openAIWSEventMayContainModel(eventType) && strings.Contains(trimmedData, mappedModel) {
+ upstreamMessage = replaceOpenAIWSMessageModel(upstreamMessage, mappedModel, originalModel)
+ }
+ if s.toolCorrector != nil && openAIWSEventMayContainToolCalls(eventType) && openAIWSMessageLikelyContainsToolCalls(upstreamMessage) {
+ if corrected, changed := s.toolCorrector.CorrectToolCallsInSSEBytes(upstreamMessage); changed {
+ upstreamMessage = corrected
+ }
+ }
+ replayCollector.AddEvent(eventType, upstreamMessage)
+
+ if !clientDisconnected {
+ if err := writeClientMessage(upstreamMessage); err != nil {
+ if isOpenAIWSClientDisconnectError(err) {
+ clientDisconnected = true
+ closeStatus, closeReason := summarizeOpenAIWSReadCloseError(err)
+ logOpenAIWSModeInfo(
+ "ingress_ws_http_bridge_client_disconnected_drain account_id=%d turn=%d close_status=%s close_reason=%s",
+ account.ID,
+ turn,
+ closeStatus,
+ truncateOpenAIWSLogValue(closeReason, openAIWSHeaderValueMaxLen),
+ )
+ } else {
+ return nil, wrapOpenAIWSIngressTurnError(
+ "write_client",
+ fmt.Errorf("write client websocket event: %w", err),
+ wroteDownstream,
+ )
+ }
+ } else {
+ wroteDownstream = true
+ }
+ }
+
+ if eventType == "error" {
+ errCodeRaw, errTypeRaw, errMsgRaw := parseOpenAIWSErrorEventFields(upstreamMessage)
+ s.persistOpenAIWSRateLimitSignal(ctx, account, resp.Header, upstreamMessage, errCodeRaw, errTypeRaw, errMsgRaw)
+ errMessage := strings.TrimSpace(errMsgRaw)
+ if errMessage == "" {
+ errMessage = "upstream error event"
+ }
+ return resultWithUsage(), errors.New(errMessage)
+ }
+ if isOpenAIWSTerminalEvent(eventType) {
+ terminalEventCount++
+ firstTokenMsValue := -1
+ if firstTokenMs != nil {
+ firstTokenMsValue = *firstTokenMs
+ }
+ logOpenAIWSModeInfo(
+ "ingress_ws_http_bridge_turn_completed account_id=%d turn=%d response_id=%s payload_bytes=%d duration_ms=%d events=%d token_events=%d terminal_events=%d first_event=%s last_event=%s first_token_ms=%d client_disconnected=%v",
+ account.ID,
+ turn,
+ truncateOpenAIWSLogValue(responseID, openAIWSIDValueMaxLen),
+ payloadBytes,
+ time.Since(turnStart).Milliseconds(),
+ eventCount,
+ tokenEventCount,
+ terminalEventCount,
+ truncateOpenAIWSLogValue(firstEventType, openAIWSLogValueMaxLen),
+ truncateOpenAIWSLogValue(lastEventType, openAIWSLogValueMaxLen),
+ firstTokenMsValue,
+ clientDisconnected,
+ )
+ return resultWithUsage(), nil
+ }
+ }
+ if err := scanner.Err(); err != nil {
+ return resultWithUsage(), fmt.Errorf("read upstream http bridge stream: %w", err)
+ }
+ if sawDone && eventCount > 0 {
+ return resultWithUsage(), nil
+ }
+ return resultWithUsage(), errors.New("upstream http bridge stream ended before terminal event")
+}
diff --git a/backend/internal/service/openai_ws_http_bridge_test.go b/backend/internal/service/openai_ws_http_bridge_test.go
new file mode 100644
index 00000000..0a1d6b56
--- /dev/null
+++ b/backend/internal/service/openai_ws_http_bridge_test.go
@@ -0,0 +1,461 @@
+package service
+
+import (
+ "context"
+ "io"
+ "net/http"
+ "net/http/httptest"
+ "strings"
+ "testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/config"
+ coderws "github.com/coder/websocket"
+ "github.com/gin-gonic/gin"
+ "github.com/stretchr/testify/require"
+ "github.com/tidwall/gjson"
+)
+
+func TestPrepareOpenAIWSHTTPBridgeBodyStripsWSFields(t *testing.T) {
+ body, err := prepareOpenAIWSHTTPBridgeBody([]byte(`{"type":"response.create","generate":true,"model":"gpt-5","stream":false,"previous_response_id":"resp_prev","input":"hi"}`))
+ require.NoError(t, err)
+ require.False(t, gjson.GetBytes(body, "type").Exists())
+ require.False(t, gjson.GetBytes(body, "generate").Exists())
+ require.False(t, gjson.GetBytes(body, "previous_response_id").Exists())
+ require.Equal(t, "gpt-5", gjson.GetBytes(body, "model").String())
+ require.True(t, gjson.GetBytes(body, "stream").Bool())
+ require.Equal(t, "hi", gjson.GetBytes(body, "input").String())
+}
+
+func TestOpenAIWSHTTPBridgeDecisionKeepsSmallFramesOnWS(t *testing.T) {
+ svc := &OpenAIGatewayService{
+ cfg: &config.Config{
+ Gateway: config.GatewayConfig{
+ OpenAIWS: config.GatewayOpenAIWSConfig{
+ HTTPBridgeEnabled: true,
+ HTTPBridgeThresholdBytes: 100,
+ },
+ },
+ },
+ }
+
+ require.False(t, svc.shouldBridgeOpenAIWSHTTP(99, ""))
+ require.True(t, svc.shouldBridgeOpenAIWSHTTP(100, ""))
+ require.False(t, svc.shouldBridgeOpenAIWSHTTP(1000, "resp_existing"))
+
+ svc.cfg.Gateway.OpenAIWS.HTTPBridgeEnabled = false
+ require.False(t, svc.shouldBridgeOpenAIWSHTTP(1000, ""))
+}
+
+func TestOpenAIWSHTTPBridgeRelaysSSEFramesAsWebSocketMessages(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ sseBody := strings.Join([]string{
+ `data: {"type":"response.created","response":{"id":"resp_bridge","model":"gpt-5"}}`,
+ "",
+ `data: {"type":"response.output_text.delta","response":{"id":"resp_bridge"},"delta":"ok"}`,
+ "",
+ `data: {"type":"response.completed","response":{"id":"resp_bridge","model":"gpt-5","usage":{"input_tokens":3,"output_tokens":2}}}`,
+ "",
+ }, "\n")
+ upstream := &httpUpstreamRecorder{resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{
+ "Content-Type": []string{"text/event-stream"},
+ "x-request-id": []string{"rid_bridge"},
+ },
+ Body: io.NopCloser(strings.NewReader(sseBody)),
+ }}
+ svc := &OpenAIGatewayService{
+ cfg: &config.Config{
+ Gateway: config.GatewayConfig{
+ MaxLineSize: defaultMaxLineSize,
+ OpenAIWS: config.GatewayOpenAIWSConfig{
+ HTTPBridgeEnabled: true,
+ HTTPBridgeThresholdBytes: 1,
+ },
+ },
+ },
+ httpUpstream: upstream,
+ toolCorrector: NewCodexToolCorrector(),
+ }
+ account := &Account{
+ ID: 7,
+ Name: "api-key",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Status: StatusActive,
+ }
+ payload := []byte(`{"type":"response.create","generate":true,"model":"gpt-5","stream":true,"input":"hi"}`)
+
+ type bridgeResult struct {
+ result *OpenAIForwardResult
+ err error
+ }
+ resultCh := make(chan bridgeResult, 1)
+ wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
+ if err != nil {
+ resultCh <- bridgeResult{err: err}
+ return
+ }
+ defer func() { _ = conn.CloseNow() }()
+
+ rec := httptest.NewRecorder()
+ ginCtx, _ := gin.CreateTestContext(rec)
+ req := r.Clone(r.Context())
+ req.Header = req.Header.Clone()
+ ginCtx.Request = req
+
+ writeClient := func(message []byte) error {
+ writeCtx, cancel := context.WithTimeout(r.Context(), 3*time.Second)
+ defer cancel()
+ return conn.Write(writeCtx, coderws.MessageText, message)
+ }
+ result, bridgeErr := svc.proxyOpenAIWSHTTPBridgeTurn(
+ r.Context(),
+ ginCtx,
+ account,
+ "sk-test",
+ payload,
+ len(payload),
+ "gpt-5",
+ "",
+ "",
+ "",
+ 1,
+ writeClient,
+ )
+ resultCh <- bridgeResult{result: result, err: bridgeErr}
+ }))
+ defer wsServer.Close()
+
+ dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
+ clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
+ cancelDial()
+ require.NoError(t, err)
+ defer func() { _ = clientConn.CloseNow() }()
+
+ readEvent := func() []byte {
+ readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
+ msgType, event, readErr := clientConn.Read(readCtx)
+ cancelRead()
+ require.NoError(t, readErr)
+ require.Equal(t, coderws.MessageText, msgType)
+ return event
+ }
+
+ created := readEvent()
+ delta := readEvent()
+ completed := readEvent()
+
+ require.Equal(t, "response.created", gjson.GetBytes(created, "type").String())
+ require.Equal(t, "response.output_text.delta", gjson.GetBytes(delta, "type").String())
+ require.Equal(t, "response.completed", gjson.GetBytes(completed, "type").String())
+
+ select {
+ case bridge := <-resultCh:
+ require.NoError(t, bridge.err)
+ require.NotNil(t, bridge.result)
+ require.Equal(t, "resp_bridge", bridge.result.RequestID)
+ require.Equal(t, 3, bridge.result.Usage.InputTokens)
+ require.Equal(t, 2, bridge.result.Usage.OutputTokens)
+ require.True(t, bridge.result.OpenAIWSMode)
+ case <-time.After(3 * time.Second):
+ t.Fatal("timed out waiting for bridge result")
+ }
+
+ require.NotNil(t, upstream.lastReq)
+ require.Equal(t, http.MethodPost, upstream.lastReq.Method)
+ require.False(t, gjson.GetBytes(upstream.lastBody, "type").Exists())
+ require.False(t, gjson.GetBytes(upstream.lastBody, "generate").Exists())
+ require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool())
+}
+
+func TestOpenAIWSHTTPBridgeAcceptsFirstFrameAboveLegacy16MiB(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ sseBody := strings.Join([]string{
+ `data: {"type":"response.created","response":{"id":"resp_large_bridge","model":"gpt-5"}}`,
+ "",
+ `data: {"type":"response.completed","response":{"id":"resp_large_bridge","model":"gpt-5","usage":{"input_tokens":9,"output_tokens":1}}}`,
+ "",
+ }, "\n")
+ upstream := &httpUpstreamRecorder{resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{
+ "Content-Type": []string{"text/event-stream"},
+ "x-request-id": []string{"rid_large_bridge"},
+ },
+ Body: io.NopCloser(strings.NewReader(sseBody)),
+ }}
+ cfg := &config.Config{
+ Gateway: config.GatewayConfig{
+ MaxLineSize: defaultMaxLineSize,
+ OpenAIWS: config.GatewayOpenAIWSConfig{
+ Enabled: true,
+ APIKeyEnabled: true,
+ ResponsesWebsocketsV2: true,
+ ClientReadLimitBytes: 64 * 1024 * 1024,
+ HTTPBridgeEnabled: true,
+ HTTPBridgeThresholdBytes: 15 * 1024 * 1024,
+ },
+ },
+ }
+ svc := &OpenAIGatewayService{
+ cfg: cfg,
+ httpUpstream: upstream,
+ toolCorrector: NewCodexToolCorrector(),
+ }
+ account := &Account{
+ ID: 9,
+ Name: "api-key",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Credentials: map[string]any{"api_key": "sk-upstream"},
+ Extra: map[string]any{
+ "openai_apikey_responses_websockets_v2_enabled": true,
+ },
+ Concurrency: 1,
+ Status: StatusActive,
+ }
+
+ payload := []byte(`{"type":"response.create","generate":true,"model":"gpt-5","stream":true,"input":"` + strings.Repeat("x", 17*1024*1024) + `"}`)
+ require.Greater(t, len(payload), 16*1024*1024)
+ require.Less(t, int64(len(payload)), ResolveOpenAIWSClientReadLimitBytes(cfg))
+
+ errCh := make(chan error, 1)
+ wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
+ if err != nil {
+ errCh <- err
+ return
+ }
+ defer func() { _ = conn.CloseNow() }()
+ conn.SetReadLimit(ResolveOpenAIWSClientReadLimitBytes(cfg))
+
+ readCtx, cancelRead := context.WithTimeout(r.Context(), 10*time.Second)
+ msgType, firstMessage, err := conn.Read(readCtx)
+ cancelRead()
+ if err != nil {
+ errCh <- err
+ return
+ }
+ if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
+ errCh <- NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "unexpected client websocket message type", nil)
+ return
+ }
+
+ rec := httptest.NewRecorder()
+ ginCtx, _ := gin.CreateTestContext(rec)
+ req := r.Clone(r.Context())
+ req.Header = req.Header.Clone()
+ req.Header.Set("User-Agent", "codex_cli_rs/0.135.0")
+ ginCtx.Request = req
+
+ proxyCtx, cancelProxy := context.WithTimeout(r.Context(), 20*time.Second)
+ defer cancelProxy()
+ errCh <- svc.ProxyResponsesWebSocketFromClient(proxyCtx, ginCtx, conn, account, "sk-test", firstMessage, nil)
+ }))
+ defer wsServer.Close()
+
+ dialCtx, cancelDial := context.WithTimeout(context.Background(), 5*time.Second)
+ clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
+ cancelDial()
+ require.NoError(t, err)
+ defer func() { _ = clientConn.CloseNow() }()
+
+ writeCtx, cancelWrite := context.WithTimeout(context.Background(), 20*time.Second)
+ err = clientConn.Write(writeCtx, coderws.MessageText, payload)
+ cancelWrite()
+ require.NoError(t, err)
+
+ var eventTypes []string
+ for {
+ readCtx, cancelRead := context.WithTimeout(context.Background(), 10*time.Second)
+ msgType, event, readErr := clientConn.Read(readCtx)
+ cancelRead()
+ require.NoError(t, readErr)
+ require.Equal(t, coderws.MessageText, msgType)
+
+ eventType := gjson.GetBytes(event, "type").String()
+ eventTypes = append(eventTypes, eventType)
+ if eventType == "response.completed" {
+ break
+ }
+ }
+ require.Contains(t, eventTypes, "response.created")
+ require.Contains(t, eventTypes, "response.completed")
+
+ require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done"))
+ select {
+ case proxyErr := <-errCh:
+ require.NoError(t, proxyErr)
+ case <-time.After(10 * time.Second):
+ t.Fatal("timed out waiting for websocket bridge proxy to finish")
+ }
+
+ require.NotNil(t, upstream.lastReq)
+ require.Equal(t, http.MethodPost, upstream.lastReq.Method)
+ require.Greater(t, len(upstream.lastBody), 16*1024*1024)
+ require.False(t, gjson.GetBytes(upstream.lastBody, "type").Exists())
+ require.False(t, gjson.GetBytes(upstream.lastBody, "generate").Exists())
+ require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool())
+ require.Equal(t, "gpt-5", gjson.GetBytes(upstream.lastBody, "model").String())
+}
+
+func TestOpenAIWSHTTPBridgeKeepsContinuationFramesOnHTTPWithoutPreviousResponseID(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ firstSSEBody := strings.Join([]string{
+ `data: {"type":"response.completed","response":{"id":"resp_bridge_first","model":"gpt-5.1","output":[{"type":"function_call","id":"fc_bridge_1","call_id":"call_bridge_1","name":"shell","arguments":"{}"}],"usage":{"input_tokens":9,"output_tokens":1}}}`,
+ "",
+ }, "\n")
+ secondSSEBody := strings.Join([]string{
+ `data: {"type":"response.completed","response":{"id":"resp_bridge_second","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`,
+ "",
+ }, "\n")
+ upstream := &httpUpstreamRecorder{responses: []*http.Response{
+ {
+ StatusCode: http.StatusOK,
+ Header: http.Header{
+ "Content-Type": []string{"text/event-stream"},
+ },
+ Body: io.NopCloser(strings.NewReader(firstSSEBody)),
+ },
+ {
+ StatusCode: http.StatusOK,
+ Header: http.Header{
+ "Content-Type": []string{"text/event-stream"},
+ },
+ Body: io.NopCloser(strings.NewReader(secondSSEBody)),
+ },
+ }}
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ cfg.Security.URLAllowlist.AllowInsecureHTTP = true
+ cfg.Gateway.OpenAIWS.Enabled = true
+ cfg.Gateway.OpenAIWS.OAuthEnabled = true
+ cfg.Gateway.OpenAIWS.APIKeyEnabled = true
+ cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
+ cfg.Gateway.OpenAIWS.HTTPBridgeEnabled = true
+ cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes = 1
+ cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1
+ cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0
+ cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1
+ cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8
+ cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
+
+ captureConn := &openAIWSCaptureConn{}
+ captureDialer := &openAIWSCaptureDialer{conn: captureConn}
+ pool := newOpenAIWSConnPool(cfg)
+ pool.setClientDialerForTest(captureDialer)
+
+ svc := &OpenAIGatewayService{
+ cfg: cfg,
+ httpUpstream: upstream,
+ cache: &stubGatewayCache{},
+ openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
+ toolCorrector: NewCodexToolCorrector(),
+ openaiWSPool: pool,
+ }
+ account := &Account{
+ ID: 19,
+ Name: "api-key-bridge-handoff",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Credentials: map[string]any{"api_key": "sk-upstream"},
+ Extra: map[string]any{
+ "responses_websockets_v2_enabled": true,
+ },
+ Concurrency: 1,
+ Status: StatusActive,
+ Schedulable: true,
+ }
+
+ errCh := make(chan error, 1)
+ wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
+ if err != nil {
+ errCh <- err
+ return
+ }
+ defer func() { _ = conn.CloseNow() }()
+
+ readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second)
+ msgType, firstMessage, err := conn.Read(readCtx)
+ cancelRead()
+ if err != nil {
+ errCh <- err
+ return
+ }
+ if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
+ errCh <- NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "unexpected client websocket message type", nil)
+ return
+ }
+
+ rec := httptest.NewRecorder()
+ ginCtx, _ := gin.CreateTestContext(rec)
+ req := r.Clone(r.Context())
+ req.Header = req.Header.Clone()
+ req.Header.Set("User-Agent", "codex_cli_rs/0.135.0")
+ ginCtx.Request = req
+
+ errCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "sk-test", firstMessage, nil)
+ }))
+ defer wsServer.Close()
+
+ dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
+ clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
+ cancelDial()
+ require.NoError(t, err)
+ defer func() { _ = clientConn.CloseNow() }()
+
+ writeMessage := func(payload string) {
+ writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
+ defer cancelWrite()
+ require.NoError(t, clientConn.Write(writeCtx, coderws.MessageText, []byte(payload)))
+ }
+ readMessage := func() []byte {
+ readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
+ defer cancelRead()
+ msgType, event, readErr := clientConn.Read(readCtx)
+ require.NoError(t, readErr)
+ require.Equal(t, coderws.MessageText, msgType)
+ return event
+ }
+
+ writeMessage(`{"type":"response.create","model":"gpt-5.1","stream":true,"input":"first"}`)
+ firstTurnEvent := readMessage()
+ require.Equal(t, "response.completed", gjson.GetBytes(firstTurnEvent, "type").String())
+ require.Equal(t, "resp_bridge_first", gjson.GetBytes(firstTurnEvent, "response.id").String())
+
+ writeMessage(`{"type":"response.create","model":"gpt-5.1","stream":false,"previous_response_id":"resp_bridge_first","input":[{"type":"function_call_output","call_id":"call_bridge_1","output":"ok"}]}`)
+ secondTurnEvent := readMessage()
+ require.Equal(t, "response.completed", gjson.GetBytes(secondTurnEvent, "type").String())
+ require.Equal(t, "resp_bridge_second", gjson.GetBytes(secondTurnEvent, "response.id").String())
+
+ require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done"))
+ select {
+ case proxyErr := <-errCh:
+ require.NoError(t, proxyErr)
+ case <-time.After(5 * time.Second):
+ t.Fatal("timed out waiting for websocket bridge proxy to finish")
+ }
+
+ require.Len(t, upstream.bodies, 2, "进入 HTTP bridge 后同一客户端 WS 连接内应保持 HTTP/SSE bridge")
+ require.False(t, gjson.GetBytes(upstream.bodies[0], "previous_response_id").Exists())
+ require.False(t, gjson.GetBytes(upstream.bodies[1], "previous_response_id").Exists())
+ secondInput := gjson.GetBytes(upstream.bodies[1], "input").Array()
+ require.Len(t, secondInput, 3)
+ require.Equal(t, "first", secondInput[0].String())
+ require.Equal(t, "function_call", secondInput[1].Get("type").String())
+ require.Equal(t, "call_bridge_1", secondInput[1].Get("call_id").String())
+ require.Equal(t, "function_call_output", secondInput[2].Get("type").String())
+ require.Equal(t, "call_bridge_1", secondInput[2].Get("call_id").String())
+ require.Equal(t, 0, captureDialer.DialCount())
+ require.Empty(t, captureConn.writes)
+}
From 2a075a85bd6746bd3d2093f56cd713d041d16663 Mon Sep 17 00:00:00 2001
From: Kanshan03 <92105202+Kanshan03@users.noreply.github.com>
Date: Mon, 1 Jun 2026 11:25:28 +0800
Subject: [PATCH 66/79] =?UTF-8?q?fix(openai):=20=E6=B3=A8=E5=85=A5=20WS=20?=
=?UTF-8?q?Codex=20=E7=94=9F=E5=9B=BE=E6=A1=A5=E6=8E=A5=E5=B7=A5=E5=85=B7?=
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit
---
.../internal/service/openai_ws_forwarder.go | 32 ++++-
...penai_ws_forwarder_ingress_session_test.go | 136 ++++++++++++++++++
2 files changed, 166 insertions(+), 2 deletions(-)
diff --git a/backend/internal/service/openai_ws_forwarder.go b/backend/internal/service/openai_ws_forwarder.go
index dd2f45fa..22ed214b 100644
--- a/backend/internal/service/openai_ws_forwarder.go
+++ b/backend/internal/service/openai_ws_forwarder.go
@@ -2477,6 +2477,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
wsPath = normalizeOpenAIWSLogValue(parsedURL.Path)
}
debugEnabled := isOpenAIWSModeDebugEnabled()
+ isCodexCLI := openai.IsCodexOfficialClientByHeaders(c.GetHeader("User-Agent"), c.GetHeader("originator")) || (s.cfg != nil && s.cfg.Gateway.ForceCodexCLI)
type openAIWSClientPayload struct {
payloadRaw []byte
@@ -2586,6 +2587,34 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
}
normalized = next
}
+ apiKey := getAPIKeyFromContext(c)
+ imageGenerationAllowed := GroupAllowsImageGeneration(apiKeyGroup(apiKey))
+ codexBridgeEnabled := isCodexCLI && imageGenerationAllowed && s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey)
+ if codexBridgeEnabled {
+ payloadMap := make(map[string]any)
+ if err := json.Unmarshal(normalized, &payloadMap); err != nil {
+ return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", err)
+ }
+ bridgeModified := false
+ if ensureOpenAIResponsesImageGenerationTool(payloadMap) {
+ bridgeModified = true
+ logOpenAIWSModeInfo("ingress_ws_codex_image_tool_injected account_id=%d", account.ID)
+ }
+ if normalizeOpenAIResponsesImageGenerationTools(payloadMap) {
+ bridgeModified = true
+ }
+ if applyCodexImageGenerationBridgeInstructions(payloadMap) {
+ bridgeModified = true
+ logOpenAIWSModeInfo("ingress_ws_codex_image_bridge_instructions_added account_id=%d", account.ID)
+ }
+ if bridgeModified {
+ rebuilt, marshalErr := json.Marshal(payloadMap)
+ if marshalErr != nil {
+ return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", marshalErr)
+ }
+ normalized = rebuilt
+ }
+ }
upstreamModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel))
if modelMissing || upstreamModel != originalModel {
next, setErr := applyPayloadMutation(normalized, "model", upstreamModel)
@@ -2595,7 +2624,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
normalized = next
}
imageIntent := IsImageGenerationIntent(openAIResponsesEndpoint, originalModel, normalized)
- if imageIntent && !GroupAllowsImageGeneration(apiKeyGroup(getAPIKeyFromContext(c))) {
+ if imageIntent && !imageGenerationAllowed {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, ImageGenerationPermissionMessage(), nil)
}
imageBillingModel := ""
@@ -2694,7 +2723,6 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
}
}
- isCodexCLI := openai.IsCodexOfficialClientByHeaders(c.GetHeader("User-Agent"), c.GetHeader("originator")) || (s.cfg != nil && s.cfg.Gateway.ForceCodexCLI)
wsHeaders, _ := s.buildOpenAIWSHeaders(c, account, token, wsDecision, isCodexCLI, turnState, strings.TrimSpace(c.GetHeader(openAIWSTurnMetadataHeader)), firstPayload.promptCacheKey)
baseAcquireReq := openAIWSAcquireRequest{
Account: account,
diff --git a/backend/internal/service/openai_ws_forwarder_ingress_session_test.go b/backend/internal/service/openai_ws_forwarder_ingress_session_test.go
index b7f1bc4f..5e4b70c2 100644
--- a/backend/internal/service/openai_ws_forwarder_ingress_session_test.go
+++ b/backend/internal/service/openai_ws_forwarder_ingress_session_test.go
@@ -298,6 +298,142 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_FollowupCreateCa
require.Equal(t, "resp_omit_model_1", gjson.Get(requestToJSONString(captureConn.writes[1]), "previous_response_id").String())
}
+func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_InjectsCodexImageBridge(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ cfg.Security.URLAllowlist.AllowInsecureHTTP = true
+ cfg.Gateway.OpenAIWS.Enabled = true
+ cfg.Gateway.OpenAIWS.OAuthEnabled = true
+ cfg.Gateway.OpenAIWS.APIKeyEnabled = true
+ cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
+ cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1
+ cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0
+ cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1
+ cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8
+ cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
+
+ captureConn := &openAIWSCaptureConn{
+ events: [][]byte{
+ []byte(`{"type":"response.completed","response":{"id":"resp_codex_image_bridge","model":"gpt-5.5","usage":{"input_tokens":1,"output_tokens":1}}}`),
+ },
+ }
+ captureDialer := &openAIWSCaptureDialer{conn: captureConn}
+ pool := newOpenAIWSConnPool(cfg)
+ pool.setClientDialerForTest(captureDialer)
+
+ svc := &OpenAIGatewayService{
+ cfg: cfg,
+ httpUpstream: &httpUpstreamRecorder{},
+ cache: &stubGatewayCache{},
+ openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
+ toolCorrector: NewCodexToolCorrector(),
+ openaiWSPool: pool,
+ }
+
+ groupID := int64(3)
+ apiKey := &APIKey{
+ ID: 1,
+ UserID: 1,
+ GroupID: &groupID,
+ Group: &Group{
+ ID: groupID,
+ AllowImageGeneration: true,
+ },
+ }
+ account := &Account{
+ ID: 31,
+ Name: "openai-codex-image-ws",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "access_token": "test-token",
+ },
+ Extra: map[string]any{
+ "openai_oauth_responses_websockets_v2_enabled": true,
+ "codex_image_generation_bridge": true,
+ },
+ }
+
+ serverErrCh := make(chan error, 1)
+ wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{
+ CompressionMode: coderws.CompressionContextTakeover,
+ })
+ if err != nil {
+ serverErrCh <- err
+ return
+ }
+ defer func() {
+ _ = conn.CloseNow()
+ }()
+
+ rec := httptest.NewRecorder()
+ ginCtx, _ := gin.CreateTestContext(rec)
+ req := r.Clone(r.Context())
+ req.Header = req.Header.Clone()
+ req.Header.Set("User-Agent", "codex_cli_rs/0.98.0")
+ ginCtx.Request = req
+ ginCtx.Set("api_key", apiKey)
+
+ readCtx, cancel := context.WithTimeout(r.Context(), 3*time.Second)
+ msgType, firstMessage, readErr := conn.Read(readCtx)
+ cancel()
+ if readErr != nil {
+ serverErrCh <- readErr
+ return
+ }
+ if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
+ serverErrCh <- errors.New("unsupported websocket client message type")
+ return
+ }
+
+ serverErrCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "test-token", firstMessage, nil)
+ }))
+ defer wsServer.Close()
+
+ dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
+ clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
+ cancelDial()
+ require.NoError(t, err)
+ defer func() {
+ _ = clientConn.CloseNow()
+ }()
+
+ writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
+ err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"gpt-5.5","stream":false,"input":"draw a cat"}`))
+ cancelWrite()
+ require.NoError(t, err)
+
+ readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
+ msgType, message, err := clientConn.Read(readCtx)
+ cancelRead()
+ require.NoError(t, err)
+ require.Equal(t, coderws.MessageText, msgType)
+ require.Equal(t, "resp_codex_image_bridge", gjson.GetBytes(message, "response.id").String())
+
+ _ = clientConn.Close(coderws.StatusNormalClosure, "done")
+
+ select {
+ case serverErr := <-serverErrCh:
+ require.NoError(t, serverErr)
+ case <-time.After(5 * time.Second):
+ t.Fatal("等待 ingress websocket 结束超时")
+ }
+
+ require.Len(t, captureConn.writes, 1)
+ upstreamPayload := requestToJSONString(captureConn.writes[0])
+ require.True(t, gjson.Get(upstreamPayload, `tools.#(type=="image_generation")`).Exists())
+ require.Equal(t, "png", gjson.Get(upstreamPayload, `tools.#(type=="image_generation").output_format`).String())
+ require.Contains(t, gjson.Get(upstreamPayload, "instructions").String(), "image_generation")
+}
+
func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_DedicatedModeDoesNotReuseConnAcrossSessions(t *testing.T) {
gin.SetMode(gin.TestMode)
From 003b2786dacfd9c1c5342c499ae57590979fef8d Mon Sep 17 00:00:00 2001
From: visa2
Date: Mon, 1 Jun 2026 12:03:15 +0800
Subject: [PATCH 67/79] test(apicompat): check type assertions in responses
stream wire tests
errcheck (check-type-assertions) flagged unchecked single-value type
assertions; switch to the comma-ok form so golangci-lint passes.
Co-Authored-By: Claude Opus 4.8
---
.../responses_stream_event_wire_test.go | 16 ++++++++++------
1 file changed, 10 insertions(+), 6 deletions(-)
diff --git a/backend/internal/pkg/apicompat/responses_stream_event_wire_test.go b/backend/internal/pkg/apicompat/responses_stream_event_wire_test.go
index fbef45af..b4f6871d 100644
--- a/backend/internal/pkg/apicompat/responses_stream_event_wire_test.go
+++ b/backend/internal/pkg/apicompat/responses_stream_event_wire_test.go
@@ -43,7 +43,8 @@ func TestWire_FunctionCallItemAlwaysComplete(t *testing.T) {
OutputIndex: 1,
Item: &ResponsesOutput{Type: "function_call", ID: "fc_1", CallID: "call_a", Name: "exec", Status: "in_progress"},
})
- item := added["item"].(map[string]any)
+ item, ok := added["item"].(map[string]any)
+ require.True(t, ok, "item must be an object")
for _, k := range []string{"call_id", "name", "arguments"} {
require.Containsf(t, item, k, "function_call item missing %q", k)
}
@@ -57,9 +58,10 @@ func TestWire_MessageItemContentAlwaysArray(t *testing.T) {
OutputIndex: 0,
Item: &ResponsesOutput{Type: "message", ID: "msg_1", Role: "assistant", Status: "in_progress"},
})
- item := m["item"].(map[string]any)
+ item, ok := m["item"].(map[string]any)
+ require.True(t, ok, "item must be an object")
require.Contains(t, item, "content")
- _, ok := item["content"].([]any)
+ _, ok = item["content"].([]any)
require.True(t, ok, "content must be an array")
}
@@ -70,9 +72,10 @@ func TestWire_ReasoningItemSummaryAlwaysArray(t *testing.T) {
OutputIndex: 0,
Item: &ResponsesOutput{Type: "reasoning", ID: "rs_1", Status: "in_progress"},
})
- item := m["item"].(map[string]any)
+ item, ok := m["item"].(map[string]any)
+ require.True(t, ok, "item must be an object")
require.Contains(t, item, "summary")
- _, ok := item["summary"].([]any)
+ _, ok = item["summary"].([]any)
require.True(t, ok, "summary must be an array")
}
@@ -82,7 +85,8 @@ func TestWire_ContentPartCarriesAnnotationsLogprobs(t *testing.T) {
Type: "response.content_part.added", OutputIndex: 0, ContentIndex: 0, ItemID: "msg_1",
Part: &ResponsesContentPart{Type: "output_text", Text: ""},
})
- part := m["part"].(map[string]any)
+ part, ok := m["part"].(map[string]any)
+ require.True(t, ok, "part must be an object")
require.Equal(t, "output_text", part["type"])
require.Contains(t, part, "text")
require.Contains(t, part, "annotations")
From 04deb819b0e549271d443040b083b2972cbc56a1 Mon Sep 17 00:00:00 2001
From: wucm667
Date: Tue, 2 Jun 2026 14:59:18 +0800
Subject: [PATCH 68/79] fix(payment): use trade_status for EasyPay query
---
backend/internal/payment/provider/easypay.go | 51 ++++++-
.../payment/provider/easypay_query_test.go | 131 ++++++++++++++++++
2 files changed, 175 insertions(+), 7 deletions(-)
create mode 100644 backend/internal/payment/provider/easypay_query_test.go
diff --git a/backend/internal/payment/provider/easypay.go b/backend/internal/payment/provider/easypay.go
index e7d8aab9..32d6b7be 100644
--- a/backend/internal/payment/provider/easypay.go
+++ b/backend/internal/payment/provider/easypay.go
@@ -213,22 +213,59 @@ func (e *EasyPay) QueryOrder(ctx context.Context, tradeNo string) (*payment.Quer
if err != nil {
return nil, fmt.Errorf("easypay query: %w", err)
}
+ type easyPayQueryData struct {
+ TradeStatus *string `json:"trade_status"`
+ Status *int `json:"status"`
+ Money *string `json:"money"`
+ TradeNo *string `json:"trade_no"`
+ }
var resp struct {
- Code int `json:"code"`
- Msg string `json:"msg"`
- Status int `json:"status"`
- Money string `json:"money"`
+ Code int `json:"code"`
+ Msg string `json:"msg"`
+ TradeStatus *string `json:"trade_status"`
+ Status *int `json:"status"`
+ Money *string `json:"money"`
+ TradeNo *string `json:"trade_no"`
+ Data easyPayQueryData `json:"data"`
}
if err := json.Unmarshal(body, &resp); err != nil {
return nil, fmt.Errorf("easypay parse query: %w", err)
}
status := payment.ProviderStatusPending
- if resp.Status == easypayStatusPaid {
+ if resp.TradeStatus != nil {
+ if *resp.TradeStatus == tradeStatusSuccess {
+ status = payment.ProviderStatusPaid
+ }
+ } else if resp.Data.TradeStatus != nil {
+ if *resp.Data.TradeStatus == tradeStatusSuccess {
+ status = payment.ProviderStatusPaid
+ }
+ } else if resp.Status != nil {
+ if *resp.Status == easypayStatusPaid {
+ status = payment.ProviderStatusPaid
+ }
+ } else if resp.Data.Status != nil && *resp.Data.Status == easypayStatusPaid {
status = payment.ProviderStatusPaid
}
- amount, _ := strconv.ParseFloat(resp.Money, 64)
+
+ money := ""
+ if resp.Money != nil {
+ money = *resp.Money
+ } else if resp.Data.Money != nil {
+ money = *resp.Data.Money
+ }
+ responseTradeNo := tradeNo
+ if resp.TradeNo != nil {
+ if *resp.TradeNo != "" {
+ responseTradeNo = *resp.TradeNo
+ }
+ } else if resp.Data.TradeNo != nil && *resp.Data.TradeNo != "" {
+ responseTradeNo = *resp.Data.TradeNo
+ }
+
+ amount, _ := strconv.ParseFloat(money, 64)
return &payment.QueryOrderResponse{
- TradeNo: tradeNo,
+ TradeNo: responseTradeNo,
Status: status,
Amount: amount,
Metadata: e.MerchantIdentityMetadata(),
diff --git a/backend/internal/payment/provider/easypay_query_test.go b/backend/internal/payment/provider/easypay_query_test.go
new file mode 100644
index 00000000..5042a94d
--- /dev/null
+++ b/backend/internal/payment/provider/easypay_query_test.go
@@ -0,0 +1,131 @@
+package provider
+
+import (
+ "context"
+ "net/http"
+ "net/http/httptest"
+ "net/url"
+ "testing"
+
+ "github.com/Wei-Shaw/sub2api/internal/payment"
+)
+
+func TestEasyPayQueryOrderStatusMapping(t *testing.T) {
+ t.Parallel()
+
+ const orderID = "order-123"
+ tests := []struct {
+ name string
+ body string
+ wantStatus string
+ wantTradeNo string
+ wantAmount float64
+ }{
+ {
+ name: "top level trade success is paid",
+ body: `{"code":1,"trade_status":"TRADE_SUCCESS","status":0,"money":"12.34","trade_no":"gateway-123"}`,
+ wantStatus: payment.ProviderStatusPaid,
+ wantTradeNo: "gateway-123",
+ wantAmount: 12.34,
+ },
+ {
+ name: "waiting trade status with paid numeric status stays pending",
+ body: `{"code":1,"trade_status":"WAITING","status":1,"money":"12.34","trade_no":"gateway-123"}`,
+ wantStatus: payment.ProviderStatusPending,
+ wantTradeNo: "gateway-123",
+ wantAmount: 12.34,
+ },
+ {
+ name: "empty trade status with paid numeric status stays pending",
+ body: `{"code":1,"trade_status":"","status":1,"money":"12.34"}`,
+ wantStatus: payment.ProviderStatusPending,
+ wantTradeNo: orderID,
+ wantAmount: 12.34,
+ },
+ {
+ name: "nested data trade success is paid",
+ body: `{"code":1,"data":{"trade_status":"TRADE_SUCCESS","status":0,"money":"9.99","trade_no":"data-456"}}`,
+ wantStatus: payment.ProviderStatusPaid,
+ wantTradeNo: "data-456",
+ wantAmount: 9.99,
+ },
+ {
+ name: "legacy numeric paid status remains compatible",
+ body: `{"code":1,"status":1,"money":"3.21"}`,
+ wantStatus: payment.ProviderStatusPaid,
+ wantTradeNo: orderID,
+ wantAmount: 3.21,
+ },
+ {
+ name: "legacy numeric non paid status is pending",
+ body: `{"code":1,"status":0,"money":"3.21"}`,
+ wantStatus: payment.ProviderStatusPending,
+ wantTradeNo: orderID,
+ wantAmount: 3.21,
+ },
+ {
+ name: "query failure with missing status is pending",
+ body: `{"code":0,"msg":"订单不存在"}`,
+ wantStatus: payment.ProviderStatusPending,
+ wantTradeNo: orderID,
+ },
+ {
+ name: "missing fields are pending",
+ body: `{}`,
+ wantStatus: payment.ProviderStatusPending,
+ wantTradeNo: orderID,
+ },
+ }
+
+ for _, tt := range tests {
+ tt := tt
+ t.Run(tt.name, func(t *testing.T) {
+ t.Parallel()
+
+ var gotForm url.Values
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ t.Errorf("method = %q, want %q", r.Method, http.MethodPost)
+ }
+ if r.URL.Path != "/api.php" {
+ t.Errorf("path = %q, want /api.php", r.URL.Path)
+ }
+ if err := r.ParseForm(); err != nil {
+ t.Errorf("ParseForm: %v", err)
+ }
+ gotForm = make(url.Values, len(r.PostForm))
+ for key, values := range r.PostForm {
+ gotForm[key] = append([]string(nil), values...)
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write([]byte(tt.body))
+ }))
+ defer server.Close()
+
+ provider := newTestEasyPay(t, server.URL)
+ resp, err := provider.QueryOrder(context.Background(), orderID)
+ if err != nil {
+ t.Fatalf("QueryOrder returned error: %v", err)
+ }
+ if resp.Status != tt.wantStatus {
+ t.Fatalf("status = %q, want %q (response=%+v)", resp.Status, tt.wantStatus, resp)
+ }
+ if resp.TradeNo != tt.wantTradeNo {
+ t.Fatalf("trade_no = %q, want %q", resp.TradeNo, tt.wantTradeNo)
+ }
+ if resp.Amount != tt.wantAmount {
+ t.Fatalf("amount = %v, want %v", resp.Amount, tt.wantAmount)
+ }
+ for key, want := range map[string]string{
+ "act": "order",
+ "pid": "pid-1",
+ "key": "pkey-1",
+ "out_trade_no": orderID,
+ } {
+ if got := gotForm.Get(key); got != want {
+ t.Fatalf("form[%s] = %q, want %q (form=%v)", key, got, want, gotForm)
+ }
+ }
+ })
+ }
+}
From c40a74d983f85f0a0ccf49e1028cefb15dc10a83 Mon Sep 17 00:00:00 2001
From: wucm667
Date: Wed, 3 Jun 2026 09:33:37 +0800
Subject: [PATCH 69/79] fix(risk-control): exempt admins from moderation
auto-ban
---
.../internal/service/content_moderation.go | 5 ++
.../service/content_moderation_test.go | 90 +++++++++++++++++++
2 files changed, 95 insertions(+)
diff --git a/backend/internal/service/content_moderation.go b/backend/internal/service/content_moderation.go
index ee1fca41..42b909c9 100644
--- a/backend/internal/service/content_moderation.go
+++ b/backend/internal/service/content_moderation.go
@@ -1656,6 +1656,11 @@ func (s *ContentModerationService) applyFlaggedAccountSideEffects(ctx context.Co
slog.Warn("content_moderation.ban_get_user_failed", "user_id", *log.UserID, "error", err)
return false
}
+ if user.IsAdmin() {
+ slog.Warn("content_moderation.autoban_skipped_admin", "user_id", *log.UserID, "role", user.Role, "count", count, "threshold", cfg.BanThreshold)
+ // TODO: Disable the triggering API key instead when API key mutation is available here.
+ return false
+ }
if user.Status != StatusDisabled {
user.Status = StatusDisabled
if err := s.userRepo.Update(ctx, user); err != nil {
diff --git a/backend/internal/service/content_moderation_test.go b/backend/internal/service/content_moderation_test.go
index 6c6fef44..9cfdc1e4 100644
--- a/backend/internal/service/content_moderation_test.go
+++ b/backend/internal/service/content_moderation_test.go
@@ -1,9 +1,11 @@
package service
import (
+ "bytes"
"context"
"encoding/json"
"fmt"
+ "log/slog"
"net/http"
"net/http/httptest"
"strings"
@@ -1484,6 +1486,94 @@ func TestContentModerationCheck_HashBlockLogsDoNotIncreaseNextViolationCount(t *
require.Equal(t, 1, logs[1].ViolationCount)
}
+func TestContentModerationAutoBanSkipsAdminAccount(t *testing.T) {
+ var slogOutput bytes.Buffer
+ previousLogger := slog.Default()
+ slog.SetDefault(slog.New(slog.NewTextHandler(&slogOutput, nil)))
+ t.Cleanup(func() {
+ slog.SetDefault(previousLogger)
+ })
+
+ cfg := defaultContentModerationConfig()
+ cfg.BanThreshold = 2
+ cfg.ViolationWindowHours = 24
+
+ userID := int64(1001)
+ repo := &contentModerationTestRepo{}
+ require.NoError(t, repo.CreateLog(context.Background(), newContentModerationFlaggedLog(userID)))
+ userRepo := &contentModerationTestUserRepo{user: &User{ID: userID, Role: RoleAdmin, Status: StatusActive}}
+ invalidator := &contentModerationTestAuthCacheInvalidator{}
+ svc := NewContentModerationService(nil, repo, nil, nil, userRepo, invalidator, nil)
+
+ svc.persistContentModerationLog(context.Background(), cfg, newContentModerationFlaggedLog(userID), "", false, true)
+
+ logs := requireContentModerationLogCount(t, repo, 2)
+ require.Equal(t, 2, logs[1].ViolationCount)
+ require.False(t, logs[1].AutoBanned)
+ require.Equal(t, StatusActive, userRepo.user.Status)
+ require.Empty(t, userRepo.updated)
+ require.Empty(t, invalidator.userIDs)
+ require.Contains(t, slogOutput.String(), "content_moderation.autoban_skipped_admin")
+ require.Contains(t, slogOutput.String(), "user_id=1001")
+ require.Contains(t, slogOutput.String(), "role=admin")
+ require.Contains(t, slogOutput.String(), "count=2")
+ require.Contains(t, slogOutput.String(), "threshold=2")
+}
+
+func TestContentModerationAutoBanDisablesRegularUserAtThreshold(t *testing.T) {
+ cfg := defaultContentModerationConfig()
+ cfg.BanThreshold = 2
+ cfg.ViolationWindowHours = 24
+
+ userID := int64(1001)
+ repo := &contentModerationTestRepo{}
+ require.NoError(t, repo.CreateLog(context.Background(), newContentModerationFlaggedLog(userID)))
+ userRepo := &contentModerationTestUserRepo{user: &User{ID: userID, Role: RoleUser, Status: StatusActive}}
+ invalidator := &contentModerationTestAuthCacheInvalidator{}
+ svc := NewContentModerationService(nil, repo, nil, nil, userRepo, invalidator, nil)
+
+ svc.persistContentModerationLog(context.Background(), cfg, newContentModerationFlaggedLog(userID), "", false, true)
+
+ logs := requireContentModerationLogCount(t, repo, 2)
+ require.Equal(t, 2, logs[1].ViolationCount)
+ require.True(t, logs[1].AutoBanned)
+ require.Len(t, userRepo.updated, 1)
+ require.Equal(t, StatusDisabled, userRepo.user.Status)
+ require.Equal(t, []int64{userID}, invalidator.userIDs)
+}
+
+func TestContentModerationAdminBelowBanThresholdRecordsViolationOnly(t *testing.T) {
+ cfg := defaultContentModerationConfig()
+ cfg.BanThreshold = 2
+ cfg.ViolationWindowHours = 24
+
+ userID := int64(1001)
+ repo := &contentModerationTestRepo{}
+ userRepo := &contentModerationTestUserRepo{user: &User{ID: userID, Role: RoleAdmin, Status: StatusActive}}
+ invalidator := &contentModerationTestAuthCacheInvalidator{}
+ svc := NewContentModerationService(nil, repo, nil, nil, userRepo, invalidator, nil)
+
+ svc.persistContentModerationLog(context.Background(), cfg, newContentModerationFlaggedLog(userID), "", false, true)
+
+ logs := requireContentModerationLogCount(t, repo, 1)
+ require.Equal(t, 1, logs[0].ViolationCount)
+ require.False(t, logs[0].AutoBanned)
+ require.Equal(t, StatusActive, userRepo.user.Status)
+ require.Empty(t, userRepo.updated)
+ require.Empty(t, invalidator.userIDs)
+}
+
+func newContentModerationFlaggedLog(userID int64) *ContentModerationLog {
+ return &ContentModerationLog{
+ UserID: &userID,
+ Action: ContentModerationActionBlock,
+ Flagged: true,
+ HighestCategory: "sexual",
+ HighestScore: 0.9,
+ CreatedAt: time.Now(),
+ }
+}
+
func TestContentModerationCheck_PreBlockFlaggedWritesRedisHashCache(t *testing.T) {
requestCount := 0
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
From 134687782ce24866f585818f46df35f0aafb5f1c Mon Sep 17 00:00:00 2001
From: wucm667
Date: Wed, 3 Jun 2026 09:48:46 +0800
Subject: [PATCH 70/79] build(go): bump toolchain to 1.26.4
---
.github/workflows/backend-ci.yml | 4 ++--
.github/workflows/release.yml | 2 +-
.github/workflows/security-scan.yml | 2 +-
Dockerfile | 2 +-
backend/Dockerfile | 2 +-
backend/go.mod | 2 +-
deploy/Dockerfile | 2 +-
7 files changed, 8 insertions(+), 8 deletions(-)
diff --git a/.github/workflows/backend-ci.yml b/.github/workflows/backend-ci.yml
index 15ff97fe..fb4d0ce6 100644
--- a/.github/workflows/backend-ci.yml
+++ b/.github/workflows/backend-ci.yml
@@ -20,7 +20,7 @@ jobs:
cache-dependency-path: backend/go.sum
- name: Verify Go version
run: |
- go version | grep -q 'go1.26.3'
+ go version | grep -q 'go1.26.4'
- name: Unit tests
working-directory: backend
run: make test-unit
@@ -60,7 +60,7 @@ jobs:
cache-dependency-path: backend/go.sum
- name: Verify Go version
run: |
- go version | grep -q 'go1.26.3'
+ go version | grep -q 'go1.26.4'
- name: golangci-lint
uses: golangci/golangci-lint-action@v9
with:
diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml
index 80bc9850..7d48131a 100644
--- a/.github/workflows/release.yml
+++ b/.github/workflows/release.yml
@@ -115,7 +115,7 @@ jobs:
- name: Verify Go version
run: |
- go version | grep -q 'go1.26.3'
+ go version | grep -q 'go1.26.4'
# Docker setup for GoReleaser
- name: Set up QEMU
diff --git a/.github/workflows/security-scan.yml b/.github/workflows/security-scan.yml
index ef8e59e5..e102b5f8 100644
--- a/.github/workflows/security-scan.yml
+++ b/.github/workflows/security-scan.yml
@@ -23,7 +23,7 @@ jobs:
cache-dependency-path: backend/go.sum
- name: Verify Go version
run: |
- go version | grep -q 'go1.26.3'
+ go version | grep -q 'go1.26.4'
- name: Run govulncheck
working-directory: backend
run: |
diff --git a/Dockerfile b/Dockerfile
index d556008b..f9a03a2b 100644
--- a/Dockerfile
+++ b/Dockerfile
@@ -7,7 +7,7 @@
# =============================================================================
ARG NODE_IMAGE=node:24-alpine
-ARG GOLANG_IMAGE=golang:1.26.3-alpine
+ARG GOLANG_IMAGE=golang:1.26.4-alpine
ARG ALPINE_IMAGE=alpine:3.21
ARG POSTGRES_IMAGE=postgres:18-alpine
ARG GOPROXY=https://goproxy.cn,direct
diff --git a/backend/Dockerfile b/backend/Dockerfile
index f153d686..26b1dc33 100644
--- a/backend/Dockerfile
+++ b/backend/Dockerfile
@@ -1,4 +1,4 @@
-FROM golang:1.26.3-alpine
+FROM golang:1.26.4-alpine
WORKDIR /app
diff --git a/backend/go.mod b/backend/go.mod
index 587d5370..62be56c8 100644
--- a/backend/go.mod
+++ b/backend/go.mod
@@ -1,6 +1,6 @@
module github.com/Wei-Shaw/sub2api
-go 1.26.3
+go 1.26.4
require (
entgo.io/ent v0.14.5
diff --git a/deploy/Dockerfile b/deploy/Dockerfile
index a947158f..d39dd17d 100644
--- a/deploy/Dockerfile
+++ b/deploy/Dockerfile
@@ -7,7 +7,7 @@
# =============================================================================
ARG NODE_IMAGE=node:24-alpine
-ARG GOLANG_IMAGE=golang:1.26.3-alpine
+ARG GOLANG_IMAGE=golang:1.26.4-alpine
ARG ALPINE_IMAGE=alpine:3.20
ARG GOPROXY=https://goproxy.cn,direct
ARG GOSUMDB=sum.golang.google.cn
From a8ffb052ca5fcaf9a3e53eba4e83bedcf2a8e8a4 Mon Sep 17 00:00:00 2001
From: ghostg00 <28946120+ghostg00@users.noreply.github.com>
Date: Wed, 3 Jun 2026 11:37:55 +0800
Subject: [PATCH 71/79] =?UTF-8?q?Revert=20"fix(usage):=20=E4=BF=AE?=
=?UTF-8?q?=E6=AD=A3=20OpenAI=205h=20=E7=94=A8=E9=87=8F=E7=99=BE=E5=88=86?=
=?UTF-8?q?=E6=AF=94=E8=AF=AD=E4=B9=89"?=
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit
This reverts commit b65dde634bf7a5338f22b46ad0f9d98a197526fc.
---
.../account_test_service_openai_test.go | 4 +--
.../service/account_usage_service_test.go | 29 +---------------
.../service/openai_gateway_service.go | 17 ++--------
...nai_gateway_service_codex_snapshot_test.go | 34 -------------------
.../service/openai_gateway_service_test.go | 2 +-
.../service/ratelimit_service_openai_test.go | 32 ++++++++---------
6 files changed, 22 insertions(+), 96 deletions(-)
diff --git a/backend/internal/service/account_test_service_openai_test.go b/backend/internal/service/account_test_service_openai_test.go
index 970c723a..910567fb 100644
--- a/backend/internal/service/account_test_service_openai_test.go
+++ b/backend/internal/service/account_test_service_openai_test.go
@@ -132,7 +132,7 @@ func TestAccountTestService_OpenAISuccessPersistsSnapshotFromHeaders(t *testing.
require.Len(t, upstream.requests, 1)
require.Equal(t, HTTPUpstreamProfileOpenAI, HTTPUpstreamProfileFromContext(upstream.requests[0].Context()))
require.NotEmpty(t, repo.updatedExtra)
- require.Equal(t, 58.0, repo.updatedExtra["codex_5h_used_percent"])
+ require.Equal(t, 42.0, repo.updatedExtra["codex_5h_used_percent"])
require.Equal(t, 88.0, repo.updatedExtra["codex_7d_used_percent"])
require.Contains(t, recorder.Body.String(), "test_complete")
}
@@ -170,7 +170,7 @@ func TestAccountTestService_OpenAI429PersistsSnapshotAndRateLimitState(t *testin
resp.Header.Set("x-codex-primary-used-percent", "100")
resp.Header.Set("x-codex-primary-reset-after-seconds", "604800")
resp.Header.Set("x-codex-primary-window-minutes", "10080")
- resp.Header.Set("x-codex-secondary-used-percent", "0")
+ resp.Header.Set("x-codex-secondary-used-percent", "100")
resp.Header.Set("x-codex-secondary-reset-after-seconds", "18000")
resp.Header.Set("x-codex-secondary-window-minutes", "300")
diff --git a/backend/internal/service/account_usage_service_test.go b/backend/internal/service/account_usage_service_test.go
index 5f37aadb..e0390c4c 100644
--- a/backend/internal/service/account_usage_service_test.go
+++ b/backend/internal/service/account_usage_service_test.go
@@ -73,7 +73,7 @@ func TestExtractOpenAICodexProbeUpdatesAccepts429WithCodexHeaders(t *testing.T)
headers.Set("x-codex-primary-used-percent", "100")
headers.Set("x-codex-primary-reset-after-seconds", "604800")
headers.Set("x-codex-primary-window-minutes", "10080")
- headers.Set("x-codex-secondary-used-percent", "0")
+ headers.Set("x-codex-secondary-used-percent", "100")
headers.Set("x-codex-secondary-reset-after-seconds", "18000")
headers.Set("x-codex-secondary-window-minutes", "300")
@@ -92,33 +92,6 @@ func TestExtractOpenAICodexProbeUpdatesAccepts429WithCodexHeaders(t *testing.T)
}
}
-func TestBuildCodexUsageProgressFromExtra_UsesCanonicalUsedPercent(t *testing.T) {
- t.Parallel()
- now := time.Date(2026, 5, 30, 7, 4, 9, 0, time.UTC)
- extra := map[string]any{
- "codex_5h_used_percent": 94.0,
- "codex_5h_reset_at": now.Add(2 * time.Hour).Format(time.RFC3339),
- "codex_7d_used_percent": 93.0,
- "codex_7d_reset_at": now.Add(5 * 24 * time.Hour).Format(time.RFC3339),
- }
-
- fiveHour := buildCodexUsageProgressFromExtra(extra, "5h", now)
- if fiveHour == nil {
- t.Fatal("expected non-nil 5h progress")
- }
- if fiveHour.Utilization != 94.0 {
- t.Fatalf("5h Utilization = %v, want 94", fiveHour.Utilization)
- }
-
- sevenDay := buildCodexUsageProgressFromExtra(extra, "7d", now)
- if sevenDay == nil {
- t.Fatal("expected non-nil 7d progress")
- }
- if sevenDay.Utilization != 93.0 {
- t.Fatalf("7d Utilization = %v, want 93", sevenDay.Utilization)
- }
-}
-
func TestAccountUsageService_PersistOpenAICodexProbeSnapshotOnlyUpdatesExtra(t *testing.T) {
t.Parallel()
diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go
index b1a95594..17ac7fc2 100644
--- a/backend/internal/service/openai_gateway_service.go
+++ b/backend/internal/service/openai_gateway_service.go
@@ -126,19 +126,6 @@ type NormalizedCodexLimits struct {
Window7dMinutes *int
}
-func normalizeCodexFiveHourUsedPercent(raw *float64) *float64 {
- if raw == nil {
- return nil
- }
- // OpenAI's 5h Codex quota header is remaining%, despite the upstream header
- // name saying "used"; the canonical codex_5h_used_percent field stores used%.
- used := 100 - *raw
- if used < 0 {
- used = 0
- }
- return &used
-}
-
// Normalize converts primary/secondary fields to canonical 5h/7d fields.
// Strategy: Compare window_minutes to determine which is 5h vs 7d.
// Returns nil if snapshot is nil or has no useful data.
@@ -197,7 +184,7 @@ func (s *OpenAICodexUsageSnapshot) Normalize() *NormalizedCodexLimits {
// Assign values
if use5hFromPrimary {
- result.Used5hPercent = normalizeCodexFiveHourUsedPercent(s.PrimaryUsedPercent)
+ result.Used5hPercent = s.PrimaryUsedPercent
result.Reset5hSeconds = s.PrimaryResetAfterSeconds
result.Window5hMinutes = s.PrimaryWindowMinutes
result.Used7dPercent = s.SecondaryUsedPercent
@@ -207,7 +194,7 @@ func (s *OpenAICodexUsageSnapshot) Normalize() *NormalizedCodexLimits {
result.Used7dPercent = s.PrimaryUsedPercent
result.Reset7dSeconds = s.PrimaryResetAfterSeconds
result.Window7dMinutes = s.PrimaryWindowMinutes
- result.Used5hPercent = normalizeCodexFiveHourUsedPercent(s.SecondaryUsedPercent)
+ result.Used5hPercent = s.SecondaryUsedPercent
result.Reset5hSeconds = s.SecondaryResetAfterSeconds
result.Window5hMinutes = s.SecondaryWindowMinutes
}
diff --git a/backend/internal/service/openai_gateway_service_codex_snapshot_test.go b/backend/internal/service/openai_gateway_service_codex_snapshot_test.go
index 22f5fa74..654dd4ca 100644
--- a/backend/internal/service/openai_gateway_service_codex_snapshot_test.go
+++ b/backend/internal/service/openai_gateway_service_codex_snapshot_test.go
@@ -104,40 +104,6 @@ func TestBuildCodexUsageExtraUpdates_UsesSnapshotUpdatedAt(t *testing.T) {
}
}
-func TestBuildCodexUsageExtraUpdates_NormalizesFiveHourRemainingToUsedPercent(t *testing.T) {
- primaryUsed := 93.0
- primaryReset := 86400
- primaryWindow := 10080
- secondaryRemaining := 6.0
- secondaryReset := 3600
- secondaryWindow := 300
-
- snapshot := &OpenAICodexUsageSnapshot{
- PrimaryUsedPercent: &primaryUsed,
- PrimaryResetAfterSeconds: &primaryReset,
- PrimaryWindowMinutes: &primaryWindow,
- SecondaryUsedPercent: &secondaryRemaining,
- SecondaryResetAfterSeconds: &secondaryReset,
- SecondaryWindowMinutes: &secondaryWindow,
- UpdatedAt: "2026-05-30T07:04:09Z",
- }
-
- updates := buildCodexUsageExtraUpdates(snapshot, time.Time{})
- if updates == nil {
- t.Fatal("expected non-nil updates")
- }
-
- if got := updates["codex_secondary_used_percent"]; got != 6.0 {
- t.Fatalf("codex_secondary_used_percent = %v, want raw upstream value 6", got)
- }
- if got := updates["codex_5h_used_percent"]; got != 94.0 {
- t.Fatalf("codex_5h_used_percent = %v, want 94", got)
- }
- if got := updates["codex_7d_used_percent"]; got != 93.0 {
- t.Fatalf("codex_7d_used_percent = %v, want 93", got)
- }
-}
-
func TestBuildCodexUsageExtraUpdates_FallbackToNowWhenUpdatedAtInvalid(t *testing.T) {
primaryUsed := 15.0
primaryReset := 30
diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go
index 5c4e979d..8aad2fa6 100644
--- a/backend/internal/service/openai_gateway_service_test.go
+++ b/backend/internal/service/openai_gateway_service_test.go
@@ -1774,7 +1774,7 @@ func TestOpenAIUpdateCodexUsageSnapshotFromHeaders(t *testing.T) {
select {
case updates := <-repo.updateExtraCalls:
- require.Equal(t, 88.0, updates["codex_5h_used_percent"])
+ require.Equal(t, 12.0, updates["codex_5h_used_percent"])
require.Equal(t, 34.0, updates["codex_7d_used_percent"])
require.Equal(t, 600, updates["codex_5h_reset_after_seconds"])
require.Equal(t, 86400, updates["codex_7d_reset_after_seconds"])
diff --git a/backend/internal/service/ratelimit_service_openai_test.go b/backend/internal/service/ratelimit_service_openai_test.go
index 107ac27e..aa5a070c 100644
--- a/backend/internal/service/ratelimit_service_openai_test.go
+++ b/backend/internal/service/ratelimit_service_openai_test.go
@@ -51,7 +51,7 @@ func TestCalculateOpenAI429ResetTime_5hExhausted(t *testing.T) {
headers.Set("x-codex-primary-used-percent", "50")
headers.Set("x-codex-primary-reset-after-seconds", "500000")
headers.Set("x-codex-primary-window-minutes", "10080") // 7 days
- headers.Set("x-codex-secondary-used-percent", "0")
+ headers.Set("x-codex-secondary-used-percent", "100")
headers.Set("x-codex-secondary-reset-after-seconds", "3600") // 1 hour
headers.Set("x-codex-secondary-window-minutes", "300") // 5 hours
@@ -122,7 +122,7 @@ func TestCalculateOpenAI429ResetTime_ReversedWindowOrder(t *testing.T) {
// Test when OpenAI sends primary as 5h and secondary as 7d (reversed)
headers := http.Header{}
- headers.Set("x-codex-primary-used-percent", "0") // This is 5h remaining%
+ headers.Set("x-codex-primary-used-percent", "100") // This is 5h
headers.Set("x-codex-primary-reset-after-seconds", "3600") // 1 hour
headers.Set("x-codex-primary-window-minutes", "300") // 5 hours - smaller!
headers.Set("x-codex-secondary-used-percent", "50")
@@ -180,7 +180,7 @@ func TestHandle429_OpenAIPersistsCodexSnapshotImmediately(t *testing.T) {
headers.Set("x-codex-primary-used-percent", "100")
headers.Set("x-codex-primary-reset-after-seconds", "604800")
headers.Set("x-codex-primary-window-minutes", "10080")
- headers.Set("x-codex-secondary-used-percent", "0")
+ headers.Set("x-codex-secondary-used-percent", "100")
headers.Set("x-codex-secondary-reset-after-seconds", "18000")
headers.Set("x-codex-secondary-window-minutes", "300")
@@ -224,7 +224,7 @@ func TestNormalizedCodexLimits(t *testing.T) {
pUsed := 100.0
pReset := 384607
pWindow := 10080
- sRemaining := 3.0
+ sUsed := 3.0
sReset := 17369
sWindow := 300
@@ -232,7 +232,7 @@ func TestNormalizedCodexLimits(t *testing.T) {
PrimaryUsedPercent: &pUsed,
PrimaryResetAfterSeconds: &pReset,
PrimaryWindowMinutes: &pWindow,
- SecondaryUsedPercent: &sRemaining,
+ SecondaryUsedPercent: &sUsed,
SecondaryResetAfterSeconds: &sReset,
SecondaryWindowMinutes: &sWindow,
}
@@ -249,8 +249,8 @@ func TestNormalizedCodexLimits(t *testing.T) {
if normalized.Reset7dSeconds == nil || *normalized.Reset7dSeconds != 384607 {
t.Errorf("expected Reset7dSeconds=384607, got %v", normalized.Reset7dSeconds)
}
- if normalized.Used5hPercent == nil || *normalized.Used5hPercent != 97.0 {
- t.Errorf("expected Used5hPercent=97, got %v", normalized.Used5hPercent)
+ if normalized.Used5hPercent == nil || *normalized.Used5hPercent != 3.0 {
+ t.Errorf("expected Used5hPercent=3, got %v", normalized.Used5hPercent)
}
if normalized.Reset5hSeconds == nil || *normalized.Reset5hSeconds != 17369 {
t.Errorf("expected Reset5hSeconds=17369, got %v", normalized.Reset5hSeconds)
@@ -338,11 +338,11 @@ func TestRateLimitService_HandleUpstreamError_403FallsBackToRawBody(t *testing.T
func TestNormalizedCodexLimits_OnlySecondaryData(t *testing.T) {
// Test when only secondary has data, no window_minutes
- sRemaining := 60.0
+ sUsed := 60.0
sReset := 3000
snapshot := &OpenAICodexUsageSnapshot{
- SecondaryUsedPercent: &sRemaining,
+ SecondaryUsedPercent: &sUsed,
SecondaryResetAfterSeconds: &sReset,
// No window_minutes, no primary data
}
@@ -354,8 +354,8 @@ func TestNormalizedCodexLimits_OnlySecondaryData(t *testing.T) {
// Legacy assumption: primary=7d, secondary=5h
// So secondary goes to 5h
- if normalized.Used5hPercent == nil || *normalized.Used5hPercent != 40.0 {
- t.Errorf("expected Used5hPercent=40, got %v", normalized.Used5hPercent)
+ if normalized.Used5hPercent == nil || *normalized.Used5hPercent != 60.0 {
+ t.Errorf("expected Used5hPercent=60, got %v", normalized.Used5hPercent)
}
if normalized.Reset5hSeconds == nil || *normalized.Reset5hSeconds != 3000 {
t.Errorf("expected Reset5hSeconds=3000, got %v", normalized.Reset5hSeconds)
@@ -370,13 +370,13 @@ func TestNormalizedCodexLimits_BothDataNoWindowMinutes(t *testing.T) {
// Test when both have data but no window_minutes
pUsed := 100.0
pReset := 400000
- sRemaining := 30.0
+ sUsed := 50.0
sReset := 10000
snapshot := &OpenAICodexUsageSnapshot{
PrimaryUsedPercent: &pUsed,
PrimaryResetAfterSeconds: &pReset,
- SecondaryUsedPercent: &sRemaining,
+ SecondaryUsedPercent: &sUsed,
SecondaryResetAfterSeconds: &sReset,
// No window_minutes
}
@@ -393,8 +393,8 @@ func TestNormalizedCodexLimits_BothDataNoWindowMinutes(t *testing.T) {
if normalized.Reset7dSeconds == nil || *normalized.Reset7dSeconds != 400000 {
t.Errorf("expected Reset7dSeconds=400000, got %v", normalized.Reset7dSeconds)
}
- if normalized.Used5hPercent == nil || *normalized.Used5hPercent != 70.0 {
- t.Errorf("expected Used5hPercent=70, got %v", normalized.Used5hPercent)
+ if normalized.Used5hPercent == nil || *normalized.Used5hPercent != 50.0 {
+ t.Errorf("expected Used5hPercent=50, got %v", normalized.Used5hPercent)
}
if normalized.Reset5hSeconds == nil || *normalized.Reset5hSeconds != 10000 {
t.Errorf("expected Reset5hSeconds=10000, got %v", normalized.Reset5hSeconds)
@@ -425,7 +425,7 @@ func TestCalculateOpenAI429ResetTime_UserProvidedScenario(t *testing.T) {
// This is the exact scenario from the user:
// codex_7d_used_percent: 100
// codex_7d_reset_after_seconds: 384607 (约4.5天后重置)
- // codex_5h_used_percent: 97 (from upstream 3% remaining)
+ // codex_5h_used_percent: 3
// codex_5h_reset_after_seconds: 17369 (约4.8小时后重置)
svc := &RateLimitService{}
From 5634cc83eb91ee4a967e4eb6c65a525b13d02884 Mon Sep 17 00:00:00 2001
From: Pluviobyte
Date: Wed, 3 Jun 2026 14:24:50 +0800
Subject: [PATCH 72/79] chore: bump Go patch version
---
.github/workflows/backend-ci.yml | 4 ++--
.github/workflows/release.yml | 2 +-
.github/workflows/security-scan.yml | 2 +-
backend/go.mod | 2 +-
4 files changed, 5 insertions(+), 5 deletions(-)
diff --git a/.github/workflows/backend-ci.yml b/.github/workflows/backend-ci.yml
index 15ff97fe..fb4d0ce6 100644
--- a/.github/workflows/backend-ci.yml
+++ b/.github/workflows/backend-ci.yml
@@ -20,7 +20,7 @@ jobs:
cache-dependency-path: backend/go.sum
- name: Verify Go version
run: |
- go version | grep -q 'go1.26.3'
+ go version | grep -q 'go1.26.4'
- name: Unit tests
working-directory: backend
run: make test-unit
@@ -60,7 +60,7 @@ jobs:
cache-dependency-path: backend/go.sum
- name: Verify Go version
run: |
- go version | grep -q 'go1.26.3'
+ go version | grep -q 'go1.26.4'
- name: golangci-lint
uses: golangci/golangci-lint-action@v9
with:
diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml
index 80bc9850..7d48131a 100644
--- a/.github/workflows/release.yml
+++ b/.github/workflows/release.yml
@@ -115,7 +115,7 @@ jobs:
- name: Verify Go version
run: |
- go version | grep -q 'go1.26.3'
+ go version | grep -q 'go1.26.4'
# Docker setup for GoReleaser
- name: Set up QEMU
diff --git a/.github/workflows/security-scan.yml b/.github/workflows/security-scan.yml
index ef8e59e5..e102b5f8 100644
--- a/.github/workflows/security-scan.yml
+++ b/.github/workflows/security-scan.yml
@@ -23,7 +23,7 @@ jobs:
cache-dependency-path: backend/go.sum
- name: Verify Go version
run: |
- go version | grep -q 'go1.26.3'
+ go version | grep -q 'go1.26.4'
- name: Run govulncheck
working-directory: backend
run: |
diff --git a/backend/go.mod b/backend/go.mod
index 587d5370..62be56c8 100644
--- a/backend/go.mod
+++ b/backend/go.mod
@@ -1,6 +1,6 @@
module github.com/Wei-Shaw/sub2api
-go 1.26.3
+go 1.26.4
require (
entgo.io/ent v0.14.5
From 60867022b64ecb81a31670b0ff38eecf5ac54edf Mon Sep 17 00:00:00 2001
From: visa2
Date: Wed, 3 Jun 2026 17:26:05 +0800
Subject: [PATCH 73/79] =?UTF-8?q?fix(apicompat):=20repair=20tool=5Fuse/too?=
=?UTF-8?q?l=5Fresult=20pairing=20on=20the=20Responses=E2=86=92Anthropic?=
=?UTF-8?q?=20path?=
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit
When an OpenAI Chat Completions client targets an Anthropic-platform group,
ForwardAsChatCompletions converts the request CC → Responses → Anthropic
(ChatCompletionsToResponses → ResponsesToAnthropicRequest) before forwarding it
upstream. The Responses→Anthropic converter emits each function_call as its own
assistant message and each function_call_output as its own user message and
relies solely on mergeConsecutiveMessages to alternate roles. That is not enough
to satisfy Anthropic's tool-pairing invariants, so a trimmed or partial tool
history produces an upstream 400, e.g.:
tool_use_id found in tool_result blocks: call_00_...
Each tool_result block must have a corresponding tool_use block in the
previous message.
The failures this leaves unrepaired:
- orphan tool_result — a client that does sliding-window context management
keeps a recent tool result but drops the assistant tool_calls message that
announced it, so the tool_result has no matching tool_use;
- unanswered/dangling tool_use — a parallel call whose sibling result never
came back, or a call left dangling, which Anthropic also rejects.
Add normalizeAnthropicToolPairing, run between two merge passes: the first merge
groups parallel calls and their results; the pairing pass indexes every
tool_result by its tool_use id, keeps only answered tool_use blocks (dropping
unanswered/dangling calls, and the assistant message entirely when nothing else
remains) and re-emits the matching tool_result blocks as the immediately
following user message; standalone/orphan tool_results are dropped from their
original position; the second merge restores alternation. This mirrors
normalizeChatMessages on the Responses→Chat path.
Tested two ways: responses_to_anthropic_tool_pairing_test.go covers the repair
on direct Responses input (developer message between call and output, parallel
both-answered kept grouped, parallel one-unanswered dropped, orphan tool_result,
dangling call, single-call baseline); responses_to_anthropic_cc_chain_test.go
drives the real ChatCompletionsToResponses → ResponsesToAnthropicRequest chain
and reproduces the production 400 (orphan and unanswered-parallel) — both fail
without the repair and pass with it. The full apicompat suite stays green.
Co-Authored-By: Claude Opus 4.8
---
.../responses_to_anthropic_cc_chain_test.go | 102 +++++++++++
.../responses_to_anthropic_request.go | 125 ++++++++++++-
...esponses_to_anthropic_tool_pairing_test.go | 165 ++++++++++++++++++
3 files changed, 391 insertions(+), 1 deletion(-)
create mode 100644 backend/internal/pkg/apicompat/responses_to_anthropic_cc_chain_test.go
create mode 100644 backend/internal/pkg/apicompat/responses_to_anthropic_tool_pairing_test.go
diff --git a/backend/internal/pkg/apicompat/responses_to_anthropic_cc_chain_test.go b/backend/internal/pkg/apicompat/responses_to_anthropic_cc_chain_test.go
new file mode 100644
index 00000000..d64680f4
--- /dev/null
+++ b/backend/internal/pkg/apicompat/responses_to_anthropic_cc_chain_test.go
@@ -0,0 +1,102 @@
+package apicompat
+
+import (
+ "encoding/json"
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+// These tests drive the exact production path for Chat Completions clients on an
+// Anthropic-platform group: ForwardAsChatCompletions runs
+// ChatCompletionsToResponses → ResponsesToAnthropicRequest
+// (gateway_forward_as_chat_completions.go), then forwards the Anthropic body
+// upstream. They assert the tool-pairing repair holds through that full chain,
+// not only for codex-style Responses input.
+func ccChainToAnthropic(t *testing.T, ccReq *ChatCompletionsRequest) []AnthropicMessage {
+ t.Helper()
+ respReq, err := ChatCompletionsToResponses(ccReq)
+ require.NoError(t, err)
+ anthReq, err := ResponsesToAnthropicRequest(respReq)
+ require.NoError(t, err)
+ assertAnthropicPairing(t, anthReq.Messages)
+ return anthReq.Messages
+}
+
+// Reproduces the production 400:
+//
+// unexpected ...content.0: tool_use_id found in tool_result blocks:
+// call_00_TgfbRvKlnD7oK6Dg00sL1661. Each tool_result block must have a
+// corresponding tool_use block in the previous message.
+//
+// A Chat Completions client trimmed its history and kept a tool result whose
+// announcing assistant tool_calls message was dropped (sliding-window context
+// management). The orphan tool_result has no matching tool_use → upstream 400.
+// The repair drops the orphan so the request is valid.
+func TestCCChain_OrphanToolResultFromTrimmedHistory(t *testing.T) {
+ orphanID := "call_00_TgfbRvKlnD7oK6Dg00sL1661"
+ msgs := ccChainToAnthropic(t, &ChatCompletionsRequest{
+ Model: "deepseek-v4-pro",
+ Messages: []ChatMessage{
+ {Role: "user", Content: json.RawMessage(`"search the web for X"`)},
+ // The assistant tool_calls message that announced orphanID was trimmed.
+ {Role: "tool", ToolCallID: orphanID, Content: json.RawMessage(`"stale search results"`)},
+ {Role: "assistant", Content: json.RawMessage(`"Here is what I found."`)},
+ {Role: "user", Content: json.RawMessage(`"thanks, now do Y"`)},
+ },
+ })
+ for _, m := range msgs {
+ require.Falsef(t, hasToolResult(parseContentBlocks(m.Content), orphanID),
+ "orphan tool_result %s should have been dropped", orphanID)
+ }
+}
+
+// A parallel web_search where one sibling's result never came back (the tool
+// failed/was skipped). The unanswered tool_use would otherwise trip Anthropic's
+// "tool_use without tool_result" check; the repair drops it.
+func TestCCChain_ParallelToolOneResultMissing(t *testing.T) {
+ msgs := ccChainToAnthropic(t, &ChatCompletionsRequest{
+ Model: "deepseek-v4-pro",
+ Messages: []ChatMessage{
+ {Role: "user", Content: json.RawMessage(`"search A and B"`)},
+ {Role: "assistant", Content: json.RawMessage(`"searching both"`), ToolCalls: []ChatToolCall{
+ {ID: "call_a", Type: "function", Function: ChatFunctionCall{Name: "web_search", Arguments: `{"q":"A"}`}},
+ {ID: "call_b", Type: "function", Function: ChatFunctionCall{Name: "web_search", Arguments: `{"q":"B"}`}},
+ }},
+ {Role: "tool", ToolCallID: "call_a", Content: json.RawMessage(`"result A"`)},
+ // call_b's result is missing.
+ },
+ })
+ for _, m := range msgs {
+ require.Falsef(t, hasToolUse(parseContentBlocks(m.Content), "call_b"),
+ "unanswered tool_use call_b should have been dropped")
+ }
+}
+
+// Baseline: a well-formed multi-round tool history (text + tool_calls per
+// assistant turn) converts and pairs correctly through the full chain.
+func TestCCChain_WellFormedMultiRound(t *testing.T) {
+ msgs := ccChainToAnthropic(t, &ChatCompletionsRequest{
+ Model: "deepseek-v4-pro",
+ Messages: []ChatMessage{
+ {Role: "user", Content: json.RawMessage(`"do A then B"`)},
+ {Role: "assistant", Content: json.RawMessage(`"running A"`), ToolCalls: []ChatToolCall{
+ {ID: "call_a", Type: "function", Function: ChatFunctionCall{Name: "exec", Arguments: `{"cmd":"A"}`}},
+ }},
+ {Role: "tool", ToolCallID: "call_a", Content: json.RawMessage(`"A ok"`)},
+ {Role: "assistant", Content: json.RawMessage(`"A done, running B"`), ToolCalls: []ChatToolCall{
+ {ID: "call_b", Type: "function", Function: ChatFunctionCall{Name: "exec", Arguments: `{"cmd":"B"}`}},
+ }},
+ {Role: "tool", ToolCallID: "call_b", Content: json.RawMessage(`"B ok"`)},
+ {Role: "assistant", Content: json.RawMessage(`"all done"`)},
+ },
+ })
+ // Both calls survive and stay paired (assertAnthropicPairing already checks).
+ var sawA, sawB bool
+ for _, m := range msgs {
+ blocks := parseContentBlocks(m.Content)
+ sawA = sawA || hasToolUse(blocks, "call_a")
+ sawB = sawB || hasToolUse(blocks, "call_b")
+ }
+ require.True(t, sawA && sawB, "both well-formed calls should be preserved")
+}
diff --git a/backend/internal/pkg/apicompat/responses_to_anthropic_request.go b/backend/internal/pkg/apicompat/responses_to_anthropic_request.go
index 8fa652f2..672ad80c 100644
--- a/backend/internal/pkg/apicompat/responses_to_anthropic_request.go
+++ b/backend/internal/pkg/apicompat/responses_to_anthropic_request.go
@@ -192,12 +192,135 @@ func convertResponsesInputToAnthropic(inputRaw json.RawMessage) (json.RawMessage
}
}
- // Merge consecutive same-role messages (Anthropic requires alternating roles)
+ // Repair tool_use/tool_result pairing, then merge consecutive same-role
+ // messages (Anthropic requires alternating roles). The first merge groups
+ // parallel calls (and their results) so the pairing pass sees them together;
+ // the pairing pass may re-split a user turn (e.g. when an injected message
+ // sat between a call and its output), so a second merge restores alternation.
+ messages = mergeConsecutiveMessages(messages)
+ messages = normalizeAnthropicToolPairing(messages)
messages = mergeConsecutiveMessages(messages)
return system, messages, nil
}
+// normalizeAnthropicToolPairing rebuilds the message sequence so it satisfies
+// Anthropic's tool_use/tool_result invariants, which the naive item-by-item
+// conversion violates whenever the Responses history interleaves anything
+// between a function_call and its function_call_output:
+//
+// - every tool_result block must have a matching tool_use in the immediately
+// preceding assistant message ("tool_result ... must have a corresponding
+// tool_use block in the previous message");
+// - every tool_use block must be answered by a tool_result in the immediately
+// following user message (Anthropic rejects unanswered tool_use ids);
+// - user/assistant turns must alternate.
+//
+// codex (Responses, store:false) re-sends the whole history each turn and
+// frequently injects items between a call and its output — a developer/approval
+// notice, or a sibling parallel call whose output never arrived. The unrepaired
+// converter emits each function_call as its own assistant message and each
+// output as its own user message, so any such interleaving breaks
+// tool_use↔tool_result adjacency and yields an upstream 400.
+//
+// The repair indexes every tool_result by its tool_use id, then for each
+// assistant message carrying tool_use blocks keeps only the answered ones
+// (dropping unanswered/dangling calls — and the assistant message entirely if it
+// has no other content) and emits the matching tool_result blocks, in call
+// order, as the very next user message. Standalone tool_result blocks are
+// dropped from their original position (re-emitted adjacent to their call);
+// orphan tool_results with no announcing tool_use are dropped. Non-tool content
+// passes through in place. This mirrors normalizeChatMessages on the
+// Responses→Chat path.
+func normalizeAnthropicToolPairing(messages []AnthropicMessage) []AnthropicMessage {
+ // Index every tool_result block by its tool_use id (last wins on dup).
+ results := make(map[string]AnthropicContentBlock)
+ for _, m := range messages {
+ if m.Role != "user" {
+ continue
+ }
+ for _, b := range parseContentBlocks(m.Content) {
+ if b.Type == "tool_result" && b.ToolUseID != "" {
+ results[b.ToolUseID] = b
+ }
+ }
+ }
+
+ out := make([]AnthropicMessage, 0, len(messages))
+ for _, m := range messages {
+ blocks := parseContentBlocks(m.Content)
+ switch m.Role {
+ case "assistant":
+ var toolUses, others []AnthropicContentBlock
+ for _, b := range blocks {
+ if b.Type == "tool_use" {
+ toolUses = append(toolUses, b)
+ } else {
+ others = append(others, b)
+ }
+ }
+ if len(toolUses) == 0 {
+ out = append(out, m)
+ continue
+ }
+ kept := make([]AnthropicContentBlock, 0, len(toolUses))
+ for _, tu := range toolUses {
+ if _, ok := results[tu.ID]; ok {
+ kept = append(kept, tu)
+ }
+ }
+ if len(kept) == 0 {
+ // No answered calls: keep any non-tool content, else drop.
+ if len(others) > 0 {
+ out = append(out, anthropicMessageFromBlocks("assistant", others))
+ }
+ continue
+ }
+ asstBlocks := make([]AnthropicContentBlock, 0, len(others)+len(kept))
+ asstBlocks = append(asstBlocks, others...)
+ asstBlocks = append(asstBlocks, kept...)
+ out = append(out, anthropicMessageFromBlocks("assistant", asstBlocks))
+
+ resBlocks := make([]AnthropicContentBlock, 0, len(kept))
+ for _, tu := range kept {
+ resBlocks = append(resBlocks, results[tu.ID])
+ }
+ out = append(out, anthropicMessageFromBlocks("user", resBlocks))
+
+ case "user":
+ var nonResult []AnthropicContentBlock
+ hasResult := false
+ for _, b := range blocks {
+ if b.Type == "tool_result" {
+ hasResult = true
+ continue
+ }
+ nonResult = append(nonResult, b)
+ }
+ if !hasResult {
+ out = append(out, m)
+ continue
+ }
+ // The tool_result blocks are re-emitted next to their call; keep any
+ // other content of this user turn in place, drop it if there is none.
+ if len(nonResult) > 0 {
+ out = append(out, anthropicMessageFromBlocks("user", nonResult))
+ }
+
+ default:
+ out = append(out, m)
+ }
+ }
+ return out
+}
+
+// anthropicMessageFromBlocks builds an AnthropicMessage whose content is the
+// marshaled block array.
+func anthropicMessageFromBlocks(role string, blocks []AnthropicContentBlock) AnthropicMessage {
+ content, _ := json.Marshal(blocks)
+ return AnthropicMessage{Role: role, Content: content}
+}
+
// extractTextFromContent extracts text from a content field that may be a
// plain string or an array of content parts.
func extractTextFromContent(raw json.RawMessage) string {
diff --git a/backend/internal/pkg/apicompat/responses_to_anthropic_tool_pairing_test.go b/backend/internal/pkg/apicompat/responses_to_anthropic_tool_pairing_test.go
new file mode 100644
index 00000000..b2522f27
--- /dev/null
+++ b/backend/internal/pkg/apicompat/responses_to_anthropic_tool_pairing_test.go
@@ -0,0 +1,165 @@
+package apicompat
+
+import (
+ "encoding/json"
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+// assertAnthropicPairing enforces the Anthropic Messages tool-pairing invariants
+// that, when violated, surface as upstream 400s.
+func assertAnthropicPairing(t *testing.T, messages []AnthropicMessage) {
+ t.Helper()
+ for i, m := range messages {
+ blocks := parseContentBlocks(m.Content)
+
+ // No two consecutive same-role messages.
+ if i > 0 {
+ require.NotEqualf(t, messages[i-1].Role, m.Role, "consecutive %s messages at %d", m.Role, i)
+ }
+
+ for _, b := range blocks {
+ switch b.Type {
+ case "tool_result":
+ // Must have a matching tool_use in the immediately previous message.
+ require.Positivef(t, i, "tool_result %s has no previous message", b.ToolUseID)
+ prev := parseContentBlocks(messages[i-1].Content)
+ require.Truef(t, hasToolUse(prev, b.ToolUseID),
+ "tool_result %s has no corresponding tool_use in previous message", b.ToolUseID)
+ case "tool_use":
+ // Must be answered by a tool_result in the immediately next message.
+ require.Lessf(t, i+1, len(messages), "tool_use %s has no following message", b.ID)
+ next := parseContentBlocks(messages[i+1].Content)
+ require.Truef(t, hasToolResult(next, b.ID),
+ "tool_use %s is not answered in the next message", b.ID)
+ }
+ }
+ }
+}
+
+func hasToolUse(blocks []AnthropicContentBlock, id string) bool {
+ for _, b := range blocks {
+ if b.Type == "tool_use" && b.ID == id {
+ return true
+ }
+ }
+ return false
+}
+
+func hasToolResult(blocks []AnthropicContentBlock, toolUseID string) bool {
+ for _, b := range blocks {
+ if b.Type == "tool_result" && b.ToolUseID == toolUseID {
+ return true
+ }
+ }
+ return false
+}
+
+func convertAnthropic(t *testing.T, input string) []AnthropicMessage {
+ t.Helper()
+ _, messages, err := convertResponsesInputToAnthropic(json.RawMessage(input))
+ require.NoError(t, err)
+ assertAnthropicPairing(t, messages)
+ return messages
+}
+
+// Tests use call_-prefixed ids because fromResponsesCallIDToAnthropic passes
+// those through unchanged (matching codex's real call_00_... ids); bare ids
+// would be rewritten to toolu_.
+
+// A developer/approval message injected between a function_call and its output
+// must be moved out of the tool_use→tool_result adjacency. This is the shape
+// that produced the production 400 "tool_result ... must have a corresponding
+// tool_use block in the previous message".
+func TestAnthropicPairing_DeveloperMessageBetween(t *testing.T) {
+ msgs := convertAnthropic(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"do it"}]},
+ {"type":"function_call","call_id":"call_A","name":"exec","arguments":"{}"},
+ {"type":"message","role":"developer","content":[{"type":"input_text","text":"Approved command prefix saved"}]},
+ {"type":"function_call_output","call_id":"call_A","output":"ok"}
+ ]`)
+ // The assistant tool_use message is immediately followed by its tool_result.
+ for i, m := range msgs {
+ if hasToolUse(parseContentBlocks(m.Content), "call_A") {
+ require.Equal(t, "user", msgs[i+1].Role)
+ require.True(t, hasToolResult(parseContentBlocks(msgs[i+1].Content), "call_A"))
+ }
+ }
+}
+
+// Parallel tool calls where both outputs arrive stay grouped: one assistant
+// message with both tool_use blocks, the next user message with both results.
+func TestAnthropicPairing_ParallelBothAnswered(t *testing.T) {
+ msgs := convertAnthropic(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"features?"}]},
+ {"type":"function_call","call_id":"call_c0","name":"exec","arguments":"{}"},
+ {"type":"function_call","call_id":"call_c1","name":"exec","arguments":"{}"},
+ {"type":"function_call_output","call_id":"call_c0","output":"log"},
+ {"type":"function_call_output","call_id":"call_c1","output":"tags"}
+ ]`)
+ var sawGrouped bool
+ for _, m := range msgs {
+ blocks := parseContentBlocks(m.Content)
+ if hasToolUse(blocks, "call_c0") && hasToolUse(blocks, "call_c1") {
+ sawGrouped = true
+ }
+ }
+ require.True(t, sawGrouped, "parallel tool_use blocks should share one assistant message")
+}
+
+// A parallel call whose sibling output never arrived must be dropped so every
+// remaining tool_use is answered.
+func TestAnthropicPairing_ParallelOneUnanswered(t *testing.T) {
+ msgs := convertAnthropic(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"q"}]},
+ {"type":"function_call","call_id":"call_A","name":"exec","arguments":"{}"},
+ {"type":"function_call","call_id":"call_B","name":"exec","arguments":"{}"},
+ {"type":"function_call_output","call_id":"call_A","output":"oa"}
+ ]`)
+ for _, m := range msgs {
+ require.Falsef(t, hasToolUse(parseContentBlocks(m.Content), "call_B"),
+ "unanswered tool_use call_B should have been dropped")
+ }
+}
+
+// An orphan tool_result whose tool_use was never announced must be dropped.
+func TestAnthropicPairing_OrphanToolResultDropped(t *testing.T) {
+ msgs := convertAnthropic(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"q"}]},
+ {"type":"function_call_output","call_id":"call_ghost","output":"orphan"}
+ ]`)
+ for _, m := range msgs {
+ require.Falsef(t, hasToolResult(parseContentBlocks(m.Content), "call_ghost"),
+ "orphan tool_result should have been dropped")
+ }
+}
+
+// A dangling tool_call at the end of the history (no output yet) drops the
+// assistant message holding only that call, leaving no tool_use behind.
+func TestAnthropicPairing_DanglingCallDropped(t *testing.T) {
+ msgs := convertAnthropic(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"q"}]},
+ {"type":"function_call","call_id":"call_A","name":"exec","arguments":"{}"}
+ ]`)
+ for _, m := range msgs {
+ require.Falsef(t, hasToolUse(parseContentBlocks(m.Content), "call_A"),
+ "dangling tool_use call_A should have been dropped")
+ }
+}
+
+// Baseline: a single answered call pairs correctly and preserves the surrounding
+// turns.
+func TestAnthropicPairing_SingleCall(t *testing.T) {
+ msgs := convertAnthropic(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"latest sha?"}]},
+ {"type":"function_call","call_id":"call_A","name":"exec","arguments":"{\"cmd\":\"git rev-parse HEAD\"}"},
+ {"type":"function_call_output","call_id":"call_A","output":"deadbeef"},
+ {"type":"message","role":"assistant","content":[{"type":"output_text","text":"It is deadbeef."}]}
+ ]`)
+ // user, assistant(tool_use), user(tool_result), assistant(text)
+ require.GreaterOrEqual(t, len(msgs), 4)
+ require.Equal(t, "user", msgs[0].Role)
+ require.True(t, hasToolUse(parseContentBlocks(msgs[1].Content), "call_A"))
+ require.True(t, hasToolResult(parseContentBlocks(msgs[2].Content), "call_A"))
+}
From bc7ce185749c543df19cfb92b1909d2f41516386 Mon Sep 17 00:00:00 2001
From: ghostg00 <28946120+ghostg00@users.noreply.github.com>
Date: Thu, 4 Jun 2026 19:48:03 +0800
Subject: [PATCH 74/79] =?UTF-8?q?fix(group):=20=E7=AE=A1=E7=90=86=E5=91=98?=
=?UTF-8?q?=E6=B8=85=E7=A9=BA=E5=88=86=E7=BB=84=E6=8F=8F=E8=BF=B0=E6=97=B6?=
=?UTF-8?q?=E6=AD=A3=E7=A1=AE=E6=8C=81=E4=B9=85=E5=8C=96?=
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit
UpdateGroup 之前用 `if input.Description != ""` 判空,
把"未提供"和"显式置空"混为一谈,导致管理员在分组编辑表单
里清空备注后保存无效。
将 UpdateGroupRequest / UpdateGroupInput 的 Description 改为
*string:nil 表示未提供(保持原值),"" 表示显式清空。
---
.../internal/handler/admin/group_handler.go | 2 +-
backend/internal/service/admin_service.go | 6 +--
.../service/admin_service_group_test.go | 42 ++++++++++++++++++-
3 files changed, 45 insertions(+), 5 deletions(-)
diff --git a/backend/internal/handler/admin/group_handler.go b/backend/internal/handler/admin/group_handler.go
index dbf6f709..102ee02f 100644
--- a/backend/internal/handler/admin/group_handler.go
+++ b/backend/internal/handler/admin/group_handler.go
@@ -123,7 +123,7 @@ type CreateGroupRequest struct {
// UpdateGroupRequest represents update group request
type UpdateGroupRequest struct {
Name string `json:"name"`
- Description string `json:"description"`
+ Description *string `json:"description"`
Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity"`
RateMultiplier *float64 `json:"rate_multiplier"`
IsExclusive *bool `json:"is_exclusive"`
diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go
index 00205d1f..ae9dd8f6 100644
--- a/backend/internal/service/admin_service.go
+++ b/backend/internal/service/admin_service.go
@@ -230,7 +230,7 @@ type CreateGroupInput struct {
type UpdateGroupInput struct {
Name string
- Description string
+ Description *string
Platform string
RateMultiplier *float64 // 使用指针以支持设置为0
IsExclusive *bool
@@ -1924,8 +1924,8 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd
if input.Name != "" {
group.Name = input.Name
}
- if input.Description != "" {
- group.Description = input.Description
+ if input.Description != nil {
+ group.Description = *input.Description
}
if input.Platform != "" {
group.Platform = input.Platform
diff --git a/backend/internal/service/admin_service_group_test.go b/backend/internal/service/admin_service_group_test.go
index 0a2020ea..eb3eff7f 100644
--- a/backend/internal/service/admin_service_group_test.go
+++ b/backend/internal/service/admin_service_group_test.go
@@ -280,8 +280,9 @@ func TestAdminService_UpdateGroup_PreservesImageGenerationControlsWhenOmitted(t
repo := &groupRepoStubForAdmin{getByID: existingGroup}
svc := &adminServiceImpl{groupRepo: repo}
+ updatedDesc := "updated"
group, err := svc.UpdateGroup(context.Background(), 1, &UpdateGroupInput{
- Description: "updated",
+ Description: &updatedDesc,
})
require.NoError(t, err)
require.NotNil(t, group)
@@ -291,6 +292,45 @@ func TestAdminService_UpdateGroup_PreservesImageGenerationControlsWhenOmitted(t
require.InDelta(t, 0.5, repo.updated.ImageRateMultiplier, 1e-12)
}
+func TestAdminService_UpdateGroup_ClearsDescriptionWhenEmptyString(t *testing.T) {
+ existingGroup := &Group{
+ ID: 1,
+ Name: "existing-group",
+ Description: "Auto-created default group",
+ Platform: PlatformOpenAI,
+ Status: StatusActive,
+ }
+ repo := &groupRepoStubForAdmin{getByID: existingGroup}
+ svc := &adminServiceImpl{groupRepo: repo}
+
+ empty := ""
+ _, err := svc.UpdateGroup(context.Background(), 1, &UpdateGroupInput{
+ Description: &empty,
+ })
+ require.NoError(t, err)
+ require.NotNil(t, repo.updated)
+ require.Equal(t, "", repo.updated.Description, "empty string should clear description")
+}
+
+func TestAdminService_UpdateGroup_PreservesDescriptionWhenNil(t *testing.T) {
+ existingGroup := &Group{
+ ID: 1,
+ Name: "existing-group",
+ Description: "keep me",
+ Platform: PlatformOpenAI,
+ Status: StatusActive,
+ }
+ repo := &groupRepoStubForAdmin{getByID: existingGroup}
+ svc := &adminServiceImpl{groupRepo: repo}
+
+ _, err := svc.UpdateGroup(context.Background(), 1, &UpdateGroupInput{
+ Description: nil,
+ })
+ require.NoError(t, err)
+ require.NotNil(t, repo.updated)
+ require.Equal(t, "keep me", repo.updated.Description, "nil should preserve existing description")
+}
+
func TestAdminService_UpdateGroup_RejectsNegativeImageRateMultiplier(t *testing.T) {
existingGroup := &Group{
ID: 1,
From 86d9b6bff982859e66382e081cae7549df3561e8 Mon Sep 17 00:00:00 2001
From: haruka <1628615876@qq.com>
Date: Thu, 4 Jun 2026 22:07:36 +0800
Subject: [PATCH 75/79] fix(openai): self-heal stale Codex used% snapshots +
lock semantics (#2994)
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit
The OpenAI/Codex 5h "used %" inversion that caused fresh accounts to show
~96-99% used (PR #2918, commit b65dde63) was already reverted in #2993, so the
stored value is now the correct "used %" again. This commit hardens that fix:
1. Regression test locking in direct "used %" semantics. The semantics have
flip-flopped twice (#2918 -> #2993) with no value-level guard — a fresh
account (secondary_used_percent=1, 5h window) must store
codex_5h_used_percent=1, not 99.
2. Stale-bounded self-heal in resolveOpenAIQuotaUtilization (the single
auto-pause chokepoint). An account poisoned with an inflated used% gets
excluded from scheduling, and a paused account never receives traffic to
refresh its snapshot — so it stayed stuck until the window's reset_at passed
(up to 5h/7d). When codex_usage_updated_at is older than 2h, the account is
no longer auto-paused on that snapshot; it gets one request whose response
headers refresh the snapshot and self-heal it. A missing timestamp is treated
as fresh (stays paused), and an actively-served exhausted account refreshes
the timestamp every response so it never crosses the bound — it cannot escape
auto-pause.
No change to Normalize(); no 100-x reintroduced; no new dependency wiring.
Co-Authored-By: Claude Opus 4.8 (1M context)
---
.../service/openai_account_scheduler_test.go | 63 +++++++++++++++++++
.../service/openai_gateway_service.go | 31 +++++++++
...nai_gateway_service_codex_snapshot_test.go | 33 ++++++++++
3 files changed, 127 insertions(+)
diff --git a/backend/internal/service/openai_account_scheduler_test.go b/backend/internal/service/openai_account_scheduler_test.go
index da5f0a66..505a5ade 100644
--- a/backend/internal/service/openai_account_scheduler_test.go
+++ b/backend/internal/service/openai_account_scheduler_test.go
@@ -909,6 +909,69 @@ func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_FreshUsageWind
require.Equal(t, int64(35602), account.ID)
}
+// Issue #2994: an account poisoned with an inflated used% (e.g. from the reverted #2918
+// inversion) gets excluded from scheduling, and a paused account never receives traffic to
+// refresh its snapshot. When the snapshot is stale (codex_usage_updated_at older than the
+// staleness bound) the account must be allowed a request so it can self-heal from the real
+// response headers — independent of the window's reset time.
+func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_StaleUsageSnapshotSkipsPause_Issue2994(t *testing.T) {
+ ctx := context.Background()
+ primary := Account{
+ ID: 35701,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 0,
+ Extra: map[string]any{
+ "codex_5h_used_percent": 99.0,
+ "auto_pause_5h_threshold": 0.95,
+ // Window has NOT reset yet, so the reset guard stays inactive.
+ "codex_5h_reset_at": time.Now().Add(time.Hour).Format(time.RFC3339),
+ // Snapshot is stale: older than openAICodexAutoPauseStaleAfter (2h).
+ "codex_usage_updated_at": time.Now().Add(-3 * time.Hour).Format(time.RFC3339),
+ },
+ }
+ secondary := Account{ID: 35702, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
+ svc := &OpenAIGatewayService{accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{primary, secondary}}, cfg: &config.Config{}}
+
+ account, err := svc.SelectAccountForModelWithExclusions(ctx, nil, "", "gpt-5.1", nil)
+ require.NoError(t, err)
+ require.NotNil(t, account)
+ require.Equal(t, int64(35701), account.ID)
+}
+
+// Issue #2994 guardrail: a genuinely-exhausted account whose snapshot was refreshed recently
+// (codex_usage_updated_at fresh) must STILL be auto-paused. The stale self-heal must not let a
+// real 99%-used account escape pause.
+func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_FreshExhaustedSnapshotStillPauses_Issue2994(t *testing.T) {
+ ctx := context.Background()
+ primary := Account{
+ ID: 35801,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 0,
+ Extra: map[string]any{
+ "codex_5h_used_percent": 99.0,
+ "auto_pause_5h_threshold": 0.95,
+ "codex_5h_reset_at": time.Now().Add(time.Hour).Format(time.RFC3339),
+ // Snapshot refreshed 1 minute ago: not stale, so the account stays paused.
+ "codex_usage_updated_at": time.Now().Add(-time.Minute).Format(time.RFC3339),
+ },
+ }
+ secondary := Account{ID: 35802, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
+ svc := &OpenAIGatewayService{accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{primary, secondary}}, cfg: &config.Config{}}
+
+ account, err := svc.SelectAccountForModelWithExclusions(ctx, nil, "", "gpt-5.1", nil)
+ require.NoError(t, err)
+ require.NotNil(t, account)
+ require.Equal(t, int64(35802), account.ID)
+}
+
func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_SkipsFreshlyRateLimitedSnapshotCandidate(t *testing.T) {
ctx := context.Background()
groupID := int64(10102)
diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go
index 17ac7fc2..81fcea9d 100644
--- a/backend/internal/service/openai_gateway_service.go
+++ b/backend/internal/service/openai_gateway_service.go
@@ -59,6 +59,10 @@ const (
codexCLIVersion = "0.125.0"
// Codex 限额快照仅用于后台展示/诊断,不需要每个成功请求都立即落库。
openAICodexSnapshotPersistMinInterval = 30 * time.Second
+ // 配额自动暂停时,超过该时长仍未刷新的 used% 快照视为陈旧,不再据此暂停账号。
+ // 被暂停的账号收不到流量,其快照永远不会从上游响应头刷新;该兜底让账号在快照
+ // 陈旧时放行一次请求,从而通过正常响应头自愈,而无需等待整个窗口(5h/7d)重置。
+ openAICodexAutoPauseStaleAfter = 2 * time.Hour
)
// OpenAI allowed headers whitelist (for non-passthrough).
@@ -1484,9 +1488,36 @@ func resolveOpenAIQuotaUtilization(extra map[string]any, window string, now time
if openAIQuotaWindowReset(extra, window, now) {
return 0, false
}
+ // 快照过于陈旧(账号长期未收到流量刷新)时,不再据此暂停。放行后下一次响应头
+ // 会刷新快照实现自愈,避免账号在错误/过期的 used% 上被永久跳过(issue #2994)。
+ if openAICodexSnapshotStaleForPause(extra, now) {
+ return 0, false
+ }
return usedPercent / 100, true
}
+// openAICodexSnapshotStaleForPause reports whether the Codex usage snapshot is stale
+// enough that it should no longer keep an account auto-paused. It anchors on
+// codex_usage_updated_at (always written by buildCodexUsageExtraUpdates). A missing or
+// unparseable timestamp returns false (treated as fresh, so the account stays paused) —
+// this is deliberate: it prevents any snapshot without a write time from silently escaping
+// auto-pause, and a genuinely-exhausted account that is actively served refreshes the
+// timestamp on every response so it never crosses the staleness bound.
+func openAICodexSnapshotStaleForPause(extra map[string]any, now time.Time) bool {
+ if len(extra) == 0 {
+ return false
+ }
+ updatedRaw, ok := extra["codex_usage_updated_at"]
+ if !ok {
+ return false
+ }
+ updatedAt, err := parseTime(fmt.Sprint(updatedRaw))
+ if err != nil {
+ return false
+ }
+ return now.Sub(updatedAt) >= openAICodexAutoPauseStaleAfter
+}
+
// openAIQuotaWindowReset reports whether the Codex usage window's reset time has
// already passed relative to now. It prefers the absolute codex__reset_at
// timestamp and falls back to codex__reset_after_seconds anchored at
diff --git a/backend/internal/service/openai_gateway_service_codex_snapshot_test.go b/backend/internal/service/openai_gateway_service_codex_snapshot_test.go
index 654dd4ca..27208b58 100644
--- a/backend/internal/service/openai_gateway_service_codex_snapshot_test.go
+++ b/backend/internal/service/openai_gateway_service_codex_snapshot_test.go
@@ -104,6 +104,39 @@ func TestBuildCodexUsageExtraUpdates_UsesSnapshotUpdatedAt(t *testing.T) {
}
}
+// TestBuildCodexUsageExtraUpdates_FreshAccountUsedPercentNotInverted_Issue2994 locks in the
+// canonical "used %" semantics for the 5h window. A fresh account reports a tiny
+// secondary-used-percent (~1%); the stored codex_5h_used_percent must equal that value
+// directly and must NOT be inverted to ~99%. Regression guard for issue #2994 / the reverted
+// commit b65dde63 (PR #2918), which applied `100 - used` and made fresh accounts look
+// exhausted, tripping auto-pause and excluding them from scheduling.
+func TestBuildCodexUsageExtraUpdates_FreshAccountUsedPercentNotInverted_Issue2994(t *testing.T) {
+ secondaryUsed := 1.0 // 5h window: barely used
+ secondaryWindow := 300
+ primaryUsed := 2.0 // 7d window: barely used
+ primaryWindow := 10080
+
+ snapshot := &OpenAICodexUsageSnapshot{
+ PrimaryUsedPercent: &primaryUsed,
+ PrimaryWindowMinutes: &primaryWindow,
+ SecondaryUsedPercent: &secondaryUsed,
+ SecondaryWindowMinutes: &secondaryWindow,
+ UpdatedAt: "2026-02-16T10:00:00Z",
+ }
+
+ updates := buildCodexUsageExtraUpdates(snapshot, time.Date(2026, 2, 16, 10, 0, 0, 0, time.UTC))
+ if updates == nil {
+ t.Fatal("expected non-nil updates")
+ }
+
+ if got := updates["codex_5h_used_percent"]; got != 1.0 {
+ t.Fatalf("codex_5h_used_percent = %v, want 1.0 (direct used%%, NOT inverted to 99)", got)
+ }
+ if got := updates["codex_7d_used_percent"]; got != 2.0 {
+ t.Fatalf("codex_7d_used_percent = %v, want 2.0 (direct used%%, NOT inverted to 98)", got)
+ }
+}
+
func TestBuildCodexUsageExtraUpdates_FallbackToNowWhenUpdatedAtInvalid(t *testing.T) {
primaryUsed := 15.0
primaryReset := 30
From ddf063352a300a3be462cd2b92fe0a28207c2042 Mon Sep 17 00:00:00 2001
From: DaydreamCoding
Date: Mon, 1 Jun 2026 16:10:35 +0800
Subject: [PATCH 76/79] =?UTF-8?q?feat(ops):=20=E9=94=99=E8=AF=AF=E6=97=A5?=
=?UTF-8?q?=E5=BF=97=20key=20=E5=BD=92=E5=9B=A0=E4=B8=8E=E6=97=A9=E9=80=80?=
=?UTF-8?q?=E5=AD=97=E6=AE=B5=E8=A1=A5=E5=85=A8?=
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit
让 /admin/ops 错误详情正确归因 API key 并补全早退场景字段,合并三项改动:
- 鉴权早退补全用户/分组/平台字段:引入 ops fallback key(ContextKeyOpsFallbackAPIKey),
apiKey 一加载成功即写入,覆盖分组停用/删除、Key 停用/过期/额度、用户停用、IP 限制等早退
路径;ops 错误日志改用 getOpsAPIKey(正式 key 优先、回退键兜底),不改「已鉴权」语义。
- 已删除 key 归因(迁移 145):删除 key 时同一事务写 deleted_api_key_audits 映射,认证失败
时用明文反查命中原所有者,错误详情展示「已删除 Key 所有者」「尝试的 Key 前缀」。
- 有效 key 报错快照前缀(迁移 147):对绑定有效 key 的错误,落库时快照明文前 8 位到
api_key_prefix(与 attempted_key_prefix 互斥),key 之后被删仍保留报错当时真实前缀。
均仅对上线后新产生的错误/删除生效。
Co-Authored-By: Claude Opus 4.8 (1M context)
---
backend/internal/handler/ops_error_logger.go | 80 ++++++++++-
.../ops_error_logger_attribution_test.go | 118 +++++++++++++++
.../internal/handler/ops_error_logger_test.go | 42 ++++++
backend/internal/repository/api_key_repo.go | 61 ++++++++
.../api_key_repo_integration_test.go | 43 ++++++
backend/internal/repository/ops_repo.go | 51 ++++++-
...po_get_error_log_by_id_integration_test.go | 94 ++++++++++++
...okup_deleted_key_audit_integration_test.go | 36 +++++
backend/internal/server/api_contract_test.go | 4 +
.../server/middleware/api_key_auth.go | 24 ++++
.../server/middleware/api_key_auth_google.go | 4 +
.../middleware/api_key_auth_google_test.go | 3 +
.../server/middleware/api_key_auth_test.go | 136 ++++++++++++++++++
.../internal/server/middleware/middleware.go | 5 +
.../service/admin_service_apikey_test.go | 3 +
backend/internal/service/api_key_service.go | 13 +-
.../service/api_key_service_cache_test.go | 4 +
.../service/api_key_service_delete_test.go | 16 ++-
.../service/api_key_service_quota_test.go | 3 +
backend/internal/service/ops_models.go | 9 ++
backend/internal/service/ops_port.go | 17 +++
.../internal/service/ops_repo_mock_test.go | 8 ++
backend/internal/service/ops_service.go | 8 ++
.../migrations/145_deleted_api_key_audit.sql | 22 +++
.../147_ops_error_log_api_key_prefix.sql | 12 ++
frontend/src/api/admin/ops.ts | 9 ++
frontend/src/i18n/locales/en.ts | 6 +-
frontend/src/i18n/locales/zh.ts | 6 +-
.../ops/components/OpsErrorDetailModal.vue | 25 ++++
29 files changed, 845 insertions(+), 17 deletions(-)
create mode 100644 backend/internal/handler/ops_error_logger_attribution_test.go
create mode 100644 backend/internal/repository/ops_repo_get_error_log_by_id_integration_test.go
create mode 100644 backend/internal/repository/ops_repo_lookup_deleted_key_audit_integration_test.go
create mode 100644 backend/migrations/145_deleted_api_key_audit.sql
create mode 100644 backend/migrations/147_ops_error_log_api_key_prefix.sql
diff --git a/backend/internal/handler/ops_error_logger.go b/backend/internal/handler/ops_error_logger.go
index 168fc271..b86c7f69 100644
--- a/backend/internal/handler/ops_error_logger.go
+++ b/backend/internal/handler/ops_error_logger.go
@@ -71,6 +71,49 @@ const (
opsErrorLogBatchSize = 32
)
+// looksLikeSystemKey 粗筛"形似本系统 key"的输入:长度 16-128 且仅含 [a-zA-Z0-9_-]。
+// 不用前缀匹配(APIKeyPrefix 可配置)。用于反查审计表前挡掉随机扫描的乱码输入。
+func looksLikeSystemKey(key string) bool {
+ if len(key) < 16 || len(key) > 128 {
+ return false
+ }
+ for _, c := range key {
+ allowed := (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') ||
+ (c >= '0' && c <= '9') || c == '_' || c == '-'
+ if !allowed {
+ return false
+ }
+ }
+ return true
+}
+
+// keyPrefix 返回脱敏前缀(前 n 个字符);不足 n 则原样返回。
+func keyPrefix(key string, n int) string {
+ if len(key) <= n {
+ return key
+ }
+ return key[:n]
+}
+
+// extractAttemptedKey 按认证中间件同样的顺序从请求头提取提交的 key 明文。
+// 与 api_key_auth.go:43-59 一致:Authorization 仅取 Bearer scheme,非 Bearer 则忽略并继续 x-api-key → x-goog-api-key。
+func extractAttemptedKey(c *gin.Context) string {
+ if h := c.GetHeader("Authorization"); h != "" {
+ parts := strings.SplitN(h, " ", 2)
+ if len(parts) == 2 && strings.EqualFold(parts[0], "Bearer") {
+ return strings.TrimSpace(parts[1])
+ }
+ // 非 Bearer:与中间件一致,忽略 Authorization,继续尝试其它 header(不在此 return)。
+ }
+ if k := c.GetHeader("x-api-key"); k != "" {
+ return strings.TrimSpace(k)
+ }
+ if k := c.GetHeader("x-goog-api-key"); k != "" {
+ return strings.TrimSpace(k)
+ }
+ return ""
+}
+
type opsErrorLogJob struct {
ops *service.OpsService
entry *service.OpsInsertErrorLogInput
@@ -546,7 +589,7 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
return
}
- apiKey, _ := middleware2.GetAPIKeyFromContext(c)
+ apiKey := getOpsAPIKey(c)
clientRequestID, _ := c.Request.Context().Value(ctxkey.ClientRequestID).(string)
model, _ := c.Get(opsModelKey)
@@ -721,6 +764,7 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
if apiKey != nil {
entry.APIKeyID = &apiKey.ID
+ entry.APIKeyPrefix = keyPrefix(apiKey.Key, 8)
if apiKey.User != nil {
entry.UserID = &apiKey.User.ID
}
@@ -765,7 +809,7 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
return
}
- apiKey, _ := middleware2.GetAPIKeyFromContext(c)
+ apiKey := getOpsAPIKey(c)
clientRequestID, _ := c.Request.Context().Value(ctxkey.ClientRequestID).(string)
@@ -911,6 +955,8 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
if apiKey != nil {
entry.APIKeyID = &apiKey.ID
+ // 有效(未删除)key 报错时快照前缀,key 之后被删也保留;与 INVALID_API_KEY 的 attempted_key_prefix 互斥。
+ entry.APIKeyPrefix = keyPrefix(apiKey.Key, 8)
if apiKey.User != nil {
entry.UserID = &apiKey.User.ID
}
@@ -929,6 +975,22 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
entry.ClientIP = &clientIP
}
+ // 已删除 key 归因:仅 INVALID_API_KEY 才尝试。响应已写出,此处不阻塞客户端。
+ if parsed.Code == opsCodeInvalidAPIKey {
+ if attemptedKey := extractAttemptedKey(c); attemptedKey != "" {
+ entry.AttemptedKeyPrefix = keyPrefix(attemptedKey, 8)
+ if looksLikeSystemKey(attemptedKey) {
+ if res, lookupErr := ops.LookupDeletedKeyAudit(c.Request.Context(), attemptedKey); lookupErr != nil {
+ log.Printf("[OpsErrorLogger] LookupDeletedKeyAudit failed: %v", lookupErr)
+ } else if res != nil {
+ owner := res.UserID
+ entry.DeletedKeyOwnerUserID = &owner
+ entry.DeletedKeyName = res.KeyName
+ }
+ }
+ }
+ }
+
enqueueOpsErrorLog(ops, entry)
}
}
@@ -1035,6 +1097,20 @@ func parseOpsErrorResponse(body []byte) parsedOpsError {
return parsedOpsError{Message: truncateString(string(body), 1024)}
}
+// getOpsAPIKey 返回用于 Ops 错误日志的 API Key:优先取已鉴权写入的正式 key;
+// 鉴权早退(分组停用/删除、Key 停用/过期/额度、用户停用、IP 限制等)时,
+// 正式 key 尚未写入,回退到 middleware 写入的 ops fallback key
+// (含 User/Group/Platform),从而让日志能展示 用户/分组/平台。
+func getOpsAPIKey(c *gin.Context) *service.APIKey {
+ if apiKey, ok := middleware2.GetAPIKeyFromContext(c); ok && apiKey != nil {
+ return apiKey
+ }
+ if apiKey, ok := middleware2.GetOpsFallbackAPIKey(c); ok && apiKey != nil {
+ return apiKey
+ }
+ return nil
+}
+
func resolveOpsPlatform(apiKey *service.APIKey, fallback string) string {
if apiKey != nil && apiKey.Group != nil && apiKey.Group.Platform != "" {
return apiKey.Group.Platform
diff --git a/backend/internal/handler/ops_error_logger_attribution_test.go b/backend/internal/handler/ops_error_logger_attribution_test.go
new file mode 100644
index 00000000..9c68d845
--- /dev/null
+++ b/backend/internal/handler/ops_error_logger_attribution_test.go
@@ -0,0 +1,118 @@
+package handler
+
+import (
+ "net/http"
+ "net/http/httptest"
+ "testing"
+
+ "github.com/gin-gonic/gin"
+)
+
+func TestLooksLikeSystemKey(t *testing.T) {
+ cases := []struct {
+ in string
+ want bool
+ }{
+ {"sk-abcdef0123456789", true},
+ {"ABCdef_-0123456789", true},
+ {"short", false},
+ {"with space xxxxxxxxxx", false},
+ {"汉字key1234567890", false},
+ {"", false},
+ }
+ for _, c := range cases {
+ if got := looksLikeSystemKey(c.in); got != c.want {
+ t.Errorf("looksLikeSystemKey(%q)=%v want %v", c.in, got, c.want)
+ }
+ }
+ long := make([]byte, 129)
+ for i := range long {
+ long[i] = 'a'
+ }
+ if looksLikeSystemKey(string(long)) {
+ t.Errorf("129-char key should be rejected")
+ }
+}
+
+func TestKeyPrefix(t *testing.T) {
+ if got := keyPrefix("sk-3f2a9c7e", 8); got != "sk-3f2a9" {
+ t.Errorf("keyPrefix=%q want %q", got, "sk-3f2a9")
+ }
+ if got := keyPrefix("abc", 8); got != "abc" {
+ t.Errorf("short key should be returned as-is, got %q", got)
+ }
+}
+
+func TestExtractAttemptedKey(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ cases := []struct {
+ name string
+ headers map[string]string
+ want string
+ }{
+ {
+ name: "Bearer in Authorization",
+ headers: map[string]string{"Authorization": "Bearer sk-testkey0123456789"},
+ want: "sk-testkey0123456789",
+ },
+ {
+ name: "Bearer case-insensitive",
+ headers: map[string]string{"Authorization": "BEARER sk-testkey0123456789"},
+ want: "sk-testkey0123456789",
+ },
+ {
+ name: "x-api-key header",
+ headers: map[string]string{"x-api-key": "sk-xapikey0123456789"},
+ want: "sk-xapikey0123456789",
+ },
+ {
+ name: "x-goog-api-key header",
+ headers: map[string]string{"x-goog-api-key": "sk-goog0123456789"},
+ want: "sk-goog0123456789",
+ },
+ {
+ name: "Authorization takes priority over x-api-key",
+ headers: map[string]string{"Authorization": "Bearer sk-auth0123456789", "x-api-key": "sk-xapi0123456789"},
+ want: "sk-auth0123456789",
+ },
+ {
+ name: "x-api-key takes priority over x-goog-api-key",
+ headers: map[string]string{"x-api-key": "sk-xapi0123456789", "x-goog-api-key": "sk-goog0123456789"},
+ want: "sk-xapi0123456789",
+ },
+ {
+ name: "no key headers",
+ headers: map[string]string{},
+ want: "",
+ },
+ {
+ name: "Bearer with leading/trailing spaces trimmed",
+ headers: map[string]string{"Authorization": "Bearer sk-trimmed0123456789 "},
+ want: "sk-trimmed0123456789",
+ },
+ {
+ // 非 Bearer Authorization 应被忽略,继续 fall-through 到 x-api-key(与认证中间件一致)
+ name: "non-Bearer Authorization falls through to x-api-key",
+ headers: map[string]string{"Authorization": "junk-not-bearer", "x-api-key": "sk-realkey1234567"},
+ want: "sk-realkey1234567",
+ },
+ }
+
+ for _, tc := range cases {
+ t.Run(tc.name, func(t *testing.T) {
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil)
+ for k, v := range tc.headers {
+ req.Header.Set(k, v)
+ }
+ c.Request = req
+
+ got := extractAttemptedKey(c)
+ if got != tc.want {
+ t.Errorf("extractAttemptedKey(%v) = %q, want %q", tc.headers, got, tc.want)
+ }
+ })
+ }
+}
diff --git a/backend/internal/handler/ops_error_logger_test.go b/backend/internal/handler/ops_error_logger_test.go
index d4e1177e..cf1685f2 100644
--- a/backend/internal/handler/ops_error_logger_test.go
+++ b/backend/internal/handler/ops_error_logger_test.go
@@ -931,3 +931,45 @@ func TestSetOpsEndpointContext_NilContext(t *testing.T) {
setOpsEndpointContext(nil, "model", int16(1))
})
}
+
+func TestGetOpsAPIKeyFallsBackToOpsFallbackKey(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+
+ // 主 key 缺席(鉴权早退场景):返回 nil。
+ require.Nil(t, getOpsAPIKey(c))
+
+ // 写入 ops 专用 fallback key 后应能取到,且带齐 user/group。
+ groupID := int64(55)
+ apiKey := &service.APIKey{
+ ID: 100,
+ GroupID: &groupID,
+ User: &service.User{ID: 7},
+ Group: &service.Group{ID: groupID, Platform: service.PlatformAnthropic},
+ }
+ c.Set(string(middleware2.ContextKeyOpsFallbackAPIKey), apiKey)
+
+ got := getOpsAPIKey(c)
+ require.NotNil(t, got)
+ require.Equal(t, int64(100), got.ID)
+ require.NotNil(t, got.User)
+ require.Equal(t, int64(7), got.User.ID)
+ require.NotNil(t, got.Group)
+ require.Equal(t, service.PlatformAnthropic, got.Group.Platform)
+}
+
+func TestGetOpsAPIKeyPrefersPrimaryContextKey(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+
+ primary := &service.APIKey{ID: 1}
+ fallback := &service.APIKey{ID: 2}
+ c.Set(string(middleware2.ContextKeyAPIKey), primary)
+ c.Set(string(middleware2.ContextKeyOpsFallbackAPIKey), fallback)
+
+ got := getOpsAPIKey(c)
+ require.NotNil(t, got)
+ require.Equal(t, int64(1), got.ID, "已鉴权请求应优先使用正式 api key")
+}
diff --git a/backend/internal/repository/api_key_repo.go b/backend/internal/repository/api_key_repo.go
index 7db35ecc..18f6878b 100644
--- a/backend/internal/repository/api_key_repo.go
+++ b/backend/internal/repository/api_key_repo.go
@@ -3,6 +3,7 @@ package repository
import (
"context"
"database/sql"
+ "errors"
"fmt"
"strings"
"time"
@@ -304,6 +305,66 @@ func (r *apiKeyRepository) Delete(ctx context.Context, id int64) error {
return nil
}
+// DeleteWithAudit 在同一事务内:
+// 1. 把(明文 key、所有者、key 名称)写入 deleted_api_key_audits;
+// 2. 软删除该 key(tombstone 覆盖 key 列以释放唯一约束)。
+//
+// 保证"被删除的 key 一定能反查到所有者"。事务模式与 group_repo.DeleteCascade 一致。
+func (r *apiKeyRepository) DeleteWithAudit(ctx context.Context, id int64) error {
+ tombstoneKey := fmt.Sprintf("__deleted__%d__%d", id, time.Now().UnixNano())
+
+ tx, err := r.client.Tx(ctx)
+ if err != nil && !errors.Is(err, dbent.ErrTxStarted) {
+ return err
+ }
+ exec := r.client
+ if err == nil {
+ defer func() { _ = tx.Rollback() }()
+ exec = tx.Client()
+ }
+ // err == dbent.ErrTxStarted 时复用当前事务(exec = r.client)。
+
+ // 1. 审计:数据源即 api_keys 当前行;WHERE deleted_at IS NULL 保证只对未删除行写一次。
+ if _, err := exec.ExecContext(ctx, `
+ INSERT INTO deleted_api_key_audits (key, api_key_id, user_id, key_name, deleted_at)
+ SELECT key, id, user_id, name, NOW()
+ FROM api_keys
+ WHERE id = $1 AND deleted_at IS NULL`, id); err != nil {
+ return err
+ }
+
+ // 2. 软删除(tombstone 覆盖 key)。
+ res, err := exec.ExecContext(ctx, `
+ UPDATE api_keys
+ SET key = $1, deleted_at = NOW(), updated_at = NOW()
+ WHERE id = $2 AND deleted_at IS NULL`, tombstoneKey, id)
+ if err != nil {
+ return err
+ }
+ affected, err := res.RowsAffected()
+ if err != nil {
+ return err
+ }
+ if affected == 0 {
+ // 并发/重复删除:记录已存在(已软删)则幂等返回 nil(defer 回滚空事务),否则 NotFound。
+ exists, existErr := r.client.APIKey.Query().
+ Where(apikey.IDEQ(id)).
+ Exist(mixins.SkipSoftDelete(ctx))
+ if existErr != nil {
+ return existErr
+ }
+ if exists {
+ return nil
+ }
+ return service.ErrAPIKeyNotFound
+ }
+
+ if tx != nil {
+ return tx.Commit()
+ }
+ return nil
+}
+
func (r *apiKeyRepository) ListByUserID(ctx context.Context, userID int64, params pagination.PaginationParams, filters service.APIKeyListFilters) ([]service.APIKey, *pagination.PaginationResult, error) {
q := r.activeQuery().Where(apikey.UserIDEQ(userID))
diff --git a/backend/internal/repository/api_key_repo_integration_test.go b/backend/internal/repository/api_key_repo_integration_test.go
index e926ed86..fdf9bc83 100644
--- a/backend/internal/repository/api_key_repo_integration_test.go
+++ b/backend/internal/repository/api_key_repo_integration_test.go
@@ -555,3 +555,46 @@ func TestIncrementQuotaUsed_Concurrent(t *testing.T) {
require.Equal(t, float64(goroutines)*increment, got.QuotaUsed,
"并发递增后总和应为 %v,实际为 %v", float64(goroutines)*increment, got.QuotaUsed)
}
+
+func (s *APIKeyRepoSuite) TestDeleteWithAudit_WritesAuditAndSoftDeletes() {
+ user := s.mustCreateUser("delwithaudit@test.com")
+ key := &service.APIKey{
+ UserID: user.ID,
+ Key: "sk-del-audit-1",
+ Name: "Audit Me",
+ Status: service.StatusActive,
+ }
+ s.Require().NoError(s.repo.Create(s.ctx, key))
+
+ s.Require().NoError(s.repo.DeleteWithAudit(s.ctx, key.ID))
+
+ _, err := s.repo.GetByID(s.ctx, key.ID)
+ s.Require().Error(err)
+
+ rows, qErr := s.client.QueryContext(s.ctx,
+ `SELECT key, key_name, user_id, api_key_id FROM deleted_api_key_audits WHERE api_key_id = $1`, key.ID)
+ s.Require().NoError(qErr)
+ defer rows.Close()
+ s.Require().True(rows.Next(), "expected one audit row")
+ var auditKey, auditName string
+ var auditUserID, auditAPIKeyID int64
+ s.Require().NoError(rows.Scan(&auditKey, &auditName, &auditUserID, &auditAPIKeyID))
+ s.Require().Equal("sk-del-audit-1", auditKey)
+ s.Require().Equal("Audit Me", auditName)
+ s.Require().Equal(user.ID, auditUserID)
+ s.Require().Equal(key.ID, auditAPIKeyID)
+}
+
+func (s *APIKeyRepoSuite) TestDeleteWithAudit_RepeatIsIdempotent() {
+ user := s.mustCreateUser("delwithaudit-idem@test.com")
+ key := &service.APIKey{UserID: user.ID, Key: "sk-del-audit-2", Name: "K", Status: service.StatusActive}
+ s.Require().NoError(s.repo.Create(s.ctx, key))
+
+ s.Require().NoError(s.repo.DeleteWithAudit(s.ctx, key.ID))
+ s.Require().NoError(s.repo.DeleteWithAudit(s.ctx, key.ID))
+}
+
+func (s *APIKeyRepoSuite) TestDeleteWithAudit_NotFound() {
+ err := s.repo.DeleteWithAudit(s.ctx, 999999)
+ s.Require().ErrorIs(err, service.ErrAPIKeyNotFound)
+}
diff --git a/backend/internal/repository/ops_repo.go b/backend/internal/repository/ops_repo.go
index 4371b8a2..a7773713 100644
--- a/backend/internal/repository/ops_repo.go
+++ b/backend/internal/repository/ops_repo.go
@@ -4,6 +4,7 @@ import (
"context"
"database/sql"
"encoding/json"
+ "errors"
"fmt"
"strings"
"time"
@@ -54,9 +55,13 @@ INSERT INTO ops_error_logs (
upstream_latency_ms,
response_latency_ms,
time_to_first_token_ms,
- created_at
+ created_at,
+ attempted_key_prefix,
+ deleted_key_owner_user_id,
+ deleted_key_name,
+ api_key_prefix
) VALUES (
- $1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23,$24,$25,$26,$27,$28,$29,$30,$31,$32,$33,$34,$35,$36,$37
+ $1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23,$24,$25,$26,$27,$28,$29,$30,$31,$32,$33,$34,$35,$36,$37,$38,$39,$40,$41
)`
func NewOpsRepository(db *sql.DB) service.OpsRepository {
@@ -165,6 +170,10 @@ func opsInsertErrorLogArgs(input *service.OpsInsertErrorLogInput) []any {
opsNullInt64(input.ResponseLatencyMs),
opsNullInt64(input.TimeToFirstTokenMs),
input.CreatedAt,
+ opsNullString(input.AttemptedKeyPrefix),
+ opsNullInt64(input.DeletedKeyOwnerUserID),
+ opsNullString(input.DeletedKeyName),
+ opsNullString(input.APIKeyPrefix),
}
}
@@ -402,11 +411,17 @@ SELECT
e.routing_latency_ms,
e.upstream_latency_ms,
e.response_latency_ms,
- e.time_to_first_token_ms
+ e.time_to_first_token_ms,
+ COALESCE(e.attempted_key_prefix, ''),
+ e.deleted_key_owner_user_id,
+ COALESCE(du.email, ''),
+ COALESCE(e.deleted_key_name, ''),
+ COALESCE(e.api_key_prefix, '')
FROM ops_error_logs e
LEFT JOIN users u ON e.user_id = u.id
LEFT JOIN accounts a ON e.account_id = a.id
LEFT JOIN groups g ON e.group_id = g.id
+LEFT JOIN users du ON e.deleted_key_owner_user_id = du.id
WHERE e.id = $1
LIMIT 1`
@@ -426,6 +441,7 @@ LIMIT 1`
var responseLatency sql.NullInt64
var ttft sql.NullInt64
var requestType sql.NullInt64
+ var deletedKeyOwnerUserID sql.NullInt64
err := r.db.QueryRowContext(ctx, q, id).Scan(
&out.ID,
@@ -471,6 +487,11 @@ LIMIT 1`
&upstreamLatency,
&responseLatency,
&ttft,
+ &out.AttemptedKeyPrefix,
+ &deletedKeyOwnerUserID,
+ &out.DeletedKeyOwnerEmail,
+ &out.DeletedKeyName,
+ &out.APIKeyPrefix,
)
if err != nil {
return nil, err
@@ -533,6 +554,10 @@ LIMIT 1`
v := int16(requestType.Int64)
out.RequestType = &v
}
+ if deletedKeyOwnerUserID.Valid {
+ v := deletedKeyOwnerUserID.Int64
+ out.DeletedKeyOwnerUserID = &v
+ }
// Normalize upstream_errors to empty string when stored as JSON null.
out.UpstreamErrors = strings.TrimSpace(out.UpstreamErrors)
@@ -543,6 +568,26 @@ LIMIT 1`
return &out, nil
}
+// LookupDeletedKeyAudit 按明文 key 反查最近一条已删除 key 审计。
+// 同一 key 可能有多条历史(反复创建/删除),取 deleted_at 最近一条(id 作同毫秒 tiebreaker)。
+// 未命中返回 (nil, nil)。
+func (r *opsRepository) LookupDeletedKeyAudit(ctx context.Context, key string) (*service.DeletedKeyAuditResult, error) {
+ var res service.DeletedKeyAuditResult
+ err := r.db.QueryRowContext(ctx, `
+ SELECT user_id, key_name
+ FROM deleted_api_key_audits
+ WHERE key = $1
+ ORDER BY deleted_at DESC, id DESC
+ LIMIT 1`, key).Scan(&res.UserID, &res.KeyName)
+ if err != nil {
+ if errors.Is(err, sql.ErrNoRows) {
+ return nil, nil
+ }
+ return nil, err
+ }
+ return &res, nil
+}
+
func (r *opsRepository) UpdateErrorResolution(ctx context.Context, errorID int64, resolved bool, resolvedByUserID *int64, resolvedAt *time.Time) error {
if r == nil || r.db == nil {
return fmt.Errorf("nil ops repository")
diff --git a/backend/internal/repository/ops_repo_get_error_log_by_id_integration_test.go b/backend/internal/repository/ops_repo_get_error_log_by_id_integration_test.go
new file mode 100644
index 00000000..470b1c0d
--- /dev/null
+++ b/backend/internal/repository/ops_repo_get_error_log_by_id_integration_test.go
@@ -0,0 +1,94 @@
+//go:build integration
+
+package repository
+
+import (
+ "context"
+ "testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/service"
+ "github.com/stretchr/testify/require"
+)
+
+// TestGetErrorLogByID_DeletedKeyOwner 验证:
+// 1. 带 deleted_key_owner_user_id 的记录能正确 JOIN users 返回 DeletedKeyOwnerEmail
+// 2. 新列全为 NULL 的普通记录 Scan 不报错,这些字段为空/nil
+func TestGetErrorLogByID_DeletedKeyOwner(t *testing.T) {
+ ctx := context.Background()
+ _, _ = integrationDB.ExecContext(ctx, "TRUNCATE ops_error_logs RESTART IDENTITY CASCADE")
+
+ repo := NewOpsRepository(integrationDB).(*opsRepository)
+
+ // ── Case 1: 带 deleted_key_owner 信息的记录 ──────────────────────────────
+ owner := mustCreateUser(t, integrationEntClient, &service.User{
+ Email: "deleted-key-owner-" + time.Now().Format("150405.000000000") + "@example.com",
+ })
+
+ var insertedID int64
+ err := integrationDB.QueryRowContext(ctx, `
+ INSERT INTO ops_error_logs (
+ error_phase, error_type, severity, status_code, created_at,
+ attempted_key_prefix, deleted_key_owner_user_id, deleted_key_name
+ ) VALUES (
+ 'auth', 'INVALID_API_KEY', 'error', 401, NOW(),
+ 'sk-test-abc', $1, 'my-deleted-key'
+ ) RETURNING id`,
+ owner.ID,
+ ).Scan(&insertedID)
+ require.NoError(t, err)
+ require.Positive(t, insertedID)
+
+ detail, err := repo.GetErrorLogByID(ctx, insertedID)
+ require.NoError(t, err)
+ require.NotNil(t, detail)
+
+ require.Equal(t, "sk-test-abc", detail.AttemptedKeyPrefix)
+ require.NotNil(t, detail.DeletedKeyOwnerUserID)
+ require.Equal(t, owner.ID, *detail.DeletedKeyOwnerUserID)
+ require.Equal(t, owner.Email, detail.DeletedKeyOwnerEmail)
+ require.Equal(t, "my-deleted-key", detail.DeletedKeyName)
+
+ // ── Case 2: 新列全为 NULL 的普通错误记录 ──────────────────────────────────
+ var plainID int64
+ err = integrationDB.QueryRowContext(ctx, `
+ INSERT INTO ops_error_logs (
+ error_phase, error_type, severity, status_code, created_at
+ ) VALUES (
+ 'upstream', 'upstream_error', 'error', 500, NOW()
+ ) RETURNING id`,
+ ).Scan(&plainID)
+ require.NoError(t, err)
+ require.Positive(t, plainID)
+
+ plain, err := repo.GetErrorLogByID(ctx, plainID)
+ require.NoError(t, err)
+ require.NotNil(t, plain)
+
+ require.Empty(t, plain.AttemptedKeyPrefix, "no prefix for plain error")
+ require.Nil(t, plain.DeletedKeyOwnerUserID, "no owner for plain error")
+ require.Empty(t, plain.DeletedKeyOwnerEmail, "no owner email for plain error")
+ require.Empty(t, plain.DeletedKeyName, "no key name for plain error")
+ require.Empty(t, plain.APIKeyPrefix, "no api key prefix for plain error")
+
+ // ── Case 3: 有效(未删除)key 报错,经 InsertErrorLog 快照 api_key_prefix ──────
+ // 走真实 InsertErrorLog 写入路径(覆盖新列 + $41 占位符),再 GetErrorLogByID 读回。
+ validID, err := repo.InsertErrorLog(ctx, &service.OpsInsertErrorLogInput{
+ ErrorPhase: "request",
+ ErrorType: "api_error",
+ Severity: "error",
+ StatusCode: 402,
+ CreatedAt: time.Now(),
+ APIKeyPrefix: "sk-valid",
+ })
+ require.NoError(t, err)
+ require.Positive(t, validID)
+
+ valid, err := repo.GetErrorLogByID(ctx, validID)
+ require.NoError(t, err)
+ require.NotNil(t, valid)
+
+ require.Equal(t, "sk-valid", valid.APIKeyPrefix)
+ require.Empty(t, valid.AttemptedKeyPrefix, "attempted prefix and api key prefix are mutually exclusive")
+ require.Nil(t, valid.DeletedKeyOwnerUserID, "valid key error has no deleted owner")
+}
diff --git a/backend/internal/repository/ops_repo_lookup_deleted_key_audit_integration_test.go b/backend/internal/repository/ops_repo_lookup_deleted_key_audit_integration_test.go
new file mode 100644
index 00000000..c77aefb9
--- /dev/null
+++ b/backend/internal/repository/ops_repo_lookup_deleted_key_audit_integration_test.go
@@ -0,0 +1,36 @@
+//go:build integration
+
+package repository
+
+import (
+ "context"
+ "testing"
+ "time"
+
+ "github.com/stretchr/testify/require"
+)
+
+func TestOpsRepositoryLookupDeletedKeyAudit(t *testing.T) {
+ ctx := context.Background()
+ _, _ = integrationDB.ExecContext(ctx, "TRUNCATE deleted_api_key_audits RESTART IDENTITY")
+ repo := NewOpsRepository(integrationDB).(*opsRepository)
+
+ // 同一 key 两条审计,取最近一条(deleted_at DESC, id DESC)
+ _, err := integrationDB.ExecContext(ctx, `
+ INSERT INTO deleted_api_key_audits (key, api_key_id, user_id, key_name, deleted_at)
+ VALUES ('sk-lookup-1', 10, 100, 'old', $1),
+ ('sk-lookup-1', 11, 200, 'new', $2)`,
+ time.Now().Add(-time.Hour), time.Now())
+ require.NoError(t, err)
+
+ res, err := repo.LookupDeletedKeyAudit(ctx, "sk-lookup-1")
+ require.NoError(t, err)
+ require.NotNil(t, res)
+ require.Equal(t, int64(200), res.UserID)
+ require.Equal(t, "new", res.KeyName)
+
+ // 未命中返回 nil
+ miss, err := repo.LookupDeletedKeyAudit(ctx, "sk-never-existed")
+ require.NoError(t, err)
+ require.Nil(t, miss)
+}
diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go
index 6bb87995..d4901da1 100644
--- a/backend/internal/server/api_contract_test.go
+++ b/backend/internal/server/api_contract_test.go
@@ -2102,6 +2102,10 @@ func (r *stubApiKeyRepo) Delete(ctx context.Context, id int64) error {
return nil
}
+func (r *stubApiKeyRepo) DeleteWithAudit(ctx context.Context, id int64) error {
+ return r.Delete(ctx, id)
+}
+
func (r *stubApiKeyRepo) ListByUserID(ctx context.Context, userID int64, params pagination.PaginationParams, _ service.APIKeyListFilters) ([]service.APIKey, *pagination.PaginationResult, error) {
ids := make([]int64, 0, len(r.byID))
for id := range r.byID {
diff --git a/backend/internal/server/middleware/api_key_auth.go b/backend/internal/server/middleware/api_key_auth.go
index d33ccbf5..ba43d126 100644
--- a/backend/internal/server/middleware/api_key_auth.go
+++ b/backend/internal/server/middleware/api_key_auth.go
@@ -76,6 +76,10 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
return
}
+ // apiKey 已加载(含 User/Group)。即便后续因分组停用/Key 停用/用户停用/
+ // IP 限制等早退中断,也让 Ops 错误日志能回退取到 user/group/platform。
+ SetOpsFallbackAPIKey(c, apiKey)
+
// ── 3. 基础鉴权(始终执行) ─────────────────────────────────
// disabled / 未知状态 → 无条件拦截(expired 和 quota_exhausted 留给计费阶段)
@@ -237,6 +241,26 @@ func GetAPIKeyFromContext(c *gin.Context) (*service.APIKey, bool) {
return apiKey, ok
}
+// SetOpsFallbackAPIKey 记录已加载的 API Key,供 Ops 错误日志在鉴权早退时回退使用。
+// 与 ContextKeyAPIKey 区分:写入它不代表请求已通过鉴权,因此不影响 handler、
+// 审计日志等对“已鉴权”的判断。
+func SetOpsFallbackAPIKey(c *gin.Context, apiKey *service.APIKey) {
+ if c == nil || apiKey == nil {
+ return
+ }
+ c.Set(string(ContextKeyOpsFallbackAPIKey), apiKey)
+}
+
+// GetOpsFallbackAPIKey 读取 Ops 错误日志专用的回退 API Key。
+func GetOpsFallbackAPIKey(c *gin.Context) (*service.APIKey, bool) {
+ value, exists := c.Get(string(ContextKeyOpsFallbackAPIKey))
+ if !exists {
+ return nil, false
+ }
+ apiKey, ok := value.(*service.APIKey)
+ return apiKey, ok
+}
+
// GetSubscriptionFromContext 从上下文中获取订阅信息
func GetSubscriptionFromContext(c *gin.Context) (*service.UserSubscription, bool) {
value, exists := c.Get(string(ContextKeySubscription))
diff --git a/backend/internal/server/middleware/api_key_auth_google.go b/backend/internal/server/middleware/api_key_auth_google.go
index 596bed52..97f3936c 100644
--- a/backend/internal/server/middleware/api_key_auth_google.go
+++ b/backend/internal/server/middleware/api_key_auth_google.go
@@ -42,6 +42,10 @@ func APIKeyAuthWithSubscriptionGoogle(apiKeyService *service.APIKeyService, subs
return
}
+ // 同 api_key_auth.go:早退中断前也写入 Ops 回退 key,便于错误日志展示
+ // user/group/platform。
+ SetOpsFallbackAPIKey(c, apiKey)
+
if !apiKey.IsActive() {
abortWithGoogleError(c, 401, "API key is disabled")
return
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 feadd27d..32e7e70f 100644
--- a/backend/internal/server/middleware/api_key_auth_google_test.go
+++ b/backend/internal/server/middleware/api_key_auth_google_test.go
@@ -56,6 +56,9 @@ func (f fakeAPIKeyRepo) Update(ctx context.Context, key *service.APIKey) error {
func (f fakeAPIKeyRepo) Delete(ctx context.Context, id int64) error {
return errors.New("not implemented")
}
+func (f fakeAPIKeyRepo) DeleteWithAudit(ctx context.Context, id int64) error {
+ return errors.New("not implemented")
+}
func (f fakeAPIKeyRepo) ListByUserID(ctx context.Context, userID int64, params pagination.PaginationParams, _ service.APIKeyListFilters) ([]service.APIKey, *pagination.PaginationResult, error) {
return nil, nil, errors.New("not implemented")
}
diff --git a/backend/internal/server/middleware/api_key_auth_test.go b/backend/internal/server/middleware/api_key_auth_test.go
index 76a24192..5d48bed2 100644
--- a/backend/internal/server/middleware/api_key_auth_test.go
+++ b/backend/internal/server/middleware/api_key_auth_test.go
@@ -419,6 +419,138 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
}
}
+func TestAPIKeyAuthSetsOpsFallbackKeyOnEarlyAbort(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ groupID := int64(101)
+ user := &service.User{
+ ID: 7,
+ Role: service.RoleUser,
+ Status: service.StatusActive,
+ Balance: 10,
+ Concurrency: 3,
+ }
+ apiKey := &service.APIKey{
+ ID: 100,
+ UserID: user.ID,
+ GroupID: &groupID,
+ Key: "test-key",
+ Status: service.StatusActive,
+ User: user,
+ Group: &service.Group{
+ ID: groupID,
+ Name: "disabled",
+ Status: service.StatusDisabled,
+ Platform: service.PlatformAnthropic,
+ Hydrated: true,
+ },
+ }
+ apiKeyRepo := &stubApiKeyRepo{
+ getByKey: func(ctx context.Context, key string) (*service.APIKey, error) {
+ if key != apiKey.Key {
+ return nil, service.ErrAPIKeyNotFound
+ }
+ clone := *apiKey
+ return &clone, nil
+ },
+ }
+ cfg := &config.Config{RunMode: config.RunModeStandard}
+ apiKeyService := service.NewAPIKeyService(apiKeyRepo, nil, nil, nil, nil, nil, cfg)
+
+ router := gin.New()
+ var fallback *service.APIKey
+ var fallbackOK bool
+ router.Use(func(c *gin.Context) {
+ c.Next()
+ fallback, fallbackOK = GetOpsFallbackAPIKey(c)
+ })
+ router.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(apiKeyService, nil, cfg)))
+ router.GET("/t", func(c *gin.Context) {
+ c.JSON(http.StatusOK, gin.H{"ok": true})
+ })
+
+ w := httptest.NewRecorder()
+ req := httptest.NewRequest(http.MethodGet, "/t", nil)
+ req.Header.Set("x-api-key", apiKey.Key)
+ router.ServeHTTP(w, req)
+
+ // 分组停用 → 早退中断,但 ops fallback key 仍应写入,含 user/group/platform。
+ require.Equal(t, http.StatusForbidden, w.Code)
+ require.Contains(t, w.Body.String(), "GROUP_DISABLED")
+ require.True(t, fallbackOK, "鉴权早退时也应写入 ops fallback api key")
+ require.NotNil(t, fallback)
+ require.Equal(t, apiKey.ID, fallback.ID)
+ require.NotNil(t, fallback.User)
+ require.Equal(t, user.ID, fallback.User.ID)
+ require.NotNil(t, fallback.GroupID)
+ require.Equal(t, groupID, *fallback.GroupID)
+ require.NotNil(t, fallback.Group)
+ require.Equal(t, service.PlatformAnthropic, fallback.Group.Platform)
+}
+
+func TestAPIKeyAuthGoogleSetsOpsFallbackKeyOnEarlyAbort(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ groupID := int64(202)
+ user := &service.User{
+ ID: 9,
+ Role: service.RoleUser,
+ Status: service.StatusActive,
+ Balance: 10,
+ Concurrency: 3,
+ }
+ apiKey := &service.APIKey{
+ ID: 200,
+ UserID: user.ID,
+ GroupID: &groupID,
+ Key: "g-key",
+ Status: service.StatusActive,
+ User: user,
+ Group: &service.Group{
+ ID: groupID,
+ Name: "disabled",
+ Status: service.StatusDisabled,
+ Platform: service.PlatformGemini,
+ Hydrated: true,
+ },
+ }
+ apiKeyRepo := &stubApiKeyRepo{
+ getByKey: func(ctx context.Context, key string) (*service.APIKey, error) {
+ if key != apiKey.Key {
+ return nil, service.ErrAPIKeyNotFound
+ }
+ clone := *apiKey
+ return &clone, nil
+ },
+ }
+ cfg := &config.Config{RunMode: config.RunModeStandard}
+ apiKeyService := service.NewAPIKeyService(apiKeyRepo, nil, nil, nil, nil, nil, cfg)
+
+ router := gin.New()
+ var fallback *service.APIKey
+ var fallbackOK bool
+ router.Use(func(c *gin.Context) {
+ c.Next()
+ fallback, fallbackOK = GetOpsFallbackAPIKey(c)
+ })
+ router.Use(gin.HandlerFunc(APIKeyAuthWithSubscriptionGoogle(apiKeyService, nil, cfg)))
+ router.GET("/t", func(c *gin.Context) {
+ c.JSON(http.StatusOK, gin.H{"ok": true})
+ })
+
+ w := httptest.NewRecorder()
+ req := httptest.NewRequest(http.MethodGet, "/t", nil)
+ req.Header.Set("x-goog-api-key", apiKey.Key)
+ router.ServeHTTP(w, req)
+
+ require.Equal(t, http.StatusForbidden, w.Code)
+ require.True(t, fallbackOK, "Google 鉴权早退时也应写入 ops fallback api key")
+ require.NotNil(t, fallback)
+ require.Equal(t, apiKey.ID, fallback.ID)
+ require.NotNil(t, fallback.User)
+ require.Equal(t, user.ID, fallback.User.ID)
+}
+
func TestRequireGroupAssignmentMarksUngroupedKeyBusinessLimited(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -761,6 +893,10 @@ func (r *stubApiKeyRepo) Delete(ctx context.Context, id int64) error {
return errors.New("not implemented")
}
+func (r *stubApiKeyRepo) DeleteWithAudit(ctx context.Context, id int64) error {
+ return errors.New("not implemented")
+}
+
func (r *stubApiKeyRepo) ListByUserID(ctx context.Context, userID int64, params pagination.PaginationParams, _ service.APIKeyListFilters) ([]service.APIKey, *pagination.PaginationResult, error) {
return nil, nil, errors.New("not implemented")
}
diff --git a/backend/internal/server/middleware/middleware.go b/backend/internal/server/middleware/middleware.go
index d42eacec..9efe78a3 100644
--- a/backend/internal/server/middleware/middleware.go
+++ b/backend/internal/server/middleware/middleware.go
@@ -24,6 +24,11 @@ const (
ContextKeySubscription ContextKey = "subscription"
// ContextKeyForcePlatform 强制平台(用于 /antigravity 路由)
ContextKeyForcePlatform ContextKey = "force_platform"
+ // ContextKeyOpsFallbackAPIKey 运维错误日志专用回退键。
+ // 鉴权早退(分组停用/删除、Key 停用/过期/额度、用户停用、IP 限制等)时,
+ // apiKey 已加载但尚未写入 ContextKeyAPIKey;该键让 Ops 错误日志仍能取到
+ // user/group/platform。仅供 Ops 错误日志读取,不代表请求已通过鉴权。
+ ContextKeyOpsFallbackAPIKey ContextKey = "ops_fallback_api_key"
)
// ForcePlatform 返回设置强制平台的中间件
diff --git a/backend/internal/service/admin_service_apikey_test.go b/backend/internal/service/admin_service_apikey_test.go
index f26fadb8..ccc8d221 100644
--- a/backend/internal/service/admin_service_apikey_test.go
+++ b/backend/internal/service/admin_service_apikey_test.go
@@ -146,6 +146,9 @@ func (s *apiKeyRepoStubForGroupUpdate) GetByKeyForAuth(context.Context, string)
panic("unexpected")
}
func (s *apiKeyRepoStubForGroupUpdate) Delete(context.Context, int64) error { panic("unexpected") }
+func (s *apiKeyRepoStubForGroupUpdate) DeleteWithAudit(context.Context, int64) error {
+ panic("unexpected")
+}
func (s *apiKeyRepoStubForGroupUpdate) ListByUserID(context.Context, int64, pagination.PaginationParams, APIKeyListFilters) ([]APIKey, *pagination.PaginationResult, error) {
panic("unexpected")
}
diff --git a/backend/internal/service/api_key_service.go b/backend/internal/service/api_key_service.go
index 48e0ab2f..dc008b8a 100644
--- a/backend/internal/service/api_key_service.go
+++ b/backend/internal/service/api_key_service.go
@@ -55,6 +55,8 @@ type APIKeyRepository interface {
GetByKeyForAuth(ctx context.Context, key string) (*APIKey, error)
Update(ctx context.Context, key *APIKey) error
Delete(ctx context.Context, id int64) error
+ // DeleteWithAudit 在同一事务内先写 deleted_api_key_audits 审计、再软删除该 key。
+ DeleteWithAudit(ctx context.Context, id int64) error
ListByUserID(ctx context.Context, userID int64, params pagination.PaginationParams, filters APIKeyListFilters) ([]APIKey, *pagination.PaginationResult, error)
VerifyOwnership(ctx context.Context, userID int64, apiKeyIDs []int64) ([]int64, error)
@@ -648,15 +650,16 @@ func (s *APIKeyService) Delete(ctx context.Context, id int64, userID int64) erro
return ErrInsufficientPerms
}
- // 清除Redis缓存(使用 userID 而非 apiKey.UserID)
+ // 事务内:写审计 + 软删除(tombstone)。
+ if err := s.apiKeyRepo.DeleteWithAudit(ctx, id); err != nil {
+ return fmt.Errorf("delete api key: %w", err)
+ }
+
+ // 删除成功后再清理缓存,避免"缓存已清但删除失败"的竞态。
if s.cache != nil {
_ = s.cache.DeleteCreateAttemptCount(ctx, userID)
}
s.InvalidateAuthCacheByKey(ctx, key)
-
- if err := s.apiKeyRepo.Delete(ctx, id); err != nil {
- return fmt.Errorf("delete api key: %w", err)
- }
s.lastUsedTouchL1.Delete(id)
return nil
diff --git a/backend/internal/service/api_key_service_cache_test.go b/backend/internal/service/api_key_service_cache_test.go
index eaac9a1c..a1dfbcb0 100644
--- a/backend/internal/service/api_key_service_cache_test.go
+++ b/backend/internal/service/api_key_service_cache_test.go
@@ -53,6 +53,10 @@ func (s *authRepoStub) Delete(ctx context.Context, id int64) error {
panic("unexpected Delete call")
}
+func (s *authRepoStub) DeleteWithAudit(ctx context.Context, id int64) error {
+ panic("unexpected DeleteWithAudit call")
+}
+
func (s *authRepoStub) ListByUserID(ctx context.Context, userID int64, params pagination.PaginationParams, filters APIKeyListFilters) ([]APIKey, *pagination.PaginationResult, error) {
panic("unexpected ListByUserID call")
}
diff --git a/backend/internal/service/api_key_service_delete_test.go b/backend/internal/service/api_key_service_delete_test.go
index 392d52b9..b8511c35 100644
--- a/backend/internal/service/api_key_service_delete_test.go
+++ b/backend/internal/service/api_key_service_delete_test.go
@@ -79,6 +79,12 @@ func (s *apiKeyRepoStub) Delete(ctx context.Context, id int64) error {
return s.deleteErr
}
+// DeleteWithAudit 与 Delete 一样记录被删除的 ID,供 service 测试断言。
+func (s *apiKeyRepoStub) DeleteWithAudit(ctx context.Context, id int64) error {
+ s.deletedIDs = append(s.deletedIDs, id)
+ return s.deleteErr
+}
+
// 以下是接口要求实现但本测试不关心的方法
func (s *apiKeyRepoStub) ListByUserID(ctx context.Context, userID int64, params pagination.PaginationParams, filters APIKeyListFilters) ([]APIKey, *pagination.PaginationResult, error) {
@@ -274,8 +280,8 @@ func TestApiKeyService_Delete_NotFound(t *testing.T) {
// 预期行为:
// - GetKeyAndOwnerID 返回正确的所有者 ID
// - 所有权验证通过
-// - 缓存被清除(在删除之前)
-// - Delete 被调用但返回错误
+// - DeleteWithAudit 被调用但返回错误
+// - 删除失败时缓存不被清除(缓存清理在删除成功后执行,消除竞态)
// - 返回包含 "delete api key" 的错误信息
func TestApiKeyService_Delete_DeleteFails(t *testing.T) {
repo := &apiKeyRepoStub{
@@ -288,7 +294,7 @@ func TestApiKeyService_Delete_DeleteFails(t *testing.T) {
err := svc.Delete(context.Background(), 3, 3) // API Key ID=3, 调用者 userID=3
require.Error(t, err)
require.ErrorContains(t, err, "delete api key")
- require.Equal(t, []int64{3}, repo.deletedIDs) // 验证删除操作被调用
- require.Equal(t, []int64{3}, cache.invalidated) // 验证缓存已被清除(即使删除失败)
- require.Equal(t, []string{svc.authCacheKey("k")}, cache.deleteAuthKeys)
+ require.Equal(t, []int64{3}, repo.deletedIDs) // 验证 DeleteWithAudit 被调用
+ require.Empty(t, cache.invalidated) // 验证删除失败时缓存未被清除(新顺序:先删后清)
+ require.Empty(t, cache.deleteAuthKeys) // 验证删除失败时 auth 缓存未被清除
}
diff --git a/backend/internal/service/api_key_service_quota_test.go b/backend/internal/service/api_key_service_quota_test.go
index cf05e16c..4d1d6f00 100644
--- a/backend/internal/service/api_key_service_quota_test.go
+++ b/backend/internal/service/api_key_service_quota_test.go
@@ -101,6 +101,9 @@ func (s *quotaBaseAPIKeyRepoStub) Update(context.Context, *APIKey) error {
func (s *quotaBaseAPIKeyRepoStub) Delete(context.Context, int64) error {
panic("unexpected Delete call")
}
+func (s *quotaBaseAPIKeyRepoStub) DeleteWithAudit(context.Context, int64) error {
+ panic("unexpected DeleteWithAudit call")
+}
func (s *quotaBaseAPIKeyRepoStub) ListByUserID(context.Context, int64, pagination.PaginationParams, APIKeyListFilters) ([]APIKey, *pagination.PaginationResult, error) {
panic("unexpected ListByUserID call")
}
diff --git a/backend/internal/service/ops_models.go b/backend/internal/service/ops_models.go
index ba735346..63c58cad 100644
--- a/backend/internal/service/ops_models.go
+++ b/backend/internal/service/ops_models.go
@@ -87,6 +87,15 @@ type OpsErrorLogDetail struct {
// vNext metric semantics
IsBusinessLimited bool `json:"is_business_limited"`
+
+ // Deleted key owner info (populated when INVALID_API_KEY and key was previously deleted)
+ AttemptedKeyPrefix string `json:"attempted_key_prefix,omitempty"`
+ DeletedKeyOwnerUserID *int64 `json:"deleted_key_owner_user_id,omitempty"`
+ DeletedKeyOwnerEmail string `json:"deleted_key_owner_email,omitempty"`
+ DeletedKeyName string `json:"deleted_key_name,omitempty"`
+
+ // Bound (non-deleted) key prefix, snapshotted at error time; mutually exclusive with AttemptedKeyPrefix.
+ APIKeyPrefix string `json:"api_key_prefix,omitempty"`
}
type OpsErrorLogFilter struct {
diff --git a/backend/internal/service/ops_port.go b/backend/internal/service/ops_port.go
index 30145ed3..0cba300d 100644
--- a/backend/internal/service/ops_port.go
+++ b/backend/internal/service/ops_port.go
@@ -10,6 +10,8 @@ type OpsRepository interface {
BatchInsertErrorLogs(ctx context.Context, inputs []*OpsInsertErrorLogInput) (int64, error)
ListErrorLogs(ctx context.Context, filter *OpsErrorLogFilter) (*OpsErrorLogList, error)
GetErrorLogByID(ctx context.Context, id int64) (*OpsErrorLogDetail, error)
+ // LookupDeletedKeyAudit 按明文 key 反查最近一条已删除 key 审计;未命中返回 (nil, nil)。
+ LookupDeletedKeyAudit(ctx context.Context, key string) (*DeletedKeyAuditResult, error)
ListRequestDetails(ctx context.Context, filter *OpsRequestDetailFilter) ([]*OpsRequestDetail, int64, error)
BatchInsertSystemLogs(ctx context.Context, inputs []*OpsInsertSystemLogInput) (int64, error)
ListSystemLogs(ctx context.Context, filter *OpsSystemLogFilter) (*OpsSystemLogList, error)
@@ -61,6 +63,12 @@ type OpsRepository interface {
GetLatestDailyBucketDate(ctx context.Context) (time.Time, bool, error)
}
+// DeletedKeyAuditResult 是按明文 key 反查 deleted_api_key_audits 的结果。
+type DeletedKeyAuditResult struct {
+ UserID int64
+ KeyName string
+}
+
type OpsInsertErrorLogInput struct {
RequestID string
ClientRequestID string
@@ -118,6 +126,15 @@ type OpsInsertErrorLogInput struct {
TimeToFirstTokenMs *int64
CreatedAt time.Time
+
+ // 已删除 key 归因(仅 INVALID_API_KEY 认证失败时可能非空)
+ AttemptedKeyPrefix string // 提交 key 的脱敏前缀(前 8 位)
+ DeletedKeyOwnerUserID *int64 // 反查命中的原所有者 user_id
+ DeletedKeyName string // 反查命中的 key 名称
+
+ // 有效(未删除)key 报错时快照的 key 脱敏前缀(前 8 位);与 AttemptedKeyPrefix 互斥。
+ // 落库快照而非读时 JOIN:key 之后被删(key 列被 tombstone 覆盖)仍保留当时前缀。
+ APIKeyPrefix string
}
type OpsInsertSystemMetricsInput struct {
diff --git a/backend/internal/service/ops_repo_mock_test.go b/backend/internal/service/ops_repo_mock_test.go
index 4138ea77..5e33bffb 100644
--- a/backend/internal/service/ops_repo_mock_test.go
+++ b/backend/internal/service/ops_repo_mock_test.go
@@ -13,6 +13,7 @@ type opsRepoMock struct {
ListSystemLogsFn func(ctx context.Context, filter *OpsSystemLogFilter) (*OpsSystemLogList, error)
DeleteSystemLogsFn func(ctx context.Context, filter *OpsSystemLogCleanupFilter) (int64, error)
InsertSystemLogCleanupAuditFn func(ctx context.Context, input *OpsSystemLogCleanupAudit) error
+ LookupDeletedKeyAuditFn func(ctx context.Context, key string) (*DeletedKeyAuditResult, error)
}
func (m *opsRepoMock) InsertErrorLog(ctx context.Context, input *OpsInsertErrorLogInput) (int64, error) {
@@ -189,4 +190,11 @@ func (m *opsRepoMock) GetLatestDailyBucketDate(ctx context.Context) (time.Time,
return time.Time{}, false, nil
}
+func (m *opsRepoMock) LookupDeletedKeyAudit(ctx context.Context, key string) (*DeletedKeyAuditResult, error) {
+ if m.LookupDeletedKeyAuditFn != nil {
+ return m.LookupDeletedKeyAuditFn(ctx, key)
+ }
+ return nil, nil
+}
+
var _ OpsRepository = (*opsRepoMock)(nil)
diff --git a/backend/internal/service/ops_service.go b/backend/internal/service/ops_service.go
index 2d7c5bd4..3d234b28 100644
--- a/backend/internal/service/ops_service.go
+++ b/backend/internal/service/ops_service.go
@@ -355,6 +355,14 @@ func (s *OpsService) GetErrorLogByID(ctx context.Context, id int64) (*OpsErrorLo
return detail, nil
}
+// LookupDeletedKeyAudit 按明文 key 反查已删除 key 的原所有者;未命中或未启用返回 (nil, nil)。
+func (s *OpsService) LookupDeletedKeyAudit(ctx context.Context, key string) (*DeletedKeyAuditResult, error) {
+ if s.opsRepo == nil {
+ return nil, nil
+ }
+ return s.opsRepo.LookupDeletedKeyAudit(ctx, key)
+}
+
func (s *OpsService) UpdateErrorResolution(ctx context.Context, errorID int64, resolved bool, resolvedByUserID *int64) error {
if err := s.RequireMonitoringEnabled(ctx); err != nil {
return err
diff --git a/backend/migrations/145_deleted_api_key_audit.sql b/backend/migrations/145_deleted_api_key_audit.sql
new file mode 100644
index 00000000..1364c094
--- /dev/null
+++ b/backend/migrations/145_deleted_api_key_audit.sql
@@ -0,0 +1,22 @@
+-- 已删除 API key 审计表:删除 key 时同步留存(明文 key、所有者、key 信息),
+-- 供认证失败(INVALID_API_KEY)反查"这个失效 key 曾属于谁"。
+-- 仅对本表上线后删除的 key 生效;此前已删的 key 原值已被 tombstone 覆盖,无法补录。
+SET LOCAL lock_timeout = '5s';
+SET LOCAL statement_timeout = '10min';
+
+CREATE TABLE IF NOT EXISTS deleted_api_key_audits (
+ id BIGSERIAL PRIMARY KEY,
+ key VARCHAR(128) NOT NULL, -- 原 key 明文(复用 api_keys 策略),非唯一
+ api_key_id BIGINT NOT NULL, -- 原 api_keys.id
+ user_id BIGINT NOT NULL, -- 原所有者(不加外键,与 ops 表设计哲学一致)
+ key_name VARCHAR(100) NOT NULL DEFAULT '', -- 原 key 名称,便于展示
+ deleted_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
+ created_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
+);
+CREATE INDEX IF NOT EXISTS deletedapikeyaudit_key ON deleted_api_key_audits (key);
+CREATE INDEX IF NOT EXISTS deletedapikeyaudit_user_id ON deleted_api_key_audits (user_id);
+
+ALTER TABLE ops_error_logs
+ ADD COLUMN IF NOT EXISTS attempted_key_prefix VARCHAR(32),
+ ADD COLUMN IF NOT EXISTS deleted_key_owner_user_id BIGINT,
+ ADD COLUMN IF NOT EXISTS deleted_key_name VARCHAR(100);
diff --git a/backend/migrations/147_ops_error_log_api_key_prefix.sql b/backend/migrations/147_ops_error_log_api_key_prefix.sql
new file mode 100644
index 00000000..bfc07489
--- /dev/null
+++ b/backend/migrations/147_ops_error_log_api_key_prefix.sql
@@ -0,0 +1,12 @@
+-- 有效(未删除)key 报错时,在 ops 落库层快照该 key 的脱敏前缀(前 8 位),
+-- 便于在 /admin/ops 错误详情识别是用户的哪一个 key 出的错。
+-- 与 attempted_key_prefix 互补且互斥:
+-- api_key_id 非空(有效 key 报错) => api_key_prefix
+-- api_key_id 为空(INVALID_API_KEY 无效) => attempted_key_prefix
+-- 落库快照(而非读时 JOIN api_keys):key 之后被删时 api_keys.key 会被 tombstone
+-- 覆盖,快照可保留报错当时的真实前缀。
+SET LOCAL lock_timeout = '5s';
+SET LOCAL statement_timeout = '10min';
+
+ALTER TABLE ops_error_logs
+ ADD COLUMN IF NOT EXISTS api_key_prefix VARCHAR(32);
diff --git a/frontend/src/api/admin/ops.ts b/frontend/src/api/admin/ops.ts
index 847fc8c9..557fd00f 100644
--- a/frontend/src/api/admin/ops.ts
+++ b/frontend/src/api/admin/ops.ts
@@ -941,6 +941,15 @@ export interface OpsErrorDetail extends OpsErrorLog {
time_to_first_token_ms?: number | null
is_business_limited: boolean
+
+ // Deleted key owner info (INVALID_API_KEY attribution)
+ attempted_key_prefix?: string | null
+ deleted_key_owner_user_id?: number | null
+ deleted_key_owner_email?: string | null
+ deleted_key_name?: string | null
+
+ // Bound (non-deleted) key prefix, snapshotted at error time
+ api_key_prefix?: string | null
}
export type OpsErrorLogsResponse = PaginatedResponse
diff --git a/frontend/src/i18n/locales/en.ts b/frontend/src/i18n/locales/en.ts
index b4e231e7..1cca8b05 100644
--- a/frontend/src/i18n/locales/en.ts
+++ b/frontend/src/i18n/locales/en.ts
@@ -4883,7 +4883,11 @@ export default {
suggestRequest: 'Client request error: ask customer to fix request parameters',
suggestAuth: 'Auth failed: verify API key/credentials',
suggestPlatform: 'Platform error: prioritize investigation and fix',
- suggestGeneric: 'See details for more context'
+ suggestGeneric: 'See details for more context',
+ apiKeyPrefix: 'Key Prefix',
+ attemptedKeyPrefix: 'Attempted Key Prefix',
+ deletedKeyOwner: 'Deleted Key Owner',
+ keyDeletedBadge: 'Key Deleted'
},
requestDetails: {
title: 'Request Details',
diff --git a/frontend/src/i18n/locales/zh.ts b/frontend/src/i18n/locales/zh.ts
index 70a7f6df..3edeb52b 100644
--- a/frontend/src/i18n/locales/zh.ts
+++ b/frontend/src/i18n/locales/zh.ts
@@ -5042,7 +5042,11 @@ export default {
suggestRequest: '⚠️ 客户端请求错误,建议:联系客户修正请求参数 / 手动标记已解决',
suggestAuth: '⚠️ 认证失败,建议:检查 API Key 是否有效 / 联系客户更新凭证',
suggestPlatform: '🚨 平台错误,建议立即排查修复',
- suggestGeneric: '查看详情了解更多信息'
+ suggestGeneric: '查看详情了解更多信息',
+ apiKeyPrefix: 'Key 前缀',
+ attemptedKeyPrefix: '尝试的 Key 前缀',
+ deletedKeyOwner: '已删除 Key 所有者',
+ keyDeletedBadge: 'Key 已删除'
},
requestDetails: {
title: '请求明细',
diff --git a/frontend/src/views/admin/ops/components/OpsErrorDetailModal.vue b/frontend/src/views/admin/ops/components/OpsErrorDetailModal.vue
index d29607e5..c346c547 100644
--- a/frontend/src/views/admin/ops/components/OpsErrorDetailModal.vue
+++ b/frontend/src/views/admin/ops/components/OpsErrorDetailModal.vue
@@ -106,6 +106,31 @@
{{ detail.message || '—' }}
+
+
+
{{ t('admin.ops.errorDetail.apiKeyPrefix') }}
+
+ {{ detail.api_key_prefix }}
+
+
+
+
+
{{ t('admin.ops.errorDetail.attemptedKeyPrefix') }}
+
+ {{ detail.attempted_key_prefix }}
+
+
+
+
+
{{ t('admin.ops.errorDetail.deletedKeyOwner') }}
+
+ {{ detail.deleted_key_owner_email }}
+ ({{ detail.deleted_key_name }})
+
+ {{ t('admin.ops.errorDetail.keyDeletedBadge') }}
+
+
+
From cfb195c7b2c91c6fa7ffb2bbcd8f816e0efcd3c9 Mon Sep 17 00:00:00 2001
From: DaydreamCoding
Date: Thu, 4 Jun 2026 19:06:24 +0800
Subject: [PATCH 77/79] =?UTF-8?q?feat(usage):=20=E8=AE=B0=E5=BD=95?=
=?UTF-8?q?=E5=B9=B6=E5=B1=95=E7=A4=BA=E5=A4=B1=E8=B4=A5=E8=AF=B7=E6=B1=82?=
=?UTF-8?q?(=E7=94=A8=E6=88=B7=E7=AB=AF+=E7=AE=A1=E7=90=86=E7=AB=AF)?=
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit
- 记录失败请求并在用户端/管理端展示;分类下拉改用统一 Select 组件
- 模型过滤改后端 ILIKE 模糊匹配;新增「Key 名称」列(含已删除标记)与按 Key 过滤;时间列移至末列
- 用户可见「已删除 key 失败请求」:OpsErrorLogFilter 加 MatchDeletedKeyOwner,用户侧归属
放宽为 (user_id OR deleted_key_owner_user_id),让 key 原所有者能看到删除 key 后继续请求
导致的认证失败记录(他人仍 NotFound,不泄露存在性)
- 迁移 148:ops_error_logs 用户+时间索引
Co-Authored-By: Claude Opus 4.8 (1M context)
---
backend/cmd/server/wire_gen.go | 104 ++++----
backend/internal/handler/admin/ops_handler.go | 16 ++
.../internal/handler/admin/setting_handler.go | 13 +-
backend/internal/handler/dto/settings.go | 5 +
backend/internal/handler/setting_handler.go | 2 +
backend/internal/handler/usage_handler.go | 131 ++++++++++-
.../handler/usage_handler_daily_test.go | 2 +-
.../usage_handler_request_type_test.go | 2 +-
.../repository/ops_error_where_test.go | 95 ++++++++
backend/internal/repository/ops_repo.go | 81 ++++++-
backend/internal/server/api_contract_test.go | 8 +-
backend/internal/server/routes/user.go | 2 +
backend/internal/service/domain_constants.go | 4 +
backend/internal/service/ops_models.go | 32 ++-
backend/internal/service/ops_service.go | 69 ++++++
.../service/ops_service_user_error_test.go | 222 ++++++++++++++++++
backend/internal/service/ops_user_error.go | 123 ++++++++++
.../internal/service/ops_user_error_test.go | 167 +++++++++++++
backend/internal/service/setting_service.go | 22 ++
.../service/setting_service_public_test.go | 13 +
...setting_service_user_error_persist_test.go | 31 +++
.../service/setting_user_error_view_test.go | 9 +
backend/internal/service/settings_view.go | 6 +
...dd_ops_error_logs_user_time_index_notx.sql | 6 +
frontend/src/api/admin/ops.ts | 2 +
frontend/src/api/admin/settings.ts | 5 +
frontend/src/api/usage.ts | 26 +-
frontend/src/components/common/Select.vue | 30 +++
.../components/user/UserErrorDetailModal.vue | 127 ++++++++++
.../user/UserErrorRequestsTable.vue | 173 ++++++++++++++
frontend/src/i18n/locales/en.ts | 29 ++-
frontend/src/i18n/locales/zh.ts | 29 ++-
frontend/src/stores/app.ts | 1 +
frontend/src/types/index.ts | 31 +++
frontend/src/views/admin/SettingsView.vue | 32 +++
frontend/src/views/admin/UsageView.vue | 93 +++++++-
frontend/src/views/user/UsageView.vue | 87 ++++++-
37 files changed, 1746 insertions(+), 84 deletions(-)
create mode 100644 backend/internal/repository/ops_error_where_test.go
create mode 100644 backend/internal/service/ops_service_user_error_test.go
create mode 100644 backend/internal/service/ops_user_error.go
create mode 100644 backend/internal/service/ops_user_error_test.go
create mode 100644 backend/internal/service/setting_service_user_error_persist_test.go
create mode 100644 backend/internal/service/setting_user_error_view_test.go
create mode 100644 backend/migrations/148_add_ops_error_logs_user_time_index_notx.sql
create mode 100644 frontend/src/components/user/UserErrorDetailModal.vue
create mode 100644 frontend/src/components/user/UserErrorRequestsTable.vue
diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go
index 10643215..814c07e6 100644
--- a/backend/cmd/server/wire_gen.go
+++ b/backend/cmd/server/wire_gen.go
@@ -91,36 +91,15 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
apiKeyHandler := handler.NewAPIKeyHandler(apiKeyService)
usageLogRepository := repository.NewUsageLogRepository(client, db)
usageService := service.NewUsageService(usageLogRepository, userRepository, client, apiKeyAuthCacheInvalidator)
- usageHandler := handler.NewUsageHandler(usageService, apiKeyService)
- redeemHandler := handler.NewRedeemHandler(redeemService)
- subscriptionHandler := handler.NewSubscriptionHandler(subscriptionService)
- announcementRepository := repository.NewAnnouncementRepository(client)
- announcementReadRepository := repository.NewAnnouncementReadRepository(client)
- announcementService := service.NewAnnouncementService(announcementRepository, announcementReadRepository, userRepository, userSubscriptionRepository)
- announcementHandler := handler.NewAnnouncementHandler(announcementService)
- channelMonitorRepository := repository.NewChannelMonitorRepository(client, db)
- channelMonitorService := service.ProvideChannelMonitorService(channelMonitorRepository, secretEncryptor)
- channelMonitorUserHandler := handler.NewChannelMonitorUserHandler(channelMonitorService, settingService)
- dashboardAggregationRepository := repository.NewDashboardAggregationRepository(db)
- dashboardStatsCache := repository.NewDashboardCache(redisClient, configConfig)
- dashboardService := service.NewDashboardService(usageLogRepository, dashboardAggregationRepository, dashboardStatsCache, configConfig)
- timingWheelService, err := service.ProvideTimingWheelService()
- if err != nil {
- return nil, err
- }
- dashboardAggregationService := service.ProvideDashboardAggregationService(dashboardAggregationRepository, timingWheelService, configConfig)
- dashboardHandler := admin.NewDashboardHandler(dashboardService, dashboardAggregationService)
+ opsRepository := repository.NewOpsRepository(db)
schedulerCache := repository.ProvideSchedulerCache(redisClient, configConfig)
accountRepository := repository.NewAccountRepository(client, db, schedulerCache)
- proxyExitInfoProber := repository.NewProxyExitInfoProber(configConfig)
- proxyLatencyCache := repository.NewProxyLatencyCache(redisClient)
- privacyClientFactory := providePrivacyClientFactory()
+ concurrencyCache := repository.ProvideConcurrencyCache(redisClient, configConfig)
+ concurrencyService := service.ProvideConcurrencyService(concurrencyCache, accountRepository, configConfig)
usageBillingRepository := repository.NewUsageBillingRepository(client, db)
gatewayCache := repository.NewGatewayCache(redisClient)
schedulerOutboxRepository := repository.NewSchedulerOutboxRepository(db)
schedulerSnapshotService := service.ProvideSchedulerSnapshotService(schedulerCache, schedulerOutboxRepository, accountRepository, groupRepository, configConfig)
- concurrencyCache := repository.ProvideConcurrencyCache(redisClient, configConfig)
- concurrencyService := service.ProvideConcurrencyService(concurrencyCache, accountRepository, configConfig)
pricingRemoteClient := repository.ProvidePricingRemoteClient(configConfig)
pricingService, err := service.ProvidePricingService(configConfig, pricingRemoteClient)
if err != nil {
@@ -134,44 +113,72 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
geminiTokenCache := repository.NewGeminiTokenCache(redisClient)
compositeTokenCacheInvalidator := service.NewCompositeTokenCacheInvalidator(geminiTokenCache)
rateLimitService := service.ProvideRateLimitService(accountRepository, usageLogRepository, configConfig, geminiQuotaService, tempUnschedCache, timeoutCounterCache, openAI403CounterCache, settingService, compositeTokenCacheInvalidator)
+ identityCache := repository.NewIdentityCache(redisClient)
+ identityService := service.NewIdentityService(identityCache)
httpUpstream := repository.NewHTTPUpstream(configConfig)
+ timingWheelService, err := service.ProvideTimingWheelService()
+ if err != nil {
+ return nil, err
+ }
deferredService := service.ProvideDeferredService(accountRepository, timingWheelService)
- openAIOAuthClient := repository.NewOpenAIOAuthClient()
- openAIOAuthService := service.ProvideOpenAIOAuthService(proxyRepository, openAIOAuthClient, privacyClientFactory)
+ claudeOAuthClient := repository.NewClaudeOAuthClient()
+ oAuthService := service.NewOAuthService(proxyRepository, claudeOAuthClient)
oAuthRefreshAPI := service.ProvideOAuthRefreshAPI(accountRepository, geminiTokenCache)
- openAITokenProvider := service.ProvideOpenAITokenProvider(accountRepository, geminiTokenCache, openAIOAuthService, oAuthRefreshAPI)
+ claudeTokenProvider := service.ProvideClaudeTokenProvider(accountRepository, geminiTokenCache, oAuthService, oAuthRefreshAPI)
+ sessionLimitCache := repository.ProvideSessionLimitCache(redisClient, configConfig)
+ rpmCache := repository.NewRPMCache(redisClient)
+ digestSessionStore := service.NewDigestSessionStore()
+ tlsFingerprintProfileRepository := repository.NewTLSFingerprintProfileRepository(client)
+ tlsFingerprintProfileCache := repository.NewTLSFingerprintProfileCache(redisClient)
+ tlsFingerprintProfileService := service.NewTLSFingerprintProfileService(tlsFingerprintProfileRepository, tlsFingerprintProfileCache)
channelRepository := repository.NewChannelRepository(db)
channelService := service.NewChannelService(channelRepository, groupRepository, apiKeyAuthCacheInvalidator, pricingService)
modelPricingResolver := service.NewModelPricingResolver(channelService, billingService)
notificationEmailService := service.NewNotificationEmailService(settingRepository, emailService)
balanceNotifyService := service.ProvideBalanceNotifyService(emailService, settingRepository, accountRepository, notificationEmailService)
+ gatewayService := service.NewGatewayService(accountRepository, groupRepository, usageLogRepository, usageBillingRepository, userRepository, userSubscriptionRepository, userGroupRateRepository, gatewayCache, configConfig, schedulerSnapshotService, concurrencyService, billingService, rateLimitService, billingCacheService, identityService, httpUpstream, deferredService, claudeTokenProvider, sessionLimitCache, rpmCache, digestSessionStore, settingService, tlsFingerprintProfileService, channelService, modelPricingResolver, balanceNotifyService, serviceUserPlatformQuotaRepository)
+ openAIOAuthClient := repository.NewOpenAIOAuthClient()
+ privacyClientFactory := providePrivacyClientFactory()
+ openAIOAuthService := service.ProvideOpenAIOAuthService(proxyRepository, openAIOAuthClient, privacyClientFactory)
+ openAITokenProvider := service.ProvideOpenAITokenProvider(accountRepository, geminiTokenCache, openAIOAuthService, oAuthRefreshAPI)
openAIGatewayService := service.NewOpenAIGatewayService(accountRepository, usageLogRepository, usageBillingRepository, userRepository, userSubscriptionRepository, userGroupRateRepository, gatewayCache, configConfig, schedulerSnapshotService, concurrencyService, billingService, rateLimitService, billingCacheService, httpUpstream, deferredService, openAITokenProvider, modelPricingResolver, channelService, balanceNotifyService, settingService, serviceUserPlatformQuotaRepository)
- adminService := service.NewAdminService(userRepository, groupRepository, accountRepository, proxyRepository, apiKeyRepository, redeemCodeRepository, userGroupRateRepository, userRPMCache, billingCacheService, proxyExitInfoProber, proxyLatencyCache, apiKeyAuthCacheInvalidator, client, settingService, subscriptionService, userSubscriptionRepository, privacyClientFactory, openAIGatewayService)
- adminUserHandler := admin.NewUserHandler(adminService, concurrencyService, serviceUserPlatformQuotaRepository, billingCache)
- sessionLimitCache := repository.ProvideSessionLimitCache(redisClient, configConfig)
- rpmCache := repository.NewRPMCache(redisClient)
- groupCapacityService := service.NewGroupCapacityService(accountRepository, groupRepository, concurrencyService, sessionLimitCache, rpmCache)
- groupHandler := admin.NewGroupHandler(adminService, dashboardService, groupCapacityService)
- claudeOAuthClient := repository.NewClaudeOAuthClient()
- oAuthService := service.NewOAuthService(proxyRepository, claudeOAuthClient)
geminiOAuthClient := repository.NewGeminiOAuthClient(configConfig)
geminiCliCodeAssistClient := repository.NewGeminiCliCodeAssistClient()
driveClient := repository.NewGeminiDriveClient()
geminiOAuthService := service.NewGeminiOAuthService(proxyRepository, geminiOAuthClient, geminiCliCodeAssistClient, driveClient, configConfig)
- antigravityOAuthService := service.NewAntigravityOAuthService(proxyRepository)
- claudeUsageFetcher := repository.NewClaudeUsageFetcher(httpUpstream)
- antigravityQuotaFetcher := service.NewAntigravityQuotaFetcher(proxyRepository)
- usageCache := service.NewUsageCache()
- identityCache := repository.NewIdentityCache(redisClient)
- tlsFingerprintProfileRepository := repository.NewTLSFingerprintProfileRepository(client)
- tlsFingerprintProfileCache := repository.NewTLSFingerprintProfileCache(redisClient)
- tlsFingerprintProfileService := service.NewTLSFingerprintProfileService(tlsFingerprintProfileRepository, tlsFingerprintProfileCache)
- accountUsageService := service.NewAccountUsageService(accountRepository, usageLogRepository, claudeUsageFetcher, geminiQuotaService, antigravityQuotaFetcher, usageCache, identityCache, tlsFingerprintProfileService)
geminiTokenProvider := service.ProvideGeminiTokenProvider(accountRepository, geminiTokenCache, geminiOAuthService, oAuthRefreshAPI)
- claudeTokenProvider := service.ProvideClaudeTokenProvider(accountRepository, geminiTokenCache, oAuthService, oAuthRefreshAPI)
+ antigravityOAuthService := service.NewAntigravityOAuthService(proxyRepository)
antigravityTokenProvider := service.ProvideAntigravityTokenProvider(accountRepository, geminiTokenCache, antigravityOAuthService, oAuthRefreshAPI, tempUnschedCache)
internal500CounterCache := repository.NewInternal500CounterCache(redisClient)
antigravityGatewayService := service.NewAntigravityGatewayService(accountRepository, gatewayCache, schedulerSnapshotService, antigravityTokenProvider, rateLimitService, httpUpstream, settingService, internal500CounterCache)
+ geminiMessagesCompatService := service.NewGeminiMessagesCompatService(accountRepository, groupRepository, gatewayCache, schedulerSnapshotService, geminiTokenProvider, rateLimitService, httpUpstream, antigravityGatewayService, configConfig)
+ opsSystemLogSink := service.ProvideOpsSystemLogSink(opsRepository)
+ opsService := service.ProvideOpsService(opsRepository, settingRepository, configConfig, accountRepository, userRepository, concurrencyService, gatewayService, openAIGatewayService, geminiMessagesCompatService, antigravityGatewayService, opsSystemLogSink, settingService)
+ usageHandler := handler.NewUsageHandler(usageService, apiKeyService, opsService, settingService)
+ redeemHandler := handler.NewRedeemHandler(redeemService)
+ subscriptionHandler := handler.NewSubscriptionHandler(subscriptionService)
+ announcementRepository := repository.NewAnnouncementRepository(client)
+ announcementReadRepository := repository.NewAnnouncementReadRepository(client)
+ announcementService := service.NewAnnouncementService(announcementRepository, announcementReadRepository, userRepository, userSubscriptionRepository)
+ announcementHandler := handler.NewAnnouncementHandler(announcementService)
+ channelMonitorRepository := repository.NewChannelMonitorRepository(client, db)
+ channelMonitorService := service.ProvideChannelMonitorService(channelMonitorRepository, secretEncryptor)
+ channelMonitorUserHandler := handler.NewChannelMonitorUserHandler(channelMonitorService, settingService)
+ dashboardAggregationRepository := repository.NewDashboardAggregationRepository(db)
+ dashboardStatsCache := repository.NewDashboardCache(redisClient, configConfig)
+ dashboardService := service.NewDashboardService(usageLogRepository, dashboardAggregationRepository, dashboardStatsCache, configConfig)
+ dashboardAggregationService := service.ProvideDashboardAggregationService(dashboardAggregationRepository, timingWheelService, configConfig)
+ dashboardHandler := admin.NewDashboardHandler(dashboardService, dashboardAggregationService)
+ proxyExitInfoProber := repository.NewProxyExitInfoProber(configConfig)
+ proxyLatencyCache := repository.NewProxyLatencyCache(redisClient)
+ adminService := service.NewAdminService(userRepository, groupRepository, accountRepository, proxyRepository, apiKeyRepository, redeemCodeRepository, userGroupRateRepository, userRPMCache, billingCacheService, proxyExitInfoProber, proxyLatencyCache, apiKeyAuthCacheInvalidator, client, settingService, subscriptionService, userSubscriptionRepository, privacyClientFactory, openAIGatewayService)
+ adminUserHandler := admin.NewUserHandler(adminService, concurrencyService, serviceUserPlatformQuotaRepository, billingCache)
+ groupCapacityService := service.NewGroupCapacityService(accountRepository, groupRepository, concurrencyService, sessionLimitCache, rpmCache)
+ groupHandler := admin.NewGroupHandler(adminService, dashboardService, groupCapacityService)
+ claudeUsageFetcher := repository.NewClaudeUsageFetcher(httpUpstream)
+ antigravityQuotaFetcher := service.NewAntigravityQuotaFetcher(proxyRepository)
+ usageCache := service.NewUsageCache()
+ accountUsageService := service.NewAccountUsageService(accountRepository, usageLogRepository, claudeUsageFetcher, geminiQuotaService, antigravityQuotaFetcher, usageCache, identityCache, tlsFingerprintProfileService)
accountTestService := service.NewAccountTestService(accountRepository, geminiTokenProvider, claudeTokenProvider, antigravityGatewayService, httpUpstream, configConfig, tlsFingerprintProfileService)
crsSyncService := service.NewCRSSyncService(accountRepository, proxyRepository, oAuthService, openAIOAuthService, geminiOAuthService, configConfig)
accountHandler := admin.NewAccountHandler(adminService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, rateLimitService, accountUsageService, accountTestService, concurrencyService, crsSyncService, sessionLimitCache, rpmCache, compositeTokenCacheInvalidator)
@@ -189,13 +196,6 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
proxyHandler := admin.NewProxyHandler(adminService)
adminRedeemHandler := admin.NewRedeemHandler(adminService, redeemService)
promoHandler := admin.NewPromoHandler(promoService)
- opsRepository := repository.NewOpsRepository(db)
- identityService := service.NewIdentityService(identityCache)
- digestSessionStore := service.NewDigestSessionStore()
- gatewayService := service.NewGatewayService(accountRepository, groupRepository, usageLogRepository, usageBillingRepository, userRepository, userSubscriptionRepository, userGroupRateRepository, gatewayCache, configConfig, schedulerSnapshotService, concurrencyService, billingService, rateLimitService, billingCacheService, identityService, httpUpstream, deferredService, claudeTokenProvider, sessionLimitCache, rpmCache, digestSessionStore, settingService, tlsFingerprintProfileService, channelService, modelPricingResolver, balanceNotifyService, serviceUserPlatformQuotaRepository)
- geminiMessagesCompatService := service.NewGeminiMessagesCompatService(accountRepository, groupRepository, gatewayCache, schedulerSnapshotService, geminiTokenProvider, rateLimitService, httpUpstream, antigravityGatewayService, configConfig)
- opsSystemLogSink := service.ProvideOpsSystemLogSink(opsRepository)
- opsService := service.ProvideOpsService(opsRepository, settingRepository, configConfig, accountRepository, userRepository, concurrencyService, gatewayService, openAIGatewayService, geminiMessagesCompatService, antigravityGatewayService, opsSystemLogSink, settingService)
encryptionKey, err := payment.ProvideEncryptionKey(configConfig)
if err != nil {
return nil, err
diff --git a/backend/internal/handler/admin/ops_handler.go b/backend/internal/handler/admin/ops_handler.go
index 418c302f..0ae93e65 100644
--- a/backend/internal/handler/admin/ops_handler.go
+++ b/backend/internal/handler/admin/ops_handler.go
@@ -137,6 +137,22 @@ func (h *OpsHandler) GetErrorLogs(c *gin.Context) {
filter.AccountID = &id
}
+ if v := strings.TrimSpace(c.Query("user_id")); v != "" {
+ id, err := strconv.ParseInt(v, 10, 64)
+ if err != nil || id <= 0 {
+ response.BadRequest(c, "Invalid user_id")
+ return
+ }
+ filter.UserID = &id
+ }
+ if v := strings.TrimSpace(c.Query("api_key_id")); v != "" {
+ id, err := strconv.ParseInt(v, 10, 64)
+ if err != nil || id <= 0 {
+ response.BadRequest(c, "Invalid api_key_id")
+ return
+ }
+ filter.APIKeyID = &id
+ }
if v := strings.TrimSpace(c.Query("resolved")); v != "" {
switch strings.ToLower(v) {
case "1", "true", "yes":
diff --git a/backend/internal/handler/admin/setting_handler.go b/backend/internal/handler/admin/setting_handler.go
index c229d340..0beb15d3 100644
--- a/backend/internal/handler/admin/setting_handler.go
+++ b/backend/internal/handler/admin/setting_handler.go
@@ -297,6 +297,8 @@ func (h *SettingHandler) GetSettings(c *gin.Context) {
AvailableChannelsEnabled: settings.AvailableChannelsEnabled,
AffiliateEnabled: settings.AffiliateEnabled,
+
+ AllowUserViewErrorRequests: settings.AllowUserViewErrorRequests,
}
// OpenAI fast policy (stored under a dedicated setting key)
@@ -658,6 +660,8 @@ type UpdateSettingsRequest struct {
AuthSourceGitHubPlatformQuotas map[string]*service.DefaultPlatformQuotaSetting `json:"auth_source_default_github_platform_quotas"`
AuthSourceGooglePlatformQuotas map[string]*service.DefaultPlatformQuotaSetting `json:"auth_source_default_google_platform_quotas"`
AuthSourceDingTalkPlatformQuotas map[string]*service.DefaultPlatformQuotaSetting `json:"auth_source_default_dingtalk_platform_quotas"`
+
+ AllowUserViewErrorRequests *bool `json:"allow_user_view_error_requests"`
}
// UpdateSettings 更新系统设置
@@ -1591,6 +1595,12 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
MaxClaudeCodeVersion: req.MaxClaudeCodeVersion,
AllowUngroupedKeyScheduling: req.AllowUngroupedKeyScheduling,
BackendModeEnabled: req.BackendModeEnabled,
+ AllowUserViewErrorRequests: func() bool {
+ if req.AllowUserViewErrorRequests != nil {
+ return *req.AllowUserViewErrorRequests
+ }
+ return previousSettings.AllowUserViewErrorRequests
+ }(),
OpsMonitoringEnabled: func() bool {
if req.OpsMonitoringEnabled != nil {
return *req.OpsMonitoringEnabled
@@ -2080,7 +2090,8 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
AffiliateEnabled: updatedSettings.AffiliateEnabled,
- RiskControlEnabled: updatedSettings.RiskControlEnabled,
+ RiskControlEnabled: updatedSettings.RiskControlEnabled,
+ AllowUserViewErrorRequests: updatedSettings.AllowUserViewErrorRequests,
}
if fastPolicy, err := h.settingService.GetOpenAIFastPolicySettings(c.Request.Context()); err != nil {
slog.Error("openai_fast_policy_settings_get_failed", "error", err)
diff --git a/backend/internal/handler/dto/settings.go b/backend/internal/handler/dto/settings.go
index 17772a2e..da89ac23 100644
--- a/backend/internal/handler/dto/settings.go
+++ b/backend/internal/handler/dto/settings.go
@@ -252,6 +252,9 @@ type SystemSettings struct {
// 系统全局默认平台配额(key = platform,nil/缺省 = 不限制)
DefaultPlatformQuotas map[string]*service.DefaultPlatformQuotaSetting `json:"default_platform_quotas,omitempty"`
+
+ // 允许终端用户在用量页查看自己的失败请求
+ AllowUserViewErrorRequests bool `json:"allow_user_view_error_requests"`
}
type DefaultSubscriptionSetting struct {
@@ -316,6 +319,8 @@ type PublicSettings struct {
AffiliateEnabled bool `json:"affiliate_enabled"`
RiskControlEnabled bool `json:"risk_control_enabled"`
+
+ AllowUserViewErrorRequests bool `json:"allow_user_view_error_requests"`
}
type LoginAgreementDocument struct {
diff --git a/backend/internal/handler/setting_handler.go b/backend/internal/handler/setting_handler.go
index 7413b840..7c79a19e 100644
--- a/backend/internal/handler/setting_handler.go
+++ b/backend/internal/handler/setting_handler.go
@@ -98,6 +98,8 @@ func (h *SettingHandler) GetPublicSettings(c *gin.Context) {
AffiliateEnabled: settings.AffiliateEnabled,
RiskControlEnabled: settings.RiskControlEnabled,
+
+ AllowUserViewErrorRequests: settings.AllowUserViewErrorRequests,
})
}
diff --git a/backend/internal/handler/usage_handler.go b/backend/internal/handler/usage_handler.go
index daa5695d..23bb62dd 100644
--- a/backend/internal/handler/usage_handler.go
+++ b/backend/internal/handler/usage_handler.go
@@ -1,6 +1,7 @@
package handler
import (
+ "net/http"
"strconv"
"strings"
"time"
@@ -18,15 +19,24 @@ import (
// UsageHandler handles usage-related requests
type UsageHandler struct {
- usageService *service.UsageService
- apiKeyService *service.APIKeyService
+ usageService *service.UsageService
+ apiKeyService *service.APIKeyService
+ opsService *service.OpsService
+ settingService *service.SettingService
}
// NewUsageHandler creates a new UsageHandler
-func NewUsageHandler(usageService *service.UsageService, apiKeyService *service.APIKeyService) *UsageHandler {
+func NewUsageHandler(
+ usageService *service.UsageService,
+ apiKeyService *service.APIKeyService,
+ opsService *service.OpsService,
+ settingService *service.SettingService,
+) *UsageHandler {
return &UsageHandler{
- usageService: usageService,
- apiKeyService: apiKeyService,
+ usageService: usageService,
+ apiKeyService: apiKeyService,
+ opsService: opsService,
+ settingService: settingService,
}
}
@@ -149,6 +159,117 @@ func (h *UsageHandler) List(c *gin.Context) {
response.Paginated(c, out, result.Total, page, pageSize)
}
+// ListErrors handles listing the current user's failed requests (redacted).
+// GET /api/v1/usage/errors
+func (h *UsageHandler) ListErrors(c *gin.Context) {
+ subject, ok := middleware2.GetAuthSubjectFromContext(c)
+ if !ok {
+ response.Unauthorized(c, "User not authenticated")
+ return
+ }
+
+ // Visibility switch (fail-closed). Defense-in-depth: frontend also hides the tab.
+ if h.settingService == nil || !h.settingService.IsUserErrorViewAllowed(c.Request.Context()) {
+ response.Forbidden(c, "Error requests view is disabled")
+ return
+ }
+ if h.opsService == nil {
+ response.Error(c, http.StatusServiceUnavailable, "Ops service not available")
+ return
+ }
+
+ page, pageSize := response.ParsePagination(c)
+ if pageSize > 100 {
+ pageSize = 100
+ }
+
+ filter := &service.OpsErrorLogFilter{Page: page, PageSize: pageSize}
+
+ // Date range (half-open [start, end)), reuse usage-list semantics.
+ userTZ := c.Query("timezone")
+ if startDateStr := c.Query("start_date"); startDateStr != "" {
+ t, err := timezone.ParseInUserLocation("2006-01-02", startDateStr, userTZ)
+ if err != nil {
+ response.BadRequest(c, "Invalid start_date format, use YYYY-MM-DD")
+ return
+ }
+ filter.StartTime = &t
+ }
+ if endDateStr := c.Query("end_date"); endDateStr != "" {
+ t, err := timezone.ParseInUserLocation("2006-01-02", endDateStr, userTZ)
+ if err != nil {
+ response.BadRequest(c, "Invalid end_date format, use YYYY-MM-DD")
+ return
+ }
+ t = t.AddDate(0, 0, 1)
+ filter.EndTime = &t
+ }
+
+ filter.Model = strings.TrimSpace(c.Query("model"))
+
+ if k := strings.TrimSpace(c.Query("api_key_id")); k != "" {
+ n, err := strconv.ParseInt(k, 10, 64)
+ if err != nil || n < 0 {
+ response.BadRequest(c, "Invalid api_key_id")
+ return
+ }
+ if n > 0 {
+ filter.APIKeyID = &n
+ }
+ }
+
+ if sc := strings.TrimSpace(c.Query("status_code")); sc != "" {
+ n, err := strconv.Atoi(sc)
+ if err != nil || n < 0 {
+ response.BadRequest(c, "Invalid status_code")
+ return
+ }
+ filter.StatusCodes = []int{n}
+ }
+
+ if cat := strings.TrimSpace(c.Query("category")); cat != "" {
+ phases, types := service.CategoryToFilter(cat)
+ filter.ErrorPhasesAny = phases
+ filter.ErrorTypesAny = types
+ }
+
+ result, err := h.opsService.ListUserErrorRequests(c.Request.Context(), subject.UserID, filter)
+ if err != nil {
+ response.ErrorFrom(c, err)
+ return
+ }
+ response.Paginated(c, result.Items, int64(result.Total), result.Page, result.PageSize)
+}
+
+// GetErrorDetail handles fetching one of the current user's failed-request details (redacted).
+// GET /api/v1/usage/errors/:id
+func (h *UsageHandler) GetErrorDetail(c *gin.Context) {
+ subject, ok := middleware2.GetAuthSubjectFromContext(c)
+ if !ok {
+ response.Unauthorized(c, "User not authenticated")
+ return
+ }
+ if h.settingService == nil || !h.settingService.IsUserErrorViewAllowed(c.Request.Context()) {
+ response.Forbidden(c, "Error requests view is disabled")
+ return
+ }
+ if h.opsService == nil {
+ response.Error(c, http.StatusServiceUnavailable, "Ops service not available")
+ return
+ }
+ id, err := strconv.ParseInt(strings.TrimSpace(c.Param("id")), 10, 64)
+ if err != nil || id <= 0 {
+ response.BadRequest(c, "Invalid id")
+ return
+ }
+ detail, err := h.opsService.GetUserErrorRequestDetail(c.Request.Context(), subject.UserID, id)
+ if err != nil {
+ response.ErrorFrom(c, err)
+ return
+ }
+ response.Success(c, detail)
+}
+
// GetByID handles getting a single usage record
// GET /api/v1/usage/:id
func (h *UsageHandler) GetByID(c *gin.Context) {
diff --git a/backend/internal/handler/usage_handler_daily_test.go b/backend/internal/handler/usage_handler_daily_test.go
index 36311fac..2a9186cf 100644
--- a/backend/internal/handler/usage_handler_daily_test.go
+++ b/backend/internal/handler/usage_handler_daily_test.go
@@ -64,7 +64,7 @@ func newDailyUsageTestRouter(usageRepo *dailyUsageRepoStub, apiKeyRepo *dailyUsa
gin.SetMode(gin.TestMode)
usageSvc := service.NewUsageService(usageRepo, nil, nil, nil)
apiKeySvc := service.NewAPIKeyService(apiKeyRepo, nil, nil, nil, nil, nil, nil)
- handler := NewUsageHandler(usageSvc, apiKeySvc)
+ handler := NewUsageHandler(usageSvc, apiKeySvc, nil, nil)
router := gin.New()
router.Use(func(c *gin.Context) {
c.Set(string(middleware2.ContextKeyUser), middleware2.AuthSubject{UserID: userID})
diff --git a/backend/internal/handler/usage_handler_request_type_test.go b/backend/internal/handler/usage_handler_request_type_test.go
index b49ed59b..ed08c5a8 100644
--- a/backend/internal/handler/usage_handler_request_type_test.go
+++ b/backend/internal/handler/usage_handler_request_type_test.go
@@ -34,7 +34,7 @@ func (s *userUsageRepoCapture) ListWithFilters(ctx context.Context, params pagin
func newUserUsageRequestTypeTestRouter(repo *userUsageRepoCapture) *gin.Engine {
gin.SetMode(gin.TestMode)
usageSvc := service.NewUsageService(repo, nil, nil, nil)
- handler := NewUsageHandler(usageSvc, nil)
+ handler := NewUsageHandler(usageSvc, nil, nil, nil)
router := gin.New()
router.Use(func(c *gin.Context) {
c.Set(string(middleware2.ContextKeyUser), middleware2.AuthSubject{UserID: 42})
diff --git a/backend/internal/repository/ops_error_where_test.go b/backend/internal/repository/ops_error_where_test.go
new file mode 100644
index 00000000..9bebb158
--- /dev/null
+++ b/backend/internal/repository/ops_error_where_test.go
@@ -0,0 +1,95 @@
+package repository
+
+import (
+ "strings"
+ "testing"
+
+ "github.com/Wei-Shaw/sub2api/internal/service"
+)
+
+func TestBuildOpsErrorLogsWhere_UserScopedFilters(t *testing.T) {
+ uid := int64(42)
+ kid := int64(7)
+ filter := &service.OpsErrorLogFilter{
+ UserID: &uid,
+ APIKeyID: &kid,
+ Model: "claude-sonnet-4-5",
+ ExcludeCountTokens: true,
+ ErrorPhasesAny: []string{"auth"},
+ ErrorTypesAny: []string{"rate_limit_error"},
+ View: "all",
+ }
+ where, args := buildOpsErrorLogsWhere(filter)
+
+ for _, want := range []string{
+ "e.user_id = $",
+ "e.api_key_id = $",
+ "COALESCE(e.requested_model, e.model, '') = $",
+ "COALESCE(e.is_count_tokens, false) = false",
+ "e.error_phase = ANY($",
+ "e.error_type = ANY($",
+ } {
+ if !strings.Contains(where, want) {
+ t.Fatalf("where missing %q\nfull: %s", want, where)
+ }
+ }
+ if len(args) != 5 {
+ t.Fatalf("expected 5 args, got %d", len(args))
+ }
+}
+
+func TestBuildOpsErrorLogsWhere_ModelFuzzy(t *testing.T) {
+ // 默认(ModelFuzzy=false)保持精确匹配
+ exact := &service.OpsErrorLogFilter{Model: "claude"}
+ whereExact, _ := buildOpsErrorLogsWhere(exact)
+ if !strings.Contains(whereExact, "COALESCE(e.requested_model, e.model, '') = $") {
+ t.Fatalf("default should be exact match, got: %s", whereExact)
+ }
+
+ // ModelFuzzy=true → ILIKE
+ fuzzy := &service.OpsErrorLogFilter{Model: "claude", ModelFuzzy: true}
+ whereFuzzy, args := buildOpsErrorLogsWhere(fuzzy)
+ if !strings.Contains(whereFuzzy, "COALESCE(e.requested_model, e.model, '') ILIKE $") {
+ t.Fatalf("ModelFuzzy should use ILIKE, got: %s", whereFuzzy)
+ }
+ if len(args) != 1 || args[0] != "%claude%" {
+ t.Fatalf("expected arg \"%%claude%%\", got %v", args)
+ }
+
+ // 通配符转义:输入含 % 应被转义为字面量
+ esc := &service.OpsErrorLogFilter{Model: "50%off", ModelFuzzy: true}
+ _, escArgs := buildOpsErrorLogsWhere(esc)
+ if len(escArgs) != 1 || escArgs[0] != `%50\%off%` {
+ t.Fatalf("expected escaped arg, got %v", escArgs)
+ }
+
+ esc2 := &service.OpsErrorLogFilter{Model: "gpt_4o", ModelFuzzy: true}
+ _, escArgs2 := buildOpsErrorLogsWhere(esc2)
+ if len(escArgs2) != 1 || escArgs2[0] != `%gpt\_4o%` {
+ t.Fatalf("underscore should be escaped, got %v", escArgs2)
+ }
+}
+
+func TestBuildOpsErrorLogsWhere_MatchDeletedKeyOwner(t *testing.T) {
+ uid := int64(42)
+
+ // 开关开启 → 归属放宽为 OR(user_id 或 deleted_key_owner_user_id),且共用同一占位符
+ on := &service.OpsErrorLogFilter{UserID: &uid, MatchDeletedKeyOwner: true}
+ whereOn, argsOn := buildOpsErrorLogsWhere(on)
+ if !strings.Contains(whereOn, "(e.user_id = $1 OR e.deleted_key_owner_user_id = $1)") {
+ t.Fatalf("MatchDeletedKeyOwner=true should widen to OR, got: %s", whereOn)
+ }
+ if len(argsOn) != 1 || argsOn[0] != uid {
+ t.Fatalf("expected single reused arg %d, got %v", uid, argsOn)
+ }
+
+ // 开关关闭(默认)→ 仅精确 user_id,绝不出现 deleted_key_owner_user_id(admin 回归)
+ off := &service.OpsErrorLogFilter{UserID: &uid}
+ whereOff, _ := buildOpsErrorLogsWhere(off)
+ if !strings.Contains(whereOff, "e.user_id = $1") {
+ t.Fatalf("default should match user_id exactly, got: %s", whereOff)
+ }
+ if strings.Contains(whereOff, "deleted_key_owner_user_id") {
+ t.Fatalf("default must NOT include deleted_key_owner_user_id, got: %s", whereOff)
+ }
+}
diff --git a/backend/internal/repository/ops_repo.go b/backend/internal/repository/ops_repo.go
index a7773713..f300a171 100644
--- a/backend/internal/repository/ops_repo.go
+++ b/backend/internal/repository/ops_repo.go
@@ -240,12 +240,16 @@ SELECT
COALESCE(e.upstream_endpoint, ''),
COALESCE(e.requested_model, ''),
COALESCE(e.upstream_model, ''),
- e.request_type
+ e.request_type,
+ COALESCE(ak.name, ''),
+ ak.deleted_at,
+ COALESCE(e.deleted_key_name, '')
FROM ops_error_logs e
LEFT JOIN accounts a ON e.account_id = a.id
LEFT JOIN groups g ON e.group_id = g.id
LEFT JOIN users u ON e.user_id = u.id
LEFT JOIN users u2 ON e.resolved_by_user_id = u2.id
+LEFT JOIN api_keys ak ON ak.id = e.api_key_id
` + where + `
ORDER BY e.created_at DESC
LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
@@ -272,6 +276,9 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
var resolvedBy sql.NullInt64
var resolvedByName string
var requestType sql.NullInt64
+ var apiKeyName string
+ var apiKeyDeletedAt sql.NullTime
+ var deletedKeyName string
if err := rows.Scan(
&item.ID,
&item.CreatedAt,
@@ -305,6 +312,9 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
&item.RequestedModel,
&item.UpstreamModel,
&requestType,
+ &apiKeyName,
+ &apiKeyDeletedAt,
+ &deletedKeyName,
); err != nil {
return nil, err
}
@@ -345,6 +355,15 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
v := int16(requestType.Int64)
item.RequestType = &v
}
+ // Key 名称:优先关联到的 ak.name(已软删的 key name 仍保留);
+ // 关联不到(api_key_id 为空 / 历史硬删)时回退错误记录里快照的 deleted_key_name。
+ if apiKeyName != "" {
+ item.APIKeyName = apiKeyName
+ } else {
+ item.APIKeyName = deletedKeyName
+ }
+ // 已删除:ak.deleted_at 非空(软删),或仅命中 deleted_key_name 兜底。
+ item.APIKeyDeleted = apiKeyDeletedAt.Valid || (apiKeyName == "" && deletedKeyName != "")
out = append(out, &item)
}
if err := rows.Err(); err != nil {
@@ -416,12 +435,15 @@ SELECT
e.deleted_key_owner_user_id,
COALESCE(du.email, ''),
COALESCE(e.deleted_key_name, ''),
- COALESCE(e.api_key_prefix, '')
+ COALESCE(e.api_key_prefix, ''),
+ COALESCE(ak.name, ''),
+ ak.deleted_at
FROM ops_error_logs e
LEFT JOIN users u ON e.user_id = u.id
LEFT JOIN accounts a ON e.account_id = a.id
LEFT JOIN groups g ON e.group_id = g.id
LEFT JOIN users du ON e.deleted_key_owner_user_id = du.id
+LEFT JOIN api_keys ak ON ak.id = e.api_key_id
WHERE e.id = $1
LIMIT 1`
@@ -442,6 +464,8 @@ LIMIT 1`
var ttft sql.NullInt64
var requestType sql.NullInt64
var deletedKeyOwnerUserID sql.NullInt64
+ var detailAPIKeyName string
+ var detailAPIKeyDeletedAt sql.NullTime
err := r.db.QueryRowContext(ctx, q, id).Scan(
&out.ID,
@@ -492,6 +516,8 @@ LIMIT 1`
&out.DeletedKeyOwnerEmail,
&out.DeletedKeyName,
&out.APIKeyPrefix,
+ &detailAPIKeyName,
+ &detailAPIKeyDeletedAt,
)
if err != nil {
return nil, err
@@ -558,6 +584,14 @@ LIMIT 1`
v := deletedKeyOwnerUserID.Int64
out.DeletedKeyOwnerUserID = &v
}
+ // Key 名称:优先关联到的 ak.name;关联不到时回退快照的 deleted_key_name。
+ if detailAPIKeyName != "" {
+ out.APIKeyName = detailAPIKeyName
+ } else {
+ out.APIKeyName = out.DeletedKeyName
+ }
+ // 已删除:ak.deleted_at 非空(软删),或仅命中 deleted_key_name 兜底。
+ out.APIKeyDeleted = detailAPIKeyDeletedAt.Valid || (detailAPIKeyName == "" && out.DeletedKeyName != "")
// Normalize upstream_errors to empty string when stored as JSON null.
out.UpstreamErrors = strings.TrimSpace(out.UpstreamErrors)
@@ -860,6 +894,14 @@ INSERT INTO ops_system_log_cleanup_audits (
return err
}
+var likePatternReplacer = strings.NewReplacer(`\`, `\\`, `%`, `\%`, `_`, `\_`)
+
+// escapeLikePattern 转义 LIKE/ILIKE 通配符(\ % _),避免用户输入被当作通配符。
+// Postgres 默认以反斜杠为转义符,无需额外 ESCAPE 子句。
+func escapeLikePattern(s string) string {
+ return likePatternReplacer.Replace(s)
+}
+
func buildOpsErrorLogsWhere(filter *service.OpsErrorLogFilter) (string, []any) {
clauses := make([]string, 0, 12)
args := make([]any, 0, 12)
@@ -972,6 +1014,41 @@ func buildOpsErrorLogsWhere(filter *service.OpsErrorLogFilter) (string, []any) {
clauses = append(clauses, "EXISTS (SELECT 1 FROM users u WHERE u.id = e.user_id AND u.email ILIKE $"+n+")")
}
+ if filter.UserID != nil && *filter.UserID > 0 {
+ args = append(args, *filter.UserID)
+ n := itoa(len(args))
+ if filter.MatchDeletedKeyOwner {
+ // 用户侧:把「删 key 后认证失败」(user_id=NULL,靠 deleted_key_owner 归因)的记录也纳入。
+ clauses = append(clauses, "(e.user_id = $"+n+" OR e.deleted_key_owner_user_id = $"+n+")")
+ } else {
+ clauses = append(clauses, "e.user_id = $"+n)
+ }
+ }
+ if filter.APIKeyID != nil && *filter.APIKeyID > 0 {
+ args = append(args, *filter.APIKeyID)
+ clauses = append(clauses, "e.api_key_id = $"+itoa(len(args)))
+ }
+ if m := strings.TrimSpace(filter.Model); m != "" {
+ if filter.ModelFuzzy {
+ args = append(args, "%"+escapeLikePattern(m)+"%")
+ clauses = append(clauses, "COALESCE(e.requested_model, e.model, '') ILIKE $"+itoa(len(args)))
+ } else {
+ args = append(args, m)
+ clauses = append(clauses, "COALESCE(e.requested_model, e.model, '') = $"+itoa(len(args)))
+ }
+ }
+ if filter.ExcludeCountTokens {
+ clauses = append(clauses, "COALESCE(e.is_count_tokens, false) = false")
+ }
+ if len(filter.ErrorPhasesAny) > 0 {
+ args = append(args, pq.Array(filter.ErrorPhasesAny))
+ clauses = append(clauses, "e.error_phase = ANY($"+itoa(len(args))+")")
+ }
+ if len(filter.ErrorTypesAny) > 0 {
+ args = append(args, pq.Array(filter.ErrorTypesAny))
+ clauses = append(clauses, "e.error_type = ANY($"+itoa(len(args))+")")
+ }
+
return "WHERE " + strings.Join(clauses, " AND "), args
}
diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go
index d4901da1..766225ff 100644
--- a/backend/internal/server/api_contract_test.go
+++ b/backend/internal/server/api_contract_test.go
@@ -896,7 +896,8 @@ func TestAPIContracts(t *testing.T) {
"wechat_connect_mobile_app_secret_configured": false,
"wechat_connect_redirect_url": "",
"wechat_connect_frontend_redirect_url": "/auth/wechat/callback",
- "wechat_connect_scopes": "snsapi_login"
+ "wechat_connect_scopes": "snsapi_login",
+ "allow_user_view_error_requests": false
}
}`,
},
@@ -1167,7 +1168,8 @@ func TestAPIContracts(t *testing.T) {
"auth_source_default_dingtalk_subscriptions": [],
"auth_source_default_dingtalk_grant_on_signup": false,
"auth_source_default_dingtalk_grant_on_first_bind": false,
- "force_email_on_third_party_signup": false
+ "force_email_on_third_party_signup": false,
+ "allow_user_view_error_requests": false
}
}`,
},
@@ -1279,7 +1281,7 @@ func newContractDeps(t *testing.T) *contractDeps {
adminService := service.NewAdminService(userRepo, groupRepo, &accountRepo, proxyRepo, apiKeyRepo, redeemRepo, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
authHandler := handler.NewAuthHandler(cfg, nil, userService, settingService, nil, redeemService, nil, nil)
apiKeyHandler := handler.NewAPIKeyHandler(apiKeyService)
- usageHandler := handler.NewUsageHandler(usageService, apiKeyService)
+ usageHandler := handler.NewUsageHandler(usageService, apiKeyService, nil, nil)
adminSettingHandler := adminhandler.NewSettingHandler(settingService, nil, nil, nil, nil, nil, nil)
adminAccountHandler := adminhandler.NewAccountHandler(adminService, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
diff --git a/backend/internal/server/routes/user.go b/backend/internal/server/routes/user.go
index 07ae33de..0f3758f7 100644
--- a/backend/internal/server/routes/user.go
+++ b/backend/internal/server/routes/user.go
@@ -82,6 +82,8 @@ func RegisterUserRoutes(
usage := authenticated.Group("/usage")
{
usage.GET("", h.Usage.List)
+ usage.GET("/errors", h.Usage.ListErrors)
+ usage.GET("/errors/:id", h.Usage.GetErrorDetail)
usage.GET("/:id", h.Usage.GetByID)
usage.GET("/stats", h.Usage.Stats)
// User dashboard endpoints
diff --git a/backend/internal/service/domain_constants.go b/backend/internal/service/domain_constants.go
index b6441238..11245d00 100644
--- a/backend/internal/service/domain_constants.go
+++ b/backend/internal/service/domain_constants.go
@@ -463,3 +463,7 @@ func SettingKeyAuthSourcePlatformQuotas(source string) string {
// AdminAPIKeyPrefix is the prefix for admin API keys (distinct from user "sk-" keys).
const AdminAPIKeyPrefix = "admin-"
+
+// SettingKeyAllowUserViewErrorRequests controls whether end users can view
+// their own failed requests on the usage page. Default false (opt-in).
+const SettingKeyAllowUserViewErrorRequests = "allow_user_view_error_requests"
diff --git a/backend/internal/service/ops_models.go b/backend/internal/service/ops_models.go
index 63c58cad..0bbe4220 100644
--- a/backend/internal/service/ops_models.go
+++ b/backend/internal/service/ops_models.go
@@ -64,6 +64,10 @@ type OpsErrorLog struct {
RequestedModel string `json:"requested_model"`
UpstreamModel string `json:"upstream_model"`
RequestType *int16 `json:"request_type"`
+
+ // 关联 api_key 名称(LEFT JOIN api_keys 取得;软删只覆盖 key 列,name 保留,故已删 key 仍有原名)。
+ APIKeyName string `json:"api_key_name,omitempty"`
+ APIKeyDeleted bool `json:"api_key_deleted,omitempty"`
}
type OpsErrorLogDetail struct {
@@ -108,7 +112,7 @@ type OpsErrorLogFilter struct {
StatusCodes []int
StatusCodesOther bool
- Phase string
+ Phase string // Special: Phase=="upstream" bypasses status>=400 clause; do not set together with ErrorPhasesAny.
Owner string
Source string
Resolved *bool
@@ -119,6 +123,32 @@ type OpsErrorLogFilter struct {
RequestID string
ClientRequestID string
+ // User-scoped filters (used by the user-facing error requests endpoint and
+ // by admin drill-down from the usage page).
+ UserID *int64
+ APIKeyID *int64
+
+ // MatchDeletedKeyOwner: 用户侧专用。UserID 设置且为 true 时,归属从 user_id=UserID
+ // 放宽为 (user_id=UserID OR deleted_key_owner_user_id=UserID),使原所有者能看到
+ // 自己「已删除 key 认证失败」的记录。admin 路径不设此开关 → 行为不变。
+ MatchDeletedKeyOwner bool
+
+ // Model matches against requested_model first, then model.
+ Model string
+ // ModelFuzzy 为 true 时 Model 走 ILIKE 模糊匹配(仅用户端启用);false(默认)保持精确 =,管理端语义不变。
+ ModelFuzzy bool
+
+ // ExcludeCountTokens drops count_tokens probe errors (is_count_tokens=true).
+ ExcludeCountTokens bool
+
+ // ErrorPhasesAny / ErrorTypesAny add plain ANY() filters WITHOUT touching the
+ // special-cased single `Phase` field (only Phase=="upstream" bypasses the status>=400 clause).
+ // NOTE: these ANY filters do NOT bypass status>=400; records with error_phase='upstream'
+ // but status_code<400 (recovered upstream errors) remain excluded.
+ // Used to map user-facing coarse categories to backend conditions.
+ ErrorPhasesAny []string
+ ErrorTypesAny []string
+
// View controls error categorization for list endpoints.
// - errors: show actionable errors (exclude business-limited / 429 / 529)
// - excluded: only show excluded errors
diff --git a/backend/internal/service/ops_service.go b/backend/internal/service/ops_service.go
index 3d234b28..a8c8a4bb 100644
--- a/backend/internal/service/ops_service.go
+++ b/backend/internal/service/ops_service.go
@@ -338,6 +338,50 @@ func (s *OpsService) GetErrorLogs(ctx context.Context, filter *OpsErrorLogFilter
return result, nil
}
+// ListUserErrorRequests 返回某个用户自己的错误请求(精简脱敏)。
+// 强制:仅当前用户、View=all(含业务限流/余额类)、排除 count_tokens 噪声。
+func (s *OpsService) ListUserErrorRequests(ctx context.Context, userID int64, filter *OpsErrorLogFilter) (*UserErrorRequestList, error) {
+ if filter == nil {
+ filter = &OpsErrorLogFilter{}
+ }
+ f := *filter // 拷贝快照,避免原地篡改调用方的 filter(slice 字段只读,浅拷贝足够)
+ filter = &f
+ uid := userID
+ filter.UserID = &uid
+ // 用户侧放宽归属:纳入「删 key 后认证失败」(user_id=NULL,靠 deleted_key_owner 归因)的记录。
+ filter.MatchDeletedKeyOwner = true
+ // APIKeyID 透传:保留 handler 传入的值。安全由 buildOpsErrorLogsWhere 的
+ // "user_id = 自己 AND api_key_id = X" 双重约束保证——传入他人 key 只会得到空集,无泄露。
+ filter.View = "all"
+ filter.ExcludeCountTokens = true
+ filter.ModelFuzzy = true // 用户端模型过滤走 ILIKE 模糊;管理端不设此字段,保持精确
+ // 防御:用户端不接受这些 admin-only / 特殊维度
+ filter.UserQuery = ""
+ filter.Owner = ""
+ filter.Source = ""
+ // 清空 Phase 是防御:Phase 是单值特殊字段,仅当其 == "upstream" 时 buildOpsErrorLogsWhere 才跳过 status>=400 子句。
+ // 用户端一律改走 category→ErrorPhasesAny/ErrorTypesAny(纯 ANY 过滤,不影响 status>=400 子句),
+ // 因此 recovered upstream(error_phase='upstream' 但 status<400,最终成功返回)记录对用户不可见——符合预期。
+ filter.Phase = ""
+
+ list, err := s.opsRepo.ListErrorLogs(ctx, filter)
+ if err != nil {
+ return nil, err
+ }
+ items := make([]*UserErrorRequest, 0, len(list.Errors))
+ for _, e := range list.Errors {
+ if r := ToUserErrorRequest(e); r != nil {
+ items = append(items, r)
+ }
+ }
+ return &UserErrorRequestList{
+ Items: items,
+ Total: list.Total,
+ Page: list.Page,
+ PageSize: list.PageSize,
+ }, nil
+}
+
func (s *OpsService) GetErrorLogByID(ctx context.Context, id int64) (*OpsErrorLogDetail, error) {
if err := s.RequireMonitoringEnabled(ctx); err != nil {
return nil, err
@@ -355,6 +399,31 @@ func (s *OpsService) GetErrorLogByID(ctx context.Context, id int64) (*OpsErrorLo
return detail, nil
}
+// GetUserErrorRequestDetail 返回某用户自己某条错误请求的脱敏详情(含 error_body)。
+// 安全:强制按用户归属校验;非本人记录一律返回 NotFound(不泄露存在性)。
+func (s *OpsService) GetUserErrorRequestDetail(ctx context.Context, userID, id int64) (*UserErrorRequestDetail, error) {
+ if s.opsRepo == nil {
+ return nil, infraerrors.NotFound("OPS_ERROR_NOT_FOUND", "ops error log not found")
+ }
+ if id <= 0 {
+ return nil, infraerrors.BadRequest("OPS_ERROR_INVALID_ID", "invalid error id")
+ }
+ detail, err := s.opsRepo.GetErrorLogByID(ctx, id)
+ if err != nil {
+ if errors.Is(err, sql.ErrNoRows) {
+ return nil, infraerrors.NotFound("OPS_ERROR_NOT_FOUND", "ops error log not found")
+ }
+ return nil, infraerrors.InternalServer("OPS_ERROR_LOAD_FAILED", "Failed to load ops error log").WithCause(err)
+ }
+ // 归属:直接归属(user_id)或经「已删除 key 归因」(deleted_key_owner_user_id)二者之一即可。
+ ownedDirectly := detail.UserID != nil && *detail.UserID == userID
+ ownedViaDeletedKey := detail.DeletedKeyOwnerUserID != nil && *detail.DeletedKeyOwnerUserID == userID
+ if !ownedDirectly && !ownedViaDeletedKey {
+ return nil, infraerrors.NotFound("OPS_ERROR_NOT_FOUND", "ops error log not found")
+ }
+ return ToUserErrorRequestDetail(detail), nil
+}
+
// LookupDeletedKeyAudit 按明文 key 反查已删除 key 的原所有者;未命中或未启用返回 (nil, nil)。
func (s *OpsService) LookupDeletedKeyAudit(ctx context.Context, key string) (*DeletedKeyAuditResult, error) {
if s.opsRepo == nil {
diff --git a/backend/internal/service/ops_service_user_error_test.go b/backend/internal/service/ops_service_user_error_test.go
new file mode 100644
index 00000000..9027ff07
--- /dev/null
+++ b/backend/internal/service/ops_service_user_error_test.go
@@ -0,0 +1,222 @@
+package service
+
+import (
+ "context"
+ "database/sql"
+ "testing"
+
+ infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
+)
+
+type stubOpsRepoForUserErr struct {
+ OpsRepository // 嵌入接口,未实现的方法 panic,仅覆盖 ListErrorLogs
+ gotFilter *OpsErrorLogFilter
+
+ // GetErrorLogByID 控制字段
+ detailToReturn *OpsErrorLogDetail
+ detailErrToReturn error
+}
+
+func (s *stubOpsRepoForUserErr) ListErrorLogs(ctx context.Context, f *OpsErrorLogFilter) (*OpsErrorLogList, error) {
+ s.gotFilter = f
+ return &OpsErrorLogList{
+ Errors: []*OpsErrorLog{{
+ Phase: "request", Type: "rate_limit_error",
+ Model: "m", RequestedModel: "rm", StatusCode: 429,
+ Message: "secret", UserEmail: "a@b.c",
+ }},
+ Total: 1, Page: 1, PageSize: 20,
+ }, nil
+}
+
+func (s *stubOpsRepoForUserErr) GetErrorLogByID(ctx context.Context, id int64) (*OpsErrorLogDetail, error) {
+ if s.detailErrToReturn != nil {
+ return nil, s.detailErrToReturn
+ }
+ return s.detailToReturn, nil
+}
+
+func TestListUserErrorRequests_ForcesScopeAndRedacts(t *testing.T) {
+ stub := &stubOpsRepoForUserErr{}
+ svc := &OpsService{opsRepo: stub}
+ uid := int64(42)
+ kid := int64(7)
+ in := &OpsErrorLogFilter{UserID: nil, View: "errors", Phase: "upstream", APIKeyID: &kid}
+ out, err := svc.ListUserErrorRequests(context.Background(), uid, in)
+ if err != nil {
+ t.Fatal(err)
+ }
+ // 强制按用户
+ if stub.gotFilter.UserID == nil || *stub.gotFilter.UserID != uid {
+ t.Fatalf("UserID not forced: %+v", stub.gotFilter.UserID)
+ }
+ // 强制 View=all(含业务限流/余额)
+ if stub.gotFilter.View != "all" {
+ t.Fatalf("View not forced to all: %q", stub.gotFilter.View)
+ }
+ // 强制排除 count_tokens
+ if !stub.gotFilter.ExcludeCountTokens {
+ t.Fatal("ExcludeCountTokens not forced")
+ }
+ // 强制清空 Phase(防止 "upstream" 绕过 status>=400 子句 + 与 ErrorPhasesAny 双重约束)
+ if stub.gotFilter.Phase != "" {
+ t.Fatalf("Phase not cleared: %q", stub.gotFilter.Phase)
+ }
+ // APIKeyID 透传保留(用户可按自己 key 过滤;越权由 user_id AND api_key_id 双重防护)
+ if stub.gotFilter.APIKeyID == nil || *stub.gotFilter.APIKeyID != kid {
+ t.Fatalf("APIKeyID should be preserved, got %v", stub.gotFilter.APIKeyID)
+ }
+ // 调用方传入的 filter 不应被原地篡改(验证 shallow copy 隔离生效)
+ if in.View != "errors" || in.UserID != nil || in.Phase != "upstream" {
+ t.Fatalf("caller filter was mutated: View=%q UserID=%v Phase=%q", in.View, in.UserID, in.Phase)
+ }
+ // 脱敏:返回条目含 message 字段
+ if len(out.Items) != 1 || out.Items[0].Category != "rate_limit" || out.Items[0].Model != "rm" {
+ t.Fatalf("bad item: %+v", out.Items)
+ }
+}
+
+func TestGetUserErrorRequestDetail_OwnershipEnforced(t *testing.T) {
+ ownerUID := int64(999)
+ callerUID := int64(1)
+ upstreamStatus := 503
+
+ detail := &OpsErrorLogDetail{
+ OpsErrorLog: OpsErrorLog{
+ ID: 42,
+ Phase: "upstream",
+ Type: "api_error",
+ Model: "gpt-4",
+ RequestedModel: "gpt-4-turbo",
+ InboundEndpoint: "/v1/chat/completions",
+ StatusCode: 502,
+ Platform: "openai",
+ Message: "upstream failed",
+ UserID: &ownerUID,
+ },
+ ErrorBody: `{"error":"upstream"}`,
+ UpstreamStatusCode: &upstreamStatus,
+ }
+
+ stub := &stubOpsRepoForUserErr{detailToReturn: detail}
+ svc := &OpsService{opsRepo: stub}
+
+ // 越权调用(callerUID=1,但记录属于 ownerUID=999)→ 应返回 NotFound,detail 为 nil
+ got, err := svc.GetUserErrorRequestDetail(context.Background(), callerUID, 42)
+ if err == nil {
+ t.Fatal("expected error for unauthorized access, got nil")
+ }
+ if got != nil {
+ t.Fatalf("expected nil detail for unauthorized access, got %+v", got)
+ }
+ // 验证错误为 NotFound(不暴露存在性)
+ if !infraerrors.IsNotFound(err) {
+ t.Fatalf("expected NotFound error, got: %v", err)
+ }
+
+ // 合法调用(callerUID=999 = ownerUID)→ 应返回 non-nil detail
+ got2, err2 := svc.GetUserErrorRequestDetail(context.Background(), ownerUID, 42)
+ if err2 != nil {
+ t.Fatalf("expected no error for legitimate access, got %v", err2)
+ }
+ if got2 == nil {
+ t.Fatal("expected non-nil detail for legitimate access")
+ }
+ if got2.ID != 42 {
+ t.Errorf("want ID=42, got %d", got2.ID)
+ }
+ if got2.ErrorBody != `{"error":"upstream"}` {
+ t.Errorf("want ErrorBody=%q, got %q", `{"error":"upstream"}`, got2.ErrorBody)
+ }
+ if got2.UpstreamStatusCode == nil || *got2.UpstreamStatusCode != 503 {
+ t.Errorf("want UpstreamStatusCode=503, got %v", got2.UpstreamStatusCode)
+ }
+ if got2.Message != "upstream failed" {
+ t.Errorf("want Message=%q, got %q", "upstream failed", got2.Message)
+ }
+}
+
+func TestGetUserErrorRequestDetail_NotFound(t *testing.T) {
+ stub := &stubOpsRepoForUserErr{detailErrToReturn: sql.ErrNoRows}
+ svc := &OpsService{opsRepo: stub}
+
+ got, err := svc.GetUserErrorRequestDetail(context.Background(), 1, 999)
+ if err == nil {
+ t.Fatal("expected error for not found, got nil")
+ }
+ if got != nil {
+ t.Fatalf("expected nil detail, got %+v", got)
+ }
+}
+
+func TestGetUserErrorRequestDetail_InvalidID(t *testing.T) {
+ stub := &stubOpsRepoForUserErr{}
+ svc := &OpsService{opsRepo: stub}
+
+ _, err := svc.GetUserErrorRequestDetail(context.Background(), 1, 0)
+ if err == nil {
+ t.Fatal("expected error for id=0")
+ }
+ _, err = svc.GetUserErrorRequestDetail(context.Background(), 1, -5)
+ if err == nil {
+ t.Fatal("expected error for id=-5")
+ }
+}
+
+func TestListUserErrorRequests_EnablesMatchDeletedKeyOwner(t *testing.T) {
+ stub := &stubOpsRepoForUserErr{}
+ svc := &OpsService{opsRepo: stub}
+ uid := int64(42)
+
+ if _, err := svc.ListUserErrorRequests(context.Background(), uid, &OpsErrorLogFilter{}); err != nil {
+ t.Fatal(err)
+ }
+ if stub.gotFilter == nil || !stub.gotFilter.MatchDeletedKeyOwner {
+ t.Fatal("ListUserErrorRequests should enable MatchDeletedKeyOwner for the user scope")
+ }
+}
+
+func TestGetUserErrorRequestDetail_DeletedKeyOwnerAccess(t *testing.T) {
+ ownerUID := int64(777)
+ otherUID := int64(2)
+
+ // 情况2:user_id=NULL,靠 deleted_key_owner_user_id 归因到 ownerUID
+ mk := func() *OpsErrorLogDetail {
+ return &OpsErrorLogDetail{
+ OpsErrorLog: OpsErrorLog{
+ ID: 55,
+ Phase: "auth",
+ Type: "api_error",
+ StatusCode: 401,
+ Message: "Invalid API key",
+ UserID: nil,
+ APIKeyName: "my-old-key",
+ APIKeyDeleted: true,
+ },
+ DeletedKeyOwnerUserID: &ownerUID,
+ }
+ }
+
+ // 原所有者(经 deleted_key 归因)→ 放行
+ svcOwner := &OpsService{opsRepo: &stubOpsRepoForUserErr{detailToReturn: mk()}}
+ got, err := svcOwner.GetUserErrorRequestDetail(context.Background(), ownerUID, 55)
+ if err != nil {
+ t.Fatalf("owner via deleted_key should be allowed, got err: %v", err)
+ }
+ if got == nil || got.ID != 55 {
+ t.Fatalf("expected detail ID=55, got %+v", got)
+ }
+ if !got.KeyDeleted || got.KeyName != "my-old-key" {
+ t.Fatalf("expected KeyDeleted=true KeyName=my-old-key, got %+v", got)
+ }
+
+ // 他人 → NotFound,不泄露存在性
+ svcOther := &OpsService{opsRepo: &stubOpsRepoForUserErr{detailToReturn: mk()}}
+ got2, err2 := svcOther.GetUserErrorRequestDetail(context.Background(), otherUID, 55)
+ if err2 == nil || got2 != nil {
+ t.Fatalf("non-owner should get (nil, NotFound), got detail=%+v err=%v", got2, err2)
+ }
+ if !infraerrors.IsNotFound(err2) {
+ t.Fatalf("expected NotFound, got %v", err2)
+ }
+}
diff --git a/backend/internal/service/ops_user_error.go b/backend/internal/service/ops_user_error.go
new file mode 100644
index 00000000..817c139a
--- /dev/null
+++ b/backend/internal/service/ops_user_error.go
@@ -0,0 +1,123 @@
+package service
+
+import "time"
+
+// UserErrorRequest 是面向终端用户的"错误请求"精简脱敏视图(白名单)。
+// 严禁包含 client_ip / user_agent / account / api_key_prefix / upstream_endpoint /
+// user_email 等敏感或内部字段。注:message(网关标准化错误描述)与 key_name
+// (用户自有 API Key 名称,KeysView 中本就可见)经产品决策对该用户开放;
+// error_body 仅在详情接口(GetUserErrorRequestDetail)按归属校验后返回。
+type UserErrorRequest struct {
+ ID int64 `json:"id"`
+ CreatedAt time.Time `json:"created_at"`
+ Model string `json:"model"`
+ InboundEndpoint string `json:"inbound_endpoint"`
+ StatusCode int `json:"status_code"`
+ Category string `json:"category"`
+ Platform string `json:"platform"`
+ Message string `json:"message"`
+ KeyName string `json:"key_name"`
+ KeyDeleted bool `json:"key_deleted"`
+}
+
+// UserErrorRequestList 是用户错误请求分页结果。
+type UserErrorRequestList struct {
+ Items []*UserErrorRequest `json:"items"`
+ Total int `json:"total"`
+ Page int `json:"page"`
+ PageSize int `json:"page_size"`
+}
+
+// MapUserErrorCategory 把后端 error_phase + error_type 映射为用户侧粗分类码。
+// 返回的是稳定的分类 code(前端做 i18n),不是展示文案。
+func MapUserErrorCategory(phase, errType string) string {
+ switch phase {
+ case "auth":
+ return "auth"
+ case "routing":
+ return "service_unavailable"
+ case "upstream", "network":
+ return "upstream"
+ case "internal":
+ return "internal"
+ case "request":
+ switch errType {
+ case "rate_limit_error":
+ return "rate_limit"
+ case "billing_error", "subscription_error":
+ return "quota"
+ case "invalid_request_error":
+ return "invalid_request"
+ }
+ }
+ return "other"
+}
+
+// CategoryToFilter 把用户侧分类码反向映射为后端过滤条件(plain ANY)。
+// 未知分类返回两个空切片(即不施加分类过滤)。
+// 注意:"other" 与未知分类都走 default 返回空切片——"other" 无对应的 phase/type 组合,无法精确反查,因此等价于不过滤。
+func CategoryToFilter(category string) (phases []string, errorTypes []string) {
+ switch category {
+ case "auth":
+ return []string{"auth"}, nil
+ case "service_unavailable":
+ return []string{"routing"}, nil
+ case "upstream":
+ return []string{"upstream", "network"}, nil
+ case "internal":
+ return []string{"internal"}, nil
+ case "rate_limit":
+ return nil, []string{"rate_limit_error"}
+ case "quota":
+ return nil, []string{"billing_error", "subscription_error"}
+ case "invalid_request":
+ return nil, []string{"invalid_request_error"}
+ default:
+ return nil, nil
+ }
+}
+
+// ToUserErrorRequest 把内部 OpsErrorLog 裁剪为用户安全视图。
+func ToUserErrorRequest(e *OpsErrorLog) *UserErrorRequest {
+ if e == nil {
+ return nil
+ }
+ model := e.RequestedModel
+ if model == "" {
+ model = e.Model
+ }
+ return &UserErrorRequest{
+ ID: e.ID,
+ CreatedAt: e.CreatedAt,
+ Model: model,
+ InboundEndpoint: e.InboundEndpoint,
+ StatusCode: e.StatusCode,
+ Category: MapUserErrorCategory(e.Phase, e.Type),
+ Platform: e.Platform,
+ Message: e.Message,
+ KeyName: e.APIKeyName,
+ KeyDeleted: e.APIKeyDeleted,
+ }
+}
+
+// UserErrorRequestDetail 是错误请求详情的脱敏视图(点击单行查看)。
+// 在 UserErrorRequest 基础上额外暴露 error_body(上游错误响应正文)与 upstream_status_code;
+// 仍严禁任何内部/敏感字段。
+type UserErrorRequestDetail struct {
+ UserErrorRequest
+ ErrorBody string `json:"error_body"`
+ UpstreamStatusCode *int `json:"upstream_status_code,omitempty"`
+}
+
+// ToUserErrorRequestDetail 把内部 OpsErrorLogDetail 裁剪为用户安全详情视图。
+func ToUserErrorRequestDetail(e *OpsErrorLogDetail) *UserErrorRequestDetail {
+ if e == nil {
+ return nil
+ }
+ base := ToUserErrorRequest(&e.OpsErrorLog)
+ return &UserErrorRequestDetail{
+ UserErrorRequest: *base,
+ ErrorBody: e.ErrorBody,
+ UpstreamStatusCode: e.UpstreamStatusCode,
+ }
+}
diff --git a/backend/internal/service/ops_user_error_test.go b/backend/internal/service/ops_user_error_test.go
new file mode 100644
index 00000000..31b0c269
--- /dev/null
+++ b/backend/internal/service/ops_user_error_test.go
@@ -0,0 +1,167 @@
+package service
+
+import (
+ "encoding/json"
+ "strings"
+ "testing"
+ "time"
+)
+
+func TestMapUserErrorCategory(t *testing.T) {
+ cases := []struct {
+ phase, etype, want string
+ }{
+ {"auth", "authentication_error", "auth"},
+ {"request", "rate_limit_error", "rate_limit"},
+ {"request", "billing_error", "quota"},
+ {"request", "subscription_error", "quota"},
+ {"request", "invalid_request_error", "invalid_request"},
+ {"routing", "api_error", "service_unavailable"},
+ {"upstream", "upstream_error", "upstream"},
+ {"network", "api_error", "upstream"},
+ {"internal", "api_error", "internal"},
+ {"weird", "weird", "other"},
+ }
+ for _, c := range cases {
+ if got := MapUserErrorCategory(c.phase, c.etype); got != c.want {
+ t.Errorf("MapUserErrorCategory(%q,%q)=%q want %q", c.phase, c.etype, got, c.want)
+ }
+ }
+}
+
+func TestCategoryToFilter(t *testing.T) {
+ phases, types := CategoryToFilter("rate_limit")
+ if len(types) != 1 || types[0] != "rate_limit_error" || len(phases) != 0 {
+ t.Fatalf("rate_limit => phases=%v types=%v", phases, types)
+ }
+ phases, types = CategoryToFilter("auth")
+ if len(phases) != 1 || phases[0] != "auth" || len(types) != 0 {
+ t.Fatalf("auth => phases=%v types=%v", phases, types)
+ }
+ phases, types = CategoryToFilter("service_unavailable")
+ if len(phases) != 1 || phases[0] != "routing" || len(types) != 0 {
+ t.Fatalf("service_unavailable => phases=%v types=%v", phases, types)
+ }
+ phases, types = CategoryToFilter("upstream")
+ if len(phases) != 2 || phases[0] != "upstream" || phases[1] != "network" || len(types) != 0 {
+ t.Fatalf("upstream => phases=%v types=%v", phases, types)
+ }
+ phases, types = CategoryToFilter("internal")
+ if len(phases) != 1 || phases[0] != "internal" || len(types) != 0 {
+ t.Fatalf("internal => phases=%v types=%v", phases, types)
+ }
+ phases, types = CategoryToFilter("quota")
+ if len(types) != 2 || types[0] != "billing_error" || types[1] != "subscription_error" || len(phases) != 0 {
+ t.Fatalf("quota => phases=%v types=%v", phases, types)
+ }
+ phases, types = CategoryToFilter("invalid_request")
+ if len(types) != 1 || types[0] != "invalid_request_error" || len(phases) != 0 {
+ t.Fatalf("invalid_request => phases=%v types=%v", phases, types)
+ }
+ phases, types = CategoryToFilter("other")
+ if len(phases) != 0 || len(types) != 0 {
+ t.Fatalf("other => phases=%v types=%v", phases, types)
+ }
+}
+
+func TestToUserErrorRequest_RedactsSensitiveFields(t *testing.T) {
+ src := &OpsErrorLog{
+ ID: 123,
+ CreatedAt: time.Unix(0, 0).UTC(),
+ Model: "m",
+ RequestedModel: "rm",
+ InboundEndpoint: "/v1/chat/completions",
+ StatusCode: 429,
+ Platform: "openai",
+ Phase: "request",
+ Type: "rate_limit_error",
+ Message: "rate limit exceeded",
+ APIKeyName: "my-key",
+ APIKeyDeleted: true,
+ }
+ out := ToUserErrorRequest(src)
+ if out.ID != 123 {
+ t.Errorf("want ID=123, got %d", out.ID)
+ }
+ if out.Model != "rm" {
+ t.Errorf("want requested_model preferred, got %q", out.Model)
+ }
+ if out.Category != "rate_limit" {
+ t.Errorf("category=%q", out.Category)
+ }
+ if out.StatusCode != 429 || out.InboundEndpoint != "/v1/chat/completions" || out.Platform != "openai" {
+ t.Errorf("basic fields wrong: %+v", out)
+ }
+ if out.Message != "rate limit exceeded" {
+ t.Errorf("want message=%q, got %q", "rate limit exceeded", out.Message)
+ }
+ if out.KeyName != "my-key" {
+ t.Errorf("want key_name=my-key, got %q", out.KeyName)
+ }
+ if !out.KeyDeleted {
+ t.Error("want key_deleted=true")
+ }
+}
+
+func TestToUserErrorRequestDetail_WhitelistAndRedacts(t *testing.T) {
+ uid := int64(42)
+ upstreamStatus := 503
+ src := &OpsErrorLogDetail{
+ OpsErrorLog: OpsErrorLog{
+ ID: 999,
+ CreatedAt: time.Unix(1000, 0).UTC(),
+ Model: "gpt-4",
+ RequestedModel: "gpt-4-turbo",
+ InboundEndpoint: "/v1/chat/completions",
+ StatusCode: 502,
+ Platform: "openai",
+ Phase: "upstream",
+ Type: "api_error",
+ Message: "upstream error",
+ UserID: &uid,
+ UserEmail: "secret@example.com",
+ ClientIP: func() *string { s := "1.2.3.4"; return &s }(),
+ UpstreamEndpoint: "https://api.openai.com/v1/chat/completions",
+ },
+ ErrorBody: `{"error":{"message":"upstream failed","type":"server_error"}}`,
+ UserAgent: "Mozilla/5.0 secret-agent",
+ UpstreamStatusCode: &upstreamStatus,
+ }
+
+ out := ToUserErrorRequestDetail(src)
+ if out == nil {
+ t.Fatal("expected non-nil detail")
+ }
+
+ // 基础字段正确映射
+ if out.ID != 999 {
+ t.Errorf("want ID=999, got %d", out.ID)
+ }
+ if out.Message != "upstream error" {
+ t.Errorf("want message=%q, got %q", "upstream error", out.Message)
+ }
+ if out.ErrorBody != src.ErrorBody {
+ t.Errorf("ErrorBody mismatch")
+ }
+ if out.UpstreamStatusCode == nil || *out.UpstreamStatusCode != 503 {
+ t.Errorf("UpstreamStatusCode mismatch")
+ }
+
+ // 序列化后不含敏感字段
+ b, err := json.Marshal(out)
+ if err != nil {
+ t.Fatalf("json.Marshal failed: %v", err)
+ }
+ raw := string(b)
+ for _, forbidden := range []string{"user_email", "client_ip", "upstream_endpoint", "user_agent"} {
+ if strings.Contains(raw, forbidden) {
+ t.Errorf("sensitive field %q leaked in JSON output: %s", forbidden, raw)
+ }
+ }
+}
+
+func TestToUserErrorRequestDetail_Nil(t *testing.T) {
+ if out := ToUserErrorRequestDetail(nil); out != nil {
+ t.Errorf("expected nil for nil input, got %+v", out)
+ }
+}
diff --git a/backend/internal/service/setting_service.go b/backend/internal/service/setting_service.go
index 98acdb80..7043736a 100644
--- a/backend/internal/service/setting_service.go
+++ b/backend/internal/service/setting_service.go
@@ -761,6 +761,7 @@ func (s *SettingService) GetPublicSettings(ctx context.Context) (*PublicSettings
SettingKeyAvailableChannelsEnabled,
SettingKeyAffiliateEnabled,
SettingKeyRiskControlEnabled,
+ SettingKeyAllowUserViewErrorRequests,
}
settings, err := s.settingRepo.GetMultiple(ctx, keys)
@@ -873,6 +874,8 @@ func (s *SettingService) GetPublicSettings(ctx context.Context) (*PublicSettings
AffiliateEnabled: settings[SettingKeyAffiliateEnabled] == "true",
RiskControlEnabled: settings[SettingKeyRiskControlEnabled] == "true",
+
+ AllowUserViewErrorRequests: settings[SettingKeyAllowUserViewErrorRequests] == "true",
}, nil
}
@@ -950,6 +953,17 @@ func (s *SettingService) GetAvailableChannelsRuntime(ctx context.Context) Availa
}
}
+// IsUserErrorViewAllowed reads the user-facing error-requests visibility switch
+// directly from the settings store. Fail-closed: on error returns false (opt-in default).
+func (s *SettingService) IsUserErrorViewAllowed(ctx context.Context) bool {
+ vals, err := s.settingRepo.GetMultiple(ctx, []string{SettingKeyAllowUserViewErrorRequests})
+ if err != nil {
+ slog.Warn("failed to get allow_user_view_error_requests setting, defaulting to false", "error", err)
+ return false
+ }
+ return vals[SettingKeyAllowUserViewErrorRequests] == "true"
+}
+
// GetAntigravityUserAgentVersion 返回 Antigravity 上游请求使用的版本号。
// 后台设置优先;为空、缺失或非法时回退到 ANTIGRAVITY_USER_AGENT_VERSION / 内置默认值。
func (s *SettingService) GetAntigravityUserAgentVersion(ctx context.Context) string {
@@ -1175,6 +1189,7 @@ type PublicSettingsInjectionPayload struct {
AvailableChannelsEnabled bool `json:"available_channels_enabled"`
AffiliateEnabled bool `json:"affiliate_enabled"`
RiskControlEnabled bool `json:"risk_control_enabled"`
+ AllowUserViewErrorRequests bool `json:"allow_user_view_error_requests"`
}
// GetPublicSettingsForInjection returns public settings in a format suitable for HTML injection.
@@ -1237,6 +1252,7 @@ func (s *SettingService) GetPublicSettingsForInjection(ctx context.Context) (any
AvailableChannelsEnabled: settings.AvailableChannelsEnabled,
AffiliateEnabled: settings.AffiliateEnabled,
RiskControlEnabled: settings.RiskControlEnabled,
+ AllowUserViewErrorRequests: settings.AllowUserViewErrorRequests,
}, nil
}
@@ -1924,6 +1940,8 @@ func (s *SettingService) buildSystemSettingsUpdates(ctx context.Context, setting
updates[SettingKeyDefaultPlatformQuotas] = string(blob)
}
+ updates[SettingKeyAllowUserViewErrorRequests] = strconv.FormatBool(settings.AllowUserViewErrorRequests)
+
return updates, nil
}
@@ -2816,6 +2834,8 @@ func (s *SettingService) InitializeDefaultSettings(ctx context.Context) error {
SettingPaymentVisibleMethodAlipayEnabled: "false",
SettingPaymentVisibleMethodWxpayEnabled: "false",
openAIAdvancedSchedulerSettingKey: "false",
+
+ SettingKeyAllowUserViewErrorRequests: "false",
}
return s.settingRepo.SetMultiple(ctx, defaults)
@@ -3373,6 +3393,8 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin
}
}
+ result.AllowUserViewErrorRequests = settings[SettingKeyAllowUserViewErrorRequests] == "true" // default false
+
return result
}
diff --git a/backend/internal/service/setting_service_public_test.go b/backend/internal/service/setting_service_public_test.go
index 2faa4d82..621620fd 100644
--- a/backend/internal/service/setting_service_public_test.go
+++ b/backend/internal/service/setting_service_public_test.go
@@ -91,6 +91,19 @@ func TestSettingService_GetPublicSettings_ExposesForceEmailOnThirdPartySignup(t
require.True(t, settings.ForceEmailOnThirdPartySignup)
}
+func TestSettingService_GetPublicSettings_ExposesAllowUserViewErrorRequests(t *testing.T) {
+ repo := &settingPublicRepoStub{
+ values: map[string]string{
+ SettingKeyAllowUserViewErrorRequests: "true",
+ },
+ }
+ svc := NewSettingService(repo, &config.Config{})
+
+ settings, err := svc.GetPublicSettings(context.Background())
+ require.NoError(t, err)
+ require.True(t, settings.AllowUserViewErrorRequests)
+}
+
func TestSettingService_GetPublicSettings_ExposesWeChatOAuthModeCapabilities(t *testing.T) {
svc := NewSettingService(&settingPublicRepoStub{
values: map[string]string{
diff --git a/backend/internal/service/setting_service_user_error_persist_test.go b/backend/internal/service/setting_service_user_error_persist_test.go
new file mode 100644
index 00000000..1ec76d46
--- /dev/null
+++ b/backend/internal/service/setting_service_user_error_persist_test.go
@@ -0,0 +1,31 @@
+//go:build unit
+
+package service
+
+import (
+ "context"
+ "testing"
+
+ "github.com/Wei-Shaw/sub2api/internal/config"
+ "github.com/stretchr/testify/require"
+)
+
+// TestAllowUserViewErrorRequests_PersistsToDB 验证 buildSystemSettingsUpdates 会将
+// AllowUserViewErrorRequests 写入 updates map(即最终落库),这是对 bug 的回归测试:
+// 该字段曾因漏写而永远无法持久化。
+func TestAllowUserViewErrorRequests_PersistsToDB(t *testing.T) {
+ // bmUpdateRepoStub 已在 setting_service_backend_mode_test.go 中定义(同 package)。
+ // 本测试不触及需要 GetValue 的设置项,getValueFn 设为 nil 即可,无需 stub。
+ repo := &bmUpdateRepoStub{}
+ svc := NewSettingService(repo, &config.Config{})
+
+ err := svc.UpdateSettings(context.Background(), &SystemSettings{
+ AllowUserViewErrorRequests: true,
+ })
+ require.NoError(t, err)
+
+ // 断言 updates 中含有该 key,且值为 "true"
+ val, ok := repo.updates[SettingKeyAllowUserViewErrorRequests]
+ require.True(t, ok, "updates map 中应包含 SettingKeyAllowUserViewErrorRequests,但未找到(bug:buildSystemSettingsUpdates 漏写)")
+ require.Equal(t, "true", val)
+}
diff --git a/backend/internal/service/setting_user_error_view_test.go b/backend/internal/service/setting_user_error_view_test.go
new file mode 100644
index 00000000..264a5a79
--- /dev/null
+++ b/backend/internal/service/setting_user_error_view_test.go
@@ -0,0 +1,9 @@
+package service
+
+import "testing"
+
+func TestSettingKeyAllowUserViewErrorRequests_Constant(t *testing.T) {
+ if SettingKeyAllowUserViewErrorRequests != "allow_user_view_error_requests" {
+ t.Fatalf("unexpected key: %s", SettingKeyAllowUserViewErrorRequests)
+ }
+}
diff --git a/backend/internal/service/settings_view.go b/backend/internal/service/settings_view.go
index 7b45ef1a..97217202 100644
--- a/backend/internal/service/settings_view.go
+++ b/backend/internal/service/settings_view.go
@@ -223,6 +223,9 @@ type SystemSettings struct {
// 系统全局默认平台配额(key = platform,nil/缺省 = 不限制)
DefaultPlatformQuotas map[string]*DefaultPlatformQuotaSetting `json:"default_platform_quotas"`
+
+ // 允许终端用户在用量页查看自己的失败请求
+ AllowUserViewErrorRequests bool
}
type DefaultSubscriptionSetting struct {
@@ -293,6 +296,9 @@ type PublicSettings struct {
// 风控中心功能开关
RiskControlEnabled bool `json:"risk_control_enabled"`
+
+ // 允许终端用户在用量页查看自己的失败请求
+ AllowUserViewErrorRequests bool `json:"allow_user_view_error_requests"`
}
type LoginAgreementDocument struct {
diff --git a/backend/migrations/148_add_ops_error_logs_user_time_index_notx.sql b/backend/migrations/148_add_ops_error_logs_user_time_index_notx.sql
new file mode 100644
index 00000000..54d73ba5
--- /dev/null
+++ b/backend/migrations/148_add_ops_error_logs_user_time_index_notx.sql
@@ -0,0 +1,6 @@
+-- 148_add_ops_error_logs_user_time_index_notx.sql
+-- 用户侧"错误请求"按 user_id 时间倒序分页所需的部分索引。
+-- 非事务迁移(_notx):CREATE INDEX CONCURRENTLY 不可在事务内执行。
+CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_ops_error_logs_user_time
+ ON ops_error_logs (user_id, created_at DESC)
+ WHERE user_id IS NOT NULL;
diff --git a/frontend/src/api/admin/ops.ts b/frontend/src/api/admin/ops.ts
index 557fd00f..defdcf2e 100644
--- a/frontend/src/api/admin/ops.ts
+++ b/frontend/src/api/admin/ops.ts
@@ -1084,6 +1084,8 @@ export type OpsErrorListQueryParams = {
platform?: string
group_id?: number | null
account_id?: number | null
+ user_id?: number
+ api_key_id?: number
phase?: string
error_owner?: string
diff --git a/frontend/src/api/admin/settings.ts b/frontend/src/api/admin/settings.ts
index 6d8e6cee..5be63076 100644
--- a/frontend/src/api/admin/settings.ts
+++ b/frontend/src/api/admin/settings.ts
@@ -612,6 +612,9 @@ export interface SystemSettings {
// OpenAI fast/flex policy
openai_fast_policy_settings?: OpenAIFastPolicySettings;
+
+ // Allow user view error requests
+ allow_user_view_error_requests: boolean;
}
export interface UpdateSettingsRequest {
@@ -842,6 +845,8 @@ export interface UpdateSettingsRequest {
// OpenAI fast/flex policy
openai_fast_policy_settings?: OpenAIFastPolicySettings;
+
+ allow_user_view_error_requests?: boolean;
}
/**
diff --git a/frontend/src/api/usage.ts b/frontend/src/api/usage.ts
index ee08ee9d..f0aec3d6 100644
--- a/frontend/src/api/usage.ts
+++ b/frontend/src/api/usage.ts
@@ -10,7 +10,10 @@ import type {
UsageStatsResponse,
PaginatedResponse,
TrendDataPoint,
- ModelStat
+ ModelStat,
+ UserErrorRequest,
+ UserErrorRequestDetail,
+ UserErrorListParams
} from '@/types'
// ==================== Dashboard Types ====================
@@ -304,6 +307,22 @@ export async function getDashboardApiKeysUsage(
return data
}
+export async function listMyErrorRequests(
+ params: UserErrorListParams,
+ config: { signal?: AbortSignal } = {}
+): Promise> {
+ const { data } = await apiClient.get>('/usage/errors', {
+ ...config,
+ params
+ })
+ return data
+}
+
+export async function getMyErrorDetail(id: number): Promise {
+ const { data } = await apiClient.get(`/usage/errors/${id}`)
+ return data
+}
+
export const usageAPI = {
list,
query,
@@ -316,7 +335,10 @@ export const usageAPI = {
getDashboardTrend,
getDashboardModels,
getMyApiKeyDailyUsage,
- getDashboardApiKeysUsage
+ getDashboardApiKeysUsage,
+ // Error requests
+ listMyErrorRequests,
+ getMyErrorDetail,
}
export default usageAPI
diff --git a/frontend/src/components/common/Select.vue b/frontend/src/components/common/Select.vue
index a8948145..fbb5c545 100644
--- a/frontend/src/components/common/Select.vue
+++ b/frontend/src/components/common/Select.vue
@@ -22,6 +22,18 @@
{{ selectedLabel }}
+
+
+
(), {
searchable: 'auto',
creatable: false,
creatablePrefix: '',
+ clearable: false,
valueKey: 'value',
labelKey: 'label'
})
@@ -239,6 +253,10 @@ const selectedLabel = computed(() => {
return placeholderText.value
})
+const hasValue = computed(
+ () => props.modelValue !== null && props.modelValue !== undefined && props.modelValue !== ''
+)
+
const filteredOptions = computed(() => {
let opts = props.options as any[]
if (isSearchable.value && searchQuery.value) {
@@ -355,6 +373,12 @@ const selectOption = (option: any) => {
triggerRef.value?.focus()
}
+const clearSelection = () => {
+ if (props.disabled) return
+ emit('update:modelValue', null)
+ emit('change', null, null)
+}
+
// Keyboards
const onTriggerKeyDown = () => {
if (!isOpen.value) {
@@ -461,6 +485,12 @@ onUnmounted(() => {
.select-icon {
@apply flex-shrink-0 text-gray-400 dark:text-dark-400;
}
+
+.select-clear {
+ @apply flex flex-shrink-0 cursor-pointer items-center justify-center;
+ @apply rounded text-gray-400 transition-colors;
+ @apply hover:text-gray-600 dark:hover:text-gray-200;
+}