From 1b6a15b4856cc7ef77800078fb2c256ab049ff58 Mon Sep 17 00:00:00 2001 From: wucm667 Date: Sat, 23 May 2026 15:04:24 +0800 Subject: [PATCH 01/79] fix(db-pool): enforce connection lifetime floors --- backend/internal/repository/db_pool.go | 44 ++++++++++++-- backend/internal/repository/db_pool_test.go | 65 +++++++++++++++++---- 2 files changed, 93 insertions(+), 16 deletions(-) diff --git a/backend/internal/repository/db_pool.go b/backend/internal/repository/db_pool.go index d7116ab1..e110068c 100644 --- a/backend/internal/repository/db_pool.go +++ b/backend/internal/repository/db_pool.go @@ -1,12 +1,26 @@ +// Package repository contains persistence infrastructure helpers. +// +// DB pool lifetimes are clamped here because lib/pq starts watchCancel +// goroutines for context-aware queries. If a cloud proxy silently drops idle +// TCP without RST/FIN, those goroutines can block in Read until database/sql +// retires the connection. This is a short-term mitigation; the long-term +// follow-up is migrating PostgreSQL access to jackc/pgx/v5/stdlib. package repository import ( "database/sql" + "log/slog" "time" "github.com/Wei-Shaw/sub2api/internal/config" ) +const ( + defaultConnMaxLifetime = 30 * time.Minute + defaultConnMaxIdleTime = 5 * time.Minute + maxConfiguredConnAge = 24 * time.Hour +) + type dbPoolSettings struct { MaxOpenConns int MaxIdleConns int @@ -14,19 +28,41 @@ type dbPoolSettings struct { ConnMaxIdleTime time.Duration } -func buildDBPoolSettings(cfg *config.Config) dbPoolSettings { +func clampDBPoolSettings(cfg *config.Config) dbPoolSettings { return dbPoolSettings{ MaxOpenConns: cfg.Database.MaxOpenConns, MaxIdleConns: cfg.Database.MaxIdleConns, - ConnMaxLifetime: time.Duration(cfg.Database.ConnMaxLifetimeMinutes) * time.Minute, - ConnMaxIdleTime: time.Duration(cfg.Database.ConnMaxIdleTimeMinutes) * time.Minute, + ConnMaxLifetime: clampDBPoolDuration("database.conn_max_lifetime_minutes", cfg.Database.ConnMaxLifetimeMinutes, defaultConnMaxLifetime), + ConnMaxIdleTime: clampDBPoolDuration("database.conn_max_idle_time_minutes", cfg.Database.ConnMaxIdleTimeMinutes, defaultConnMaxIdleTime), } } +func clampDBPoolDuration(key string, minutes int, fallback time.Duration) time.Duration { + if minutes <= 0 || minutes > int(maxConfiguredConnAge/time.Minute) { + slog.Warn("database connection pool duration clamped", + "key", key, + "before", minutes, + "after", int(fallback/time.Minute), + ) + return fallback + } + + return time.Duration(minutes) * time.Minute +} + func applyDBPoolSettings(db *sql.DB, cfg *config.Config) { - settings := buildDBPoolSettings(cfg) + settings := clampDBPoolSettings(cfg) db.SetMaxOpenConns(settings.MaxOpenConns) db.SetMaxIdleConns(settings.MaxIdleConns) db.SetConnMaxLifetime(settings.ConnMaxLifetime) db.SetConnMaxIdleTime(settings.ConnMaxIdleTime) + + slog.Info("database connection pool configured", + slog.Group("effective", + slog.Int("max_open", settings.MaxOpenConns), + slog.Int("max_idle", settings.MaxIdleConns), + slog.Duration("max_lifetime", settings.ConnMaxLifetime), + slog.Duration("max_idle_time", settings.ConnMaxIdleTime), + ), + ) } diff --git a/backend/internal/repository/db_pool_test.go b/backend/internal/repository/db_pool_test.go index 3868106a..2757f97c 100644 --- a/backend/internal/repository/db_pool_test.go +++ b/backend/internal/repository/db_pool_test.go @@ -11,21 +11,62 @@ import ( _ "github.com/lib/pq" ) -func TestBuildDBPoolSettings(t *testing.T) { - cfg := &config.Config{ - Database: config.DatabaseConfig{ - MaxOpenConns: 50, - MaxIdleConns: 10, - ConnMaxLifetimeMinutes: 30, - ConnMaxIdleTimeMinutes: 5, +func TestClampDBPoolSettings(t *testing.T) { + tests := []struct { + name string + connMaxLifetime int + connMaxIdleTime int + wantMaxLifetime time.Duration + wantConnMaxIdleTime time.Duration + }{ + { + name: "zero values fall back to safe defaults", + connMaxLifetime: 0, + connMaxIdleTime: 0, + wantMaxLifetime: 30 * time.Minute, + wantConnMaxIdleTime: 5 * time.Minute, + }, + { + name: "negative values fall back to safe defaults", + connMaxLifetime: -1, + connMaxIdleTime: -5, + wantMaxLifetime: 30 * time.Minute, + wantConnMaxIdleTime: 5 * time.Minute, + }, + { + name: "reasonable values pass through", + connMaxLifetime: 15, + connMaxIdleTime: 3, + wantMaxLifetime: 15 * time.Minute, + wantConnMaxIdleTime: 3 * time.Minute, + }, + { + name: "values over twenty four hours fall back to safe defaults", + connMaxLifetime: 24*60 + 1, + connMaxIdleTime: 24*60 + 1, + wantMaxLifetime: 30 * time.Minute, + wantConnMaxIdleTime: 5 * time.Minute, }, } - settings := buildDBPoolSettings(cfg) - require.Equal(t, 50, settings.MaxOpenConns) - require.Equal(t, 10, settings.MaxIdleConns) - require.Equal(t, 30*time.Minute, settings.ConnMaxLifetime) - require.Equal(t, 5*time.Minute, settings.ConnMaxIdleTime) + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cfg := &config.Config{ + Database: config.DatabaseConfig{ + MaxOpenConns: 50, + MaxIdleConns: 10, + ConnMaxLifetimeMinutes: tt.connMaxLifetime, + ConnMaxIdleTimeMinutes: tt.connMaxIdleTime, + }, + } + + settings := clampDBPoolSettings(cfg) + require.Equal(t, 50, settings.MaxOpenConns) + require.Equal(t, 10, settings.MaxIdleConns) + require.Equal(t, tt.wantMaxLifetime, settings.ConnMaxLifetime) + require.Equal(t, tt.wantConnMaxIdleTime, settings.ConnMaxIdleTime) + }) + } } func TestApplyDBPoolSettings(t *testing.T) { From b6a38ddab75c335d1a3881f7ebba23896e23f44c Mon Sep 17 00:00:00 2001 From: wucm667 Date: Tue, 26 May 2026 19:59:12 +0800 Subject: [PATCH 02/79] =?UTF-8?q?feat(admin):=20=E8=B4=A6=E5=8F=B7?= =?UTF-8?q?=E7=AE=A1=E7=90=86=E5=88=97=E8=A1=A8=E6=96=B0=E5=A2=9E=E5=88=9B?= =?UTF-8?q?=E5=BB=BA=E6=97=B6=E9=97=B4=E5=88=97?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../admin/account_handler_list_test.go | 52 +++++++++++++ frontend/src/i18n/locales/en.ts | 1 + frontend/src/i18n/locales/zh.ts | 1 + frontend/src/views/admin/AccountsView.vue | 5 ++ .../__tests__/AccountsView.bulkEdit.spec.ts | 77 ++++++++++++++++++- 5 files changed, 135 insertions(+), 1 deletion(-) create mode 100644 backend/internal/handler/admin/account_handler_list_test.go diff --git a/backend/internal/handler/admin/account_handler_list_test.go b/backend/internal/handler/admin/account_handler_list_test.go new file mode 100644 index 00000000..4d628365 --- /dev/null +++ b/backend/internal/handler/admin/account_handler_list_test.go @@ -0,0 +1,52 @@ +package admin + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func setupAccountListRouter() (*gin.Engine, *stubAdminService) { + gin.SetMode(gin.TestMode) + router := gin.New() + adminSvc := newStubAdminService() + handler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + router.GET("/api/v1/admin/accounts", handler.List) + return router, adminSvc +} + +func TestAccountHandlerListIncludesCreatedAt(t *testing.T) { + router, adminSvc := setupAccountListRouter() + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/accounts?page=1&page_size=20&sort_by=created_at&sort_order=desc", nil) + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, "created_at", adminSvc.lastListAccounts.sortBy) + + var payload struct { + Data struct { + Items []struct { + ID int64 `json:"id"` + CreatedAt string `json:"created_at"` + } `json:"items"` + } `json:"data"` + } + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &payload)) + require.Len(t, payload.Data.Items, 1) + + createdAt := payload.Data.Items[0].CreatedAt + require.NotEmpty(t, createdAt) + require.True(t, strings.HasSuffix(createdAt, "Z"), "created_at should be serialized as UTC") + parsed, err := time.Parse(time.RFC3339Nano, createdAt) + require.NoError(t, err) + _, offset := parsed.Zone() + require.Equal(t, 0, offset) +} diff --git a/frontend/src/i18n/locales/en.ts b/frontend/src/i18n/locales/en.ts index 1d94fa29..16c4abad 100644 --- a/frontend/src/i18n/locales/en.ts +++ b/frontend/src/i18n/locales/en.ts @@ -3072,6 +3072,7 @@ export default { usageWindows: 'Usage Windows', proxy: 'Proxy', lastUsed: 'Last Used', + createdAt: 'Created', expiresAt: 'Expires At', actions: 'Actions' }, diff --git a/frontend/src/i18n/locales/zh.ts b/frontend/src/i18n/locales/zh.ts index 8fa15e72..0578e8ce 100644 --- a/frontend/src/i18n/locales/zh.ts +++ b/frontend/src/i18n/locales/zh.ts @@ -3110,6 +3110,7 @@ export default { usageWindows: '用量窗口', proxy: '代理', lastUsed: '最近使用', + createdAt: '创建时间', expiresAt: '过期时间', actions: '操作' }, diff --git a/frontend/src/views/admin/AccountsView.vue b/frontend/src/views/admin/AccountsView.vue index 51137c8b..c602225c 100644 --- a/frontend/src/views/admin/AccountsView.vue +++ b/frontend/src/views/admin/AccountsView.vue @@ -301,6 +301,9 @@ + @@ -646,6 +646,109 @@ +
+
+
+ +

+ {{ t("admin.groups.modelsList.hint") }} +

+
+ +
+
+
+ + 已选 {{ createModelsListSelectedCount }} / + {{ createModelsListState.items.length }} + +
+ + +
+
+
+

+ {{ t("admin.groups.modelsList.loading") }} +

+

+ {{ t("admin.groups.modelsList.empty") }} +

+
+ + + {{ item.id }} + + + +
+
+
+
+
+
+
+ +

+ {{ t("admin.groups.modelsList.hint") }} +

+
+ +
+
+
+ + 已选 {{ editModelsListSelectedCount }} / + {{ editModelsListState.items.length }} + +
+ + +
+
+
+

+ {{ t("admin.groups.modelsList.loading") }} +

+

+ {{ t("admin.groups.modelsList.empty") }} +

+
+ + + {{ item.id }} + + + +
+
+
+
+
(null); const sortableGroups = ref([]); const createMessagesDispatchDefaults = createDefaultMessagesDispatchFormState(); const editMessagesDispatchDefaults = createDefaultMessagesDispatchFormState(); +const createModelsListState = reactive(createInitialModelsListState()); +const editModelsListState = reactive(createInitialModelsListState()); +const createModelsListLoading = ref(false); +const editModelsListLoading = ref(false); +const modelsListCandidatesTracker = createModelsListCandidatesTracker(); +const createModelsListSelectedCount = computed( + () => createModelsListState.items.filter((item) => item.selected).length, +); +const editModelsListSelectedCount = computed( + () => editModelsListState.items.filter((item) => item.selected).length, +); const createForm = reactive({ name: "", @@ -3335,6 +3561,52 @@ const removeEditRoutingRule = (rule: ModelRoutingRule) => { editModelRoutingRules.value.splice(index, 1); }; +const resetModelsListState = ( + state: typeof createModelsListState, + config?: Parameters[0], +) => { + const fresh = createInitialModelsListState(config); + state.enabled = fresh.enabled; + state.savedModels = fresh.savedModels; + state.items = fresh.items; +}; + +const loadModelsListCandidates = async ( + mode: "create" | "edit", + groupID: number, + platform: GroupPlatform, +) => { + const request = { mode, groupID, platform }; + const requestID = modelsListCandidatesTracker.next(request); + const state = mode === "create" ? createModelsListState : editModelsListState; + const loadingRef = mode === "create" ? createModelsListLoading : editModelsListLoading; + loadingRef.value = true; + try { + const models = await adminAPI.groups.getModelsListCandidates(groupID, platform); + if (!modelsListCandidatesTracker.isCurrent(requestID, request)) { + return; + } + setModelsListCandidates(state, models); + } catch (error) { + if (!modelsListCandidatesTracker.isCurrent(requestID, request)) { + return; + } + console.error("Error loading group models list candidates:", error); + } finally { + if (modelsListCandidatesTracker.isCurrent(requestID, request)) { + loadingRef.value = false; + } + } +}; + +const moveCreateModelsListItem = (fromIndex: number, toIndex: number) => { + moveModelsListItem(createModelsListState, fromIndex, toIndex); +}; + +const moveEditModelsListItem = (fromIndex: number, toIndex: number) => { + moveModelsListItem(editModelsListState, fromIndex, toIndex); +}; + // 将 UI 格式的路由规则转换为 API 格式 const convertRoutingRulesToApiFormat = ( rules: ModelRoutingRule[], @@ -3624,6 +3896,11 @@ const handleSort = (key: string, order: 'asc' | 'desc') => { loadGroups(); }; +const openCreateModal = () => { + showCreateModal.value = true; + loadModelsListCandidates("create", 0, createForm.platform); +}; + const closeCreateModal = () => { showCreateModal.value = false; createModelRoutingRules.value.forEach((rule) => { @@ -3654,6 +3931,8 @@ const closeCreateModal = () => { createForm.supported_model_scopes = ["claude", "gemini_text", "gemini_image"]; createForm.mcp_xml_inject = true; createForm.copy_accounts_from_group_ids = []; + createForm.rpm_limit = 0; + resetModelsListState(createModelsListState); createModelRoutingRules.value = []; }; @@ -3708,6 +3987,7 @@ const handleCreateGroup = async () => { model_routing: convertRoutingRulesToApiFormat( createModelRoutingRules.value, ), + models_list_config: buildModelsListConfig(createModelsListState), supported_model_scopes: normalizeSupportedModelScopesForPlatform( createForm.platform, createForm.supported_model_scopes, @@ -3794,10 +4074,12 @@ const handleEdit = async (group: AdminGroup) => { editForm.mcp_xml_inject = group.mcp_xml_inject ?? true; editForm.copy_accounts_from_group_ids = []; // 复制账号字段每次编辑时重置为空 editForm.rpm_limit = group.rpm_limit ?? 0; + resetModelsListState(editModelsListState, group.models_list_config); // 加载模型路由规则(异步加载账号名称) editModelRoutingRules.value = await convertApiFormatToRoutingRules( group.model_routing, ); + loadModelsListCandidates("edit", group.id, group.platform); showEditModal.value = true; }; @@ -3811,6 +4093,7 @@ const closeEditModal = () => { editModelRoutingRules.value = []; editForm.copy_accounts_from_group_ids = []; resetMessagesDispatchFormState(editForm); + resetModelsListState(editModelsListState); }; const handleUpdateGroup = async () => { @@ -3843,6 +4126,7 @@ const handleUpdateGroup = async () => { model_routing: convertRoutingRulesToApiFormat( editModelRoutingRules.value, ), + models_list_config: buildModelsListConfig(editModelsListState), supported_model_scopes: normalizeSupportedModelScopesForPlatform( editForm.platform, editForm.supported_model_scopes, @@ -3960,6 +4244,8 @@ watch( createForm.require_oauth_only = false; createForm.require_privacy_set = false; } + resetModelsListState(createModelsListState); + loadModelsListCandidates("create", 0, newVal); }, ); @@ -3976,6 +4262,10 @@ watch( editForm.require_oauth_only = false; editForm.require_privacy_set = false; } + if (editingGroup.value) { + resetModelsListState(editModelsListState, editForm.platform === editingGroup.value.platform ? editingGroup.value.models_list_config : undefined); + loadModelsListCandidates("edit", editingGroup.value.id, newVal); + } }, ); @@ -4049,6 +4339,7 @@ const saveSortOrder = async () => { onMounted(() => { loadGroups(); + loadModelsListCandidates("create", 0, createForm.platform); document.addEventListener("click", handleClickOutside); }); diff --git a/frontend/src/views/admin/__tests__/groupsModelsList.spec.ts b/frontend/src/views/admin/__tests__/groupsModelsList.spec.ts new file mode 100644 index 00000000..ae50c861 --- /dev/null +++ b/frontend/src/views/admin/__tests__/groupsModelsList.spec.ts @@ -0,0 +1,125 @@ +import { describe, expect, it } from "vitest"; + +import { + buildModelsListConfig, + createModelsListState, + hydrateModelsListState, + invertModelsListSelection, + moveModelsListItem, + selectAllModelsListItems, + setModelsListCandidates, + toggleModelsListItem, +} from "../groupsModelsList"; + +describe("groupsModelsList", () => { + it("selects all default candidates for a new disabled config", () => { + const state = createModelsListState(); + + setModelsListCandidates(state, ["gpt-5.5", "gpt-5.4"]); + + expect(state.enabled).toBe(false); + expect(state.items).toEqual([ + { id: "gpt-5.5", selected: true }, + { id: "gpt-5.4", selected: true }, + ]); + }); + + it("keeps saved selections and marks new candidates as unselected when editing", () => { + const state = createModelsListState({ + enabled: true, + models: ["gpt-5.5", "gpt-5.4"], + }); + + setModelsListCandidates(state, ["gpt-5.4", "legacy-gpt", "gpt-5.5"]); + + expect(state.enabled).toBe(true); + expect(state.items).toEqual([ + { id: "gpt-5.5", selected: true }, + { id: "gpt-5.4", selected: true }, + { id: "legacy-gpt", selected: false }, + ]); + }); + + it("preserves explicitly unselected saved candidates when candidates refresh", () => { + const state = createModelsListState({ + enabled: true, + models: ["gpt-5.5"], + }); + + setModelsListCandidates(state, ["gpt-5.5", "gpt-5.4"]); + + expect(state.items).toEqual([ + { id: "gpt-5.5", selected: true }, + { id: "gpt-5.4", selected: false }, + ]); + }); + + it("builds config with selected models in current display order", () => { + const state = hydrateModelsListState({ + enabled: true, + models: ["gpt-5.5", "gpt-5.4", "legacy-gpt"], + }, ["gpt-5.5", "gpt-5.4", "legacy-gpt"]); + + toggleModelsListItem(state, "legacy-gpt"); + moveModelsListItem(state, 1, 0); + + expect(buildModelsListConfig(state)).toEqual({ + enabled: true, + models: ["gpt-5.4", "gpt-5.5"], + }); + }); + + it("keeps selected models in payload even when disabled so reopening can restore choices", () => { + const state = hydrateModelsListState({ + enabled: false, + models: ["gpt-5.5"], + }, ["gpt-5.5", "gpt-5.4"]); + + expect(buildModelsListConfig(state)).toEqual({ + enabled: false, + models: ["gpt-5.5"], + }); + }); + + it("preserves saved models when candidates have not loaded yet", () => { + const state = createModelsListState({ + enabled: true, + models: ["gpt-5.5", "gpt-5.4"], + }); + + expect(buildModelsListConfig(state)).toEqual({ + enabled: true, + models: ["gpt-5.5", "gpt-5.4"], + }); + }); + + it("selects all candidate models from the toolbar action", () => { + const state = hydrateModelsListState({ + enabled: true, + models: ["gpt-5.5"], + }, ["gpt-5.5", "gpt-5.4", "gpt-5.4-mini"]); + + selectAllModelsListItems(state); + + expect(state.items).toEqual([ + { id: "gpt-5.5", selected: true }, + { id: "gpt-5.4", selected: true }, + { id: "gpt-5.4-mini", selected: true }, + ]); + }); + + it("inverts selected models from the toolbar action", () => { + const state = hydrateModelsListState({ + enabled: true, + models: ["gpt-5.5"], + }, ["gpt-5.5", "gpt-5.4", "gpt-5.4-mini"]); + + invertModelsListSelection(state); + + expect(state.items).toEqual([ + { id: "gpt-5.5", selected: false }, + { id: "gpt-5.4", selected: true }, + { id: "gpt-5.4-mini", selected: true }, + ]); + }); +}); diff --git a/frontend/src/views/admin/__tests__/groupsModelsListCandidates.spec.ts b/frontend/src/views/admin/__tests__/groupsModelsListCandidates.spec.ts new file mode 100644 index 00000000..ec292c63 --- /dev/null +++ b/frontend/src/views/admin/__tests__/groupsModelsListCandidates.spec.ts @@ -0,0 +1,65 @@ +import { describe, expect, it } from "vitest"; + +import { + createModelsListCandidatesTracker, +} from "../groupsModelsListCandidates"; + +describe("groupsModelsListCandidates", () => { + it("rejects stale candidate responses after a newer platform request starts", () => { + const tracker = createModelsListCandidatesTracker(); + const first = { + mode: "create" as const, + groupID: 0, + platform: "openai" as const, + }; + const second = { + mode: "create" as const, + groupID: 0, + platform: "anthropic" as const, + }; + + const firstID = tracker.next(first); + const secondID = tracker.next(second); + + expect(tracker.isCurrent(firstID, first)).toBe(false); + expect(tracker.isCurrent(secondID, second)).toBe(true); + }); + + it("rejects responses for a previous edit group even with the same platform", () => { + const tracker = createModelsListCandidatesTracker(); + const first = { + mode: "edit" as const, + groupID: 10, + platform: "openai" as const, + }; + const second = { + mode: "edit" as const, + groupID: 11, + platform: "openai" as const, + }; + + const firstID = tracker.next(first); + tracker.next(second); + + expect(tracker.isCurrent(firstID, first)).toBe(false); + }); + + it("tracks create and edit requests independently", () => { + const tracker = createModelsListCandidatesTracker(); + const editRequest = { + mode: "edit" as const, + groupID: 10, + platform: "openai" as const, + }; + const createRequest = { + mode: "create" as const, + groupID: 0, + platform: "anthropic" as const, + }; + + const editID = tracker.next(editRequest); + tracker.next(createRequest); + + expect(tracker.isCurrent(editID, editRequest)).toBe(true); + }); +}); diff --git a/frontend/src/views/admin/__tests__/groupsModelsListLayout.spec.ts b/frontend/src/views/admin/__tests__/groupsModelsListLayout.spec.ts new file mode 100644 index 00000000..6ac3d769 --- /dev/null +++ b/frontend/src/views/admin/__tests__/groupsModelsListLayout.spec.ts @@ -0,0 +1,19 @@ +import { readFileSync } from "node:fs"; +import { fileURLToPath } from "node:url"; +import { dirname, resolve } from "node:path"; + +import { describe, expect, it } from "vitest"; + +const currentDir = dirname(fileURLToPath(import.meta.url)); +const groupsViewSource = readFileSync( + resolve(currentDir, "../GroupsView.vue"), + "utf8", +); + +describe("groups models list layout", () => { + it("keeps the toolbar outside of the scrolling list content", () => { + expect(groupsViewSource).toContain("overflow-hidden rounded-lg border"); + expect(groupsViewSource).toContain("max-h-64 space-y-2 overflow-y-auto p-2"); + expect(groupsViewSource).not.toContain("sticky top-0"); + }); +}); diff --git a/frontend/src/views/admin/groupsModelsList.ts b/frontend/src/views/admin/groupsModelsList.ts new file mode 100644 index 00000000..790268fe --- /dev/null +++ b/frontend/src/views/admin/groupsModelsList.ts @@ -0,0 +1,121 @@ +export interface ModelsListConfig { + enabled: boolean + models: string[] +} + +export interface ModelsListItem { + id: string + selected: boolean +} + +export interface ModelsListState { + enabled: boolean + savedModels: string[] + items: ModelsListItem[] +} + +export const createModelsListState = ( + config?: Partial | null, +): ModelsListState => ({ + enabled: config?.enabled ?? false, + savedModels: normalizeModels(config?.models ?? []), + items: [], +}) + +export const hydrateModelsListState = ( + config: Partial | null | undefined, + candidates: string[], +): ModelsListState => { + const state = createModelsListState(config) + setModelsListCandidates(state, candidates) + return state +} + +export const setModelsListCandidates = ( + state: ModelsListState, + candidates: string[], +) => { + const normalizedCandidates = normalizeModels(candidates) + const currentSelected = new Set( + state.items.filter(item => item.selected).map(item => item.id), + ) + const currentKnown = new Set(state.items.map(item => item.id)) + const savedSelected = new Set(state.savedModels) + const hasExistingItems = state.items.length > 0 + const selectionOrder = normalizeModels([ + ...state.items.map(item => item.id), + ...state.savedModels, + ...normalizedCandidates, + ]) + + state.items = selectionOrder.map(id => { + const selected = hasExistingItems + ? currentSelected.has(id) + : state.savedModels.length > 0 + ? savedSelected.has(id) + : normalizedCandidates.includes(id) + + return { + id, + selected: selected && (currentKnown.has(id) || savedSelected.has(id) || state.savedModels.length === 0), + } + }) +} + +export const toggleModelsListItem = (state: ModelsListState, modelID: string) => { + const item = state.items.find(item => item.id === modelID) + if (item) { + item.selected = !item.selected + } +} + +export const selectAllModelsListItems = (state: ModelsListState) => { + state.items.forEach(item => { + item.selected = true + }) +} + +export const invertModelsListSelection = (state: ModelsListState) => { + state.items.forEach(item => { + item.selected = !item.selected + }) +} + +export const moveModelsListItem = ( + state: ModelsListState, + fromIndex: number, + toIndex: number, +) => { + if ( + fromIndex === toIndex || + fromIndex < 0 || + toIndex < 0 || + fromIndex >= state.items.length || + toIndex >= state.items.length + ) { + return + } + const [item] = state.items.splice(fromIndex, 1) + state.items.splice(toIndex, 0, item) +} + +export const buildModelsListConfig = (state: ModelsListState): ModelsListConfig => ({ + enabled: state.enabled, + models: state.items.length > 0 + ? state.items.filter(item => item.selected).map(item => item.id) + : [...state.savedModels], +}) + +const normalizeModels = (models: string[]): string[] => { + const seen = new Set() + const out: string[] = [] + for (const raw of models) { + const model = raw.trim() + if (!model || seen.has(model)) { + continue + } + seen.add(model) + out.push(model) + } + return out +} diff --git a/frontend/src/views/admin/groupsModelsListCandidates.ts b/frontend/src/views/admin/groupsModelsListCandidates.ts new file mode 100644 index 00000000..2c722af8 --- /dev/null +++ b/frontend/src/views/admin/groupsModelsListCandidates.ts @@ -0,0 +1,41 @@ +import type { GroupPlatform } from "@/types"; + +export type ModelsListCandidatesMode = "create" | "edit"; + +export interface ModelsListCandidatesRequest { + mode: ModelsListCandidatesMode; + groupID: number; + platform: GroupPlatform; +} + +export interface ModelsListCandidatesTracker { + next(request: ModelsListCandidatesRequest): number; + isCurrent(requestID: number, request: ModelsListCandidatesRequest): boolean; +} + +export const createModelsListCandidatesTracker = (): ModelsListCandidatesTracker => { + let currentRequestID = 0; + const currentByMode: Partial> = {}; + + return { + next(request) { + currentRequestID += 1; + currentByMode[request.mode] = { + id: currentRequestID, + request: { ...request }, + }; + return currentRequestID; + }, + isCurrent(requestID, request) { + const current = currentByMode[request.mode]; + return ( + current?.id === requestID && + current.request.groupID === request.groupID && + current.request.platform === request.platform + ); + }, + }; +}; From 32ea9cfe1f4ccc708a6c369c2a49986500e39d90 Mon Sep 17 00:00:00 2001 From: haichuan Date: Wed, 27 May 2026 20:24:52 +0800 Subject: [PATCH 12/79] fix: fallback to SSE body for API key responses --- .../service/openai_gateway_service.go | 12 ++++---- .../service/openai_gateway_service_test.go | 29 +++++++++++++++++++ 2 files changed, 36 insertions(+), 5 deletions(-) diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index d4921511..640a3810 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -4903,20 +4903,22 @@ func (s *OpenAIGatewayService) handleNonStreamingResponse(ctx context.Context, r if isEventStreamResponse(resp.Header) { return s.handleSSEToJSON(resp, c, body, originalModel, mappedModel) } + bodyLooksLikeSSE := bytes.Contains(body, []byte("data:")) || bytes.Contains(body, []byte("event:")) + // For OAuth accounts, also fall back to a body-content heuristic because // the upstream may omit the Content-Type header while still sending SSE. // This heuristic is NOT applied to API-key accounts to avoid false // positives on JSON responses that coincidentally contain "data:" or // "event:" in their text content. - if account.Type == AccountTypeOAuth { - bodyLooksLikeSSE := bytes.Contains(body, []byte("data:")) || bytes.Contains(body, []byte("event:")) - if bodyLooksLikeSSE { - return s.handleSSEToJSON(resp, c, body, originalModel, mappedModel) - } + if account.Type == AccountTypeOAuth && bodyLooksLikeSSE { + return s.handleSSEToJSON(resp, c, body, originalModel, mappedModel) } usageValue, usageOK := extractOpenAIUsageFromJSONBytes(body) if !usageOK { + if bodyLooksLikeSSE { + return s.handleSSEToJSON(resp, c, body, originalModel, mappedModel) + } return nil, fmt.Errorf("parse response: invalid json response") } usage := &usageValue diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index ef35aa1a..8bed920d 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -2281,6 +2281,35 @@ func TestHandleSSEToJSON_CompletedEventReturnsJSON(t *testing.T) { require.NotContains(t, rec.Body.String(), "data:") } +func TestHandleNonStreamingResponse_APIKeyFallsBackToSSEBodyWhenContentTypeIsWrong(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + + svc := &OpenAIGatewayService{cfg: &config.Config{}} + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(strings.Join([]string{ + `data: {"type":"response.output_text.delta","delta":"hel"}`, + `data: {"type":"response.output_text.delta","delta":"lo"}`, + `data: {"type":"response.completed","response":{"id":"resp_api_key_sse","object":"response","model":"gpt-5.4","status":"completed","output":[],"usage":{"input_tokens":3,"output_tokens":2,"total_tokens":5}}}`, + `data: [DONE]`, + }, "\n"))), + } + account := &Account{ID: 1, Type: AccountTypeAPIKey} + + result, err := svc.handleNonStreamingResponse(context.Background(), resp, c, account, "gpt-5.4", "gpt-5.4") + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, 3, result.InputTokens) + require.Equal(t, 2, result.OutputTokens) + require.NotContains(t, rec.Body.String(), "data:") + require.Equal(t, "resp_api_key_sse", gjson.Get(rec.Body.String(), "id").String()) + require.Equal(t, "hello", gjson.Get(rec.Body.String(), "output.0.content.0.text").String()) +} + func TestHandleSSEToJSON_ReconstructsImageGenerationOutputItemDone(t *testing.T) { gin.SetMode(gin.TestMode) rec := httptest.NewRecorder() From 89d96f4b25c6e5ead427f67b54a8aae5fdf993e4 Mon Sep 17 00:00:00 2001 From: "github-actions[bot]" <41898282+github-actions[bot]@users.noreply.github.com> Date: Wed, 27 May 2026 14:28:22 +0000 Subject: [PATCH 13/79] chore: sync VERSION to 0.1.132 [skip ci] --- backend/cmd/server/VERSION | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/backend/cmd/server/VERSION b/backend/cmd/server/VERSION index 66c01044..7b9dfc4d 100644 --- a/backend/cmd/server/VERSION +++ b/backend/cmd/server/VERSION @@ -1 +1 @@ -0.1.131 +0.1.132 From 89dffdd2e1915bb84f578ac3bdd97ad91ae349f3 Mon Sep 17 00:00:00 2001 From: stabey <36232531+stabey@users.noreply.github.com> Date: Wed, 27 May 2026 22:36:52 +0800 Subject: [PATCH 14/79] fix(apicompat): emit OpenAI-semantic input_tokens when converting Anthropic to Responses Anthropic Messages reports input_tokens excluding cache_read/cache_creation, but OpenAI Responses input_tokens is the total including cached tokens. The reverse converter passed Anthropic's input_tokens straight through, so client-facing prompt_tokens/input_tokens were short by the cached count and cache_creation was dropped entirely. Fix the non-stream path and the streaming state machine to add cache_read + cache_creation back into input_tokens, and track CacheCreationInputTokens on the streaming state. Six downstream paths benefit (Anthropic->Responses, Anthropic->ChatCompletions, Gemini->ChatCompletions, each sync + stream). Co-Authored-By: Claude Opus 4.7 --- .../pkg/apicompat/anthropic_responses_test.go | 136 ++++++++++++++++++ .../anthropic_to_responses_response.go | 40 ++++-- 2 files changed, 168 insertions(+), 8 deletions(-) diff --git a/backend/internal/pkg/apicompat/anthropic_responses_test.go b/backend/internal/pkg/apicompat/anthropic_responses_test.go index bb566081..8997835c 100644 --- a/backend/internal/pkg/apicompat/anthropic_responses_test.go +++ b/backend/internal/pkg/apicompat/anthropic_responses_test.go @@ -1597,3 +1597,139 @@ func TestAnthropicToResponses_TemperatureStrippedForAllGpt5Variants(t *testing.T }) } } + +// --------------------------------------------------------------------------- +// AnthropicToResponsesResponse: Anthropic input_tokens excludes cached tokens +// while OpenAI Responses input_tokens is the total including cached tokens. +// --------------------------------------------------------------------------- + +func TestAnthropicToResponsesResponse_CacheTokensUseOpenAIInputSemantics(t *testing.T) { + resp := &AnthropicResponse{ + ID: "msg_cache", + Model: "claude-sonnet-4-5-20250929", + Content: []AnthropicContentBlock{ + {Type: "text", Text: "ok"}, + }, + StopReason: "end_turn", + Usage: AnthropicUsage{ + InputTokens: 3318, + OutputTokens: 123, + CacheReadInputTokens: 50688, + CacheCreationInputTokens: 200, + }, + } + + out := AnthropicToResponsesResponse(resp) + require.NotNil(t, out.Usage) + // 3318 (uncached) + 50688 (read) + 200 (creation) = 54206 + assert.Equal(t, 54206, out.Usage.InputTokens) + assert.Equal(t, 123, out.Usage.OutputTokens) + assert.Equal(t, 54329, out.Usage.TotalTokens) + require.NotNil(t, out.Usage.InputTokensDetails) + assert.Equal(t, 50688, out.Usage.InputTokensDetails.CachedTokens) +} + +func TestAnthropicToResponsesResponse_NoCacheTokens(t *testing.T) { + resp := &AnthropicResponse{ + ID: "msg_nocache", + Model: "claude-sonnet-4-5-20250929", + Content: []AnthropicContentBlock{ + {Type: "text", Text: "ok"}, + }, + StopReason: "end_turn", + Usage: AnthropicUsage{ + InputTokens: 100, + OutputTokens: 50, + }, + } + + out := AnthropicToResponsesResponse(resp) + require.NotNil(t, out.Usage) + assert.Equal(t, 100, out.Usage.InputTokens) + assert.Equal(t, 50, out.Usage.OutputTokens) + assert.Equal(t, 150, out.Usage.TotalTokens) + assert.Nil(t, out.Usage.InputTokensDetails) +} + +func TestAnthropicEventToResponses_CacheTokensRoundTripFromMessageStart(t *testing.T) { + state := NewAnthropicEventToResponsesState() + + // message_start carries cache fields on the initial Usage object. + AnthropicEventToResponsesEvents(&AnthropicStreamEvent{ + Type: "message_start", + Message: &AnthropicResponse{ + ID: "msg_stream_cache", + Model: "claude-sonnet-4-5-20250929", + Usage: AnthropicUsage{ + InputTokens: 12, + CacheReadInputTokens: 9, + CacheCreationInputTokens: 3, + }, + }, + }, state) + + AnthropicEventToResponsesEvents(&AnthropicStreamEvent{ + Type: "message_delta", + Usage: &AnthropicUsage{ + OutputTokens: 7, + }, + }, state) + + events := AnthropicEventToResponsesEvents(&AnthropicStreamEvent{Type: "message_stop"}, state) + + // The terminal response.completed event must include OpenAI-semantic usage. + var completed *ResponsesStreamEvent + for i := range events { + if events[i].Type == "response.completed" { + completed = &events[i] + } + } + require.NotNil(t, completed, "response.completed event must be emitted") + require.NotNil(t, completed.Response) + require.NotNil(t, completed.Response.Usage) + // 12 (uncached) + 9 (read) + 3 (creation) = 24 + assert.Equal(t, 24, completed.Response.Usage.InputTokens) + assert.Equal(t, 7, completed.Response.Usage.OutputTokens) + assert.Equal(t, 31, completed.Response.Usage.TotalTokens) + require.NotNil(t, completed.Response.Usage.InputTokensDetails) + assert.Equal(t, 9, completed.Response.Usage.InputTokensDetails.CachedTokens) +} + +func TestAnthropicEventToResponses_CacheTokensFromMessageDelta(t *testing.T) { + state := NewAnthropicEventToResponsesState() + + AnthropicEventToResponsesEvents(&AnthropicStreamEvent{ + Type: "message_start", + Message: &AnthropicResponse{ + ID: "msg_delta_cache", + Model: "claude-sonnet-4-5-20250929", + Usage: AnthropicUsage{InputTokens: 20}, + }, + }, state) + + // Some upstreams only emit cache fields on the final message_delta. + AnthropicEventToResponsesEvents(&AnthropicStreamEvent{ + Type: "message_delta", + Usage: &AnthropicUsage{ + OutputTokens: 8, + CacheReadInputTokens: 11, + CacheCreationInputTokens: 4, + }, + }, state) + + events := AnthropicEventToResponsesEvents(&AnthropicStreamEvent{Type: "message_stop"}, state) + + var completed *ResponsesStreamEvent + for i := range events { + if events[i].Type == "response.completed" { + completed = &events[i] + } + } + require.NotNil(t, completed) + require.NotNil(t, completed.Response.Usage) + // 20 (uncached) + 11 (read) + 4 (creation) = 35 + assert.Equal(t, 35, completed.Response.Usage.InputTokens) + assert.Equal(t, 8, completed.Response.Usage.OutputTokens) + require.NotNil(t, completed.Response.Usage.InputTokensDetails) + assert.Equal(t, 11, completed.Response.Usage.InputTokensDetails.CachedTokens) +} diff --git a/backend/internal/pkg/apicompat/anthropic_to_responses_response.go b/backend/internal/pkg/apicompat/anthropic_to_responses_response.go index 9290e399..de8ab78d 100644 --- a/backend/internal/pkg/apicompat/anthropic_to_responses_response.go +++ b/backend/internal/pkg/apicompat/anthropic_to_responses_response.go @@ -95,10 +95,16 @@ func AnthropicToResponsesResponse(resp *AnthropicResponse) *ResponsesResponse { } // Usage + // Anthropic's input_tokens excludes cache_read/cache_creation, while OpenAI + // Responses' input_tokens is the total including cached tokens. Add them back + // when converting so downstream consumers see OpenAI semantics. + totalInputTokens := resp.Usage.InputTokens + + resp.Usage.CacheReadInputTokens + + resp.Usage.CacheCreationInputTokens out.Usage = &ResponsesUsage{ - InputTokens: resp.Usage.InputTokens, + InputTokens: totalInputTokens, OutputTokens: resp.Usage.OutputTokens, - TotalTokens: resp.Usage.InputTokens + resp.Usage.OutputTokens, + TotalTokens: totalInputTokens + resp.Usage.OutputTokens, } if resp.Usage.CacheReadInputTokens > 0 { out.Usage.InputTokensDetails = &ResponsesInputTokensDetails{ @@ -150,10 +156,13 @@ type AnthropicEventToResponsesState struct { CurrentCallID string CurrentName string - // Usage from message_delta - InputTokens int - OutputTokens int - CacheReadInputTokens int + // Usage from message_start / message_delta. InputTokens here follows + // Anthropic semantics (excludes cached tokens); they are added back when + // emitting the OpenAI Responses usage. + InputTokens int + OutputTokens int + CacheReadInputTokens int + CacheCreationInputTokens int } // NewAnthropicEventToResponsesState returns an initialised stream state. @@ -225,6 +234,12 @@ func anthToResHandleMessageStart(evt *AnthropicStreamEvent, state *AnthropicEven if evt.Message.Usage.InputTokens > 0 { state.InputTokens = evt.Message.Usage.InputTokens } + if evt.Message.Usage.CacheReadInputTokens > 0 { + state.CacheReadInputTokens = evt.Message.Usage.CacheReadInputTokens + } + if evt.Message.Usage.CacheCreationInputTokens > 0 { + state.CacheCreationInputTokens = evt.Message.Usage.CacheCreationInputTokens + } } if state.CreatedSent { @@ -392,9 +407,15 @@ func anthToResHandleMessageDelta(evt *AnthropicStreamEvent, state *AnthropicEven // Update usage if evt.Usage != nil { state.OutputTokens = evt.Usage.OutputTokens + if evt.Usage.InputTokens > 0 { + state.InputTokens = evt.Usage.InputTokens + } if evt.Usage.CacheReadInputTokens > 0 { state.CacheReadInputTokens = evt.Usage.CacheReadInputTokens } + if evt.Usage.CacheCreationInputTokens > 0 { + state.CacheCreationInputTokens = evt.Usage.CacheCreationInputTokens + } } return nil @@ -472,10 +493,13 @@ func makeResponsesCompletedEvent( seq := state.SequenceNumber state.SequenceNumber++ + // Anthropic's input_tokens excludes cache_read/cache_creation; add them + // back to match OpenAI Responses semantics where input_tokens is the total. + totalInputTokens := state.InputTokens + state.CacheReadInputTokens + state.CacheCreationInputTokens usage := &ResponsesUsage{ - InputTokens: state.InputTokens, + InputTokens: totalInputTokens, OutputTokens: state.OutputTokens, - TotalTokens: state.InputTokens + state.OutputTokens, + TotalTokens: totalInputTokens + state.OutputTokens, } if state.CacheReadInputTokens > 0 { usage.InputTokensDetails = &ResponsesInputTokensDetails{ From 56908d3c4cfc2ca6d705321e40dd15e0e32ba68e Mon Sep 17 00:00:00 2001 From: DaydreamCoding <22166516+DaydreamCoding@users.noreply.github.com> Date: Wed, 27 May 2026 19:42:35 +0800 Subject: [PATCH 15/79] =?UTF-8?q?feat(openai):=20codex=5Fcli=5Fonly=20?= =?UTF-8?q?=E6=96=B0=E5=A2=9E=E6=94=BE=E8=A1=8C=20Claude=20Code=20Codex=20?= =?UTF-8?q?=E6=8F=92=E4=BB=B6=E7=9A=84=E6=9C=BA=E5=88=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 适用场景:在 Claude Code 中使用 https://github.com/openai/codex-plugin-cc 插件时,插件经官方 codex app-server 以 clientInfo.name="Claude Code" 完成 initialize 握手,请求头被设为 originator=Claude Code、User-Agent 含 "Claude Code/",不在官方客户端白名单内,原本会被 codex_cli_only 拦截 403。 在官方客户端白名单未命中时评估两层独立放行(OR 语义): - 按账号:account.Extra.codex_cli_only_allowed_clients 引用命名预设 (目前仅 claude_code),detector reason=allowed_client_matched - 全局开关:/admin/settings 网关服务 OpenAI 区块新增 openai_allow_claude_code_codex_plugin(默认 false),开启后对所有 codex_cli_only 账号统一放行,detector reason=global_allowed_client_matched 签名仍要求 originator=Claude Code 精确等值 + UA 含 "Claude Code/"。 上游转发保持透传不变。 Co-Authored-By: Claude Opus 4.7 (1M context) --- .../internal/handler/admin/setting_handler.go | 12 ++ backend/internal/handler/dto/settings.go | 1 + backend/internal/pkg/openai/allowed_client.go | 78 +++++++++ .../pkg/openai/allowed_client_test.go | 95 +++++++++++ backend/internal/server/api_contract_test.go | 2 + backend/internal/service/account.go | 32 ++++ ...unt_codex_cli_only_allowed_clients_test.go | 68 ++++++++ backend/internal/service/domain_constants.go | 3 + .../openai_client_restriction_detector.go | 28 +++- ...openai_client_restriction_detector_test.go | 151 +++++++++++++++++- .../service/openai_gateway_service.go | 13 +- ...nai_gateway_service_codex_cli_only_test.go | 4 +- backend/internal/service/setting_service.go | 90 +++++++++-- ...g_service_openai_allow_claude_code_test.go | 55 +++++++ backend/internal/service/settings_view.go | 1 + frontend/src/api/admin/settings.ts | 2 + .../account/BulkEditAccountModal.vue | 54 +++++++ .../components/account/CreateAccountModal.vue | 39 +++++ .../components/account/EditAccountModal.vue | 37 +++++ .../__tests__/BulkEditAccountModal.spec.ts | 19 +++ frontend/src/i18n/locales/en.ts | 6 + frontend/src/i18n/locales/zh.ts | 5 + frontend/src/views/admin/SettingsView.vue | 15 ++ 23 files changed, 787 insertions(+), 23 deletions(-) create mode 100644 backend/internal/pkg/openai/allowed_client.go create mode 100644 backend/internal/pkg/openai/allowed_client_test.go create mode 100644 backend/internal/service/account_codex_cli_only_allowed_clients_test.go create mode 100644 backend/internal/service/setting_service_openai_allow_claude_code_test.go diff --git a/backend/internal/handler/admin/setting_handler.go b/backend/internal/handler/admin/setting_handler.go index 3c7fe581..c229d340 100644 --- a/backend/internal/handler/admin/setting_handler.go +++ b/backend/internal/handler/admin/setting_handler.go @@ -256,6 +256,7 @@ func (h *SettingHandler) GetSettings(c *gin.Context) { RewriteMessageCacheControl: settings.RewriteMessageCacheControl, AntigravityUserAgentVersion: settings.AntigravityUserAgentVersion, OpenAICodexUserAgent: settings.OpenAICodexUserAgent, + OpenAIAllowClaudeCodeCodexPlugin: settings.OpenAIAllowClaudeCodeCodexPlugin, WebSearchEmulationEnabled: settings.WebSearchEmulationEnabled, PaymentVisibleMethodAlipaySource: settings.PaymentVisibleMethodAlipaySource, PaymentVisibleMethodWxpaySource: settings.PaymentVisibleMethodWxpaySource, @@ -584,6 +585,7 @@ type UpdateSettingsRequest struct { RewriteMessageCacheControl *bool `json:"rewrite_message_cache_control"` AntigravityUserAgentVersion *string `json:"antigravity_user_agent_version"` OpenAICodexUserAgent *string `json:"openai_codex_user_agent"` + OpenAIAllowClaudeCodeCodexPlugin *bool `json:"openai_allow_claude_code_codex_plugin"` // Payment visible method routing PaymentVisibleMethodAlipaySource *string `json:"payment_visible_method_alipay_source"` @@ -1655,6 +1657,12 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { } return previousSettings.OpenAICodexUserAgent }(), + OpenAIAllowClaudeCodeCodexPlugin: func() bool { + if req.OpenAIAllowClaudeCodeCodexPlugin != nil { + return *req.OpenAIAllowClaudeCodeCodexPlugin + } + return previousSettings.OpenAIAllowClaudeCodeCodexPlugin + }(), PaymentVisibleMethodAlipaySource: func() string { if req.PaymentVisibleMethodAlipaySource != nil { return strings.TrimSpace(*req.PaymentVisibleMethodAlipaySource) @@ -2031,6 +2039,7 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { RewriteMessageCacheControl: updatedSettings.RewriteMessageCacheControl, AntigravityUserAgentVersion: updatedSettings.AntigravityUserAgentVersion, OpenAICodexUserAgent: updatedSettings.OpenAICodexUserAgent, + OpenAIAllowClaudeCodeCodexPlugin: updatedSettings.OpenAIAllowClaudeCodeCodexPlugin, PaymentVisibleMethodAlipaySource: updatedSettings.PaymentVisibleMethodAlipaySource, PaymentVisibleMethodWxpaySource: updatedSettings.PaymentVisibleMethodWxpaySource, PaymentVisibleMethodAlipayEnabled: updatedSettings.PaymentVisibleMethodAlipayEnabled, @@ -2500,6 +2509,9 @@ func diffSettings(before *service.SystemSettings, after *service.SystemSettings, if before.OpenAICodexUserAgent != after.OpenAICodexUserAgent { changed = append(changed, "openai_codex_user_agent") } + if before.OpenAIAllowClaudeCodeCodexPlugin != after.OpenAIAllowClaudeCodeCodexPlugin { + changed = append(changed, "openai_allow_claude_code_codex_plugin") + } if before.PaymentVisibleMethodAlipaySource != after.PaymentVisibleMethodAlipaySource { changed = append(changed, "payment_visible_method_alipay_source") } diff --git a/backend/internal/handler/dto/settings.go b/backend/internal/handler/dto/settings.go index eecf98ac..17772a2e 100644 --- a/backend/internal/handler/dto/settings.go +++ b/backend/internal/handler/dto/settings.go @@ -185,6 +185,7 @@ type SystemSettings struct { RewriteMessageCacheControl bool `json:"rewrite_message_cache_control"` AntigravityUserAgentVersion string `json:"antigravity_user_agent_version"` OpenAICodexUserAgent string `json:"openai_codex_user_agent"` + OpenAIAllowClaudeCodeCodexPlugin bool `json:"openai_allow_claude_code_codex_plugin"` // Web Search Emulation WebSearchEmulationEnabled bool `json:"web_search_emulation_enabled"` diff --git a/backend/internal/pkg/openai/allowed_client.go b/backend/internal/pkg/openai/allowed_client.go new file mode 100644 index 00000000..d4ca14ee --- /dev/null +++ b/backend/internal/pkg/openai/allowed_client.go @@ -0,0 +1,78 @@ +package openai + +import "strings" + +// 命名预设 ID。账号侧 codex_cli_only_allowed_clients 只能引用这些预设键, +// 具体匹配规则固化在下方 registry 中,配置只能「选择启用哪些预设」、不能自定义规则, +// 以防该白名单退化为可任意放宽的后门。 +const ( + // AllowedClientClaudeCode 对应 Claude Code CLI 的 codex 插件。 + AllowedClientClaudeCode = "claude_code" +) + +// AllowedClientEntry 描述一个被额外放行的非官方 Codex 客户端签名。 +// Originator 必须精确等值匹配(归一化后)。 +// UAContains 为必填字段:列表为空,或列表中存在任何空白 marker,均视为非法配置, +// 整体安全失败(return false);每一项都必须出现在 User-Agent 中。 +// 这确保双因子匹配不会因缺失 UA 声明而退化为仅凭可伪造的 originator 单因子放行。 +type AllowedClientEntry struct { + Originator string + UAContains []string +} + +// allowedClientRegistry 固化各命名预设的签名规则。 +// +// Claude Code codex 插件签名来源:插件以 clientInfo.name="Claude Code" 完成 app-server +// initialize 握手,codex 据此把 originator 设为 "Claude Code",User-Agent 前缀同样为 +// "Claude Code/"(两者同源)。若上游 Claude Code 插件更改 clientInfo.name,此处需同步更新。 +var allowedClientRegistry = map[string]AllowedClientEntry{ + AllowedClientClaudeCode: { + Originator: "Claude Code", + UAContains: []string{"Claude Code/"}, + }, +} + +// IsAllowedClientMatch 判断请求头是否命中给定的额外客户端签名。 +// originator 必须精确等值(归一化后);UAContains 中每一项都必须出现在 UA 中。 +// UAContains 为必填:列表为空或含任何空白 marker 均视为非法配置,整体安全失败。 +func IsAllowedClientMatch(userAgent, originator string, entry AllowedClientEntry) bool { + wantOriginator := normalizeCodexClientHeader(entry.Originator) + if wantOriginator == "" { + return false + } + if normalizeCodexClientHeader(originator) != wantOriginator { + return false + } + // 预设必须声明 UA 特征:否则将退化为仅凭可伪造的 originator 单因子匹配。 + if len(entry.UAContains) == 0 { + return false + } + ua := normalizeCodexClientHeader(userAgent) + for _, marker := range entry.UAContains { + normalizedMarker := normalizeCodexClientHeader(marker) + if normalizedMarker == "" { + // 空白 marker 让该项失去校验能力,会让双因子退化为仅 originator + // 单因子;视为非法配置,安全失败。 + return false + } + if !strings.Contains(ua, normalizedMarker) { + return false + } + } + return true +} + +// MatchAllowedClients 判断请求头是否命中 clientIDs 引用的任一预设签名。 +// 未知预设 ID 会被忽略;空列表恒不放行(默认拒绝)。 +func MatchAllowedClients(userAgent, originator string, clientIDs []string) bool { + for _, id := range clientIDs { + entry, ok := allowedClientRegistry[normalizeCodexClientHeader(id)] + if !ok { + continue + } + if IsAllowedClientMatch(userAgent, originator, entry) { + return true + } + } + return false +} diff --git a/backend/internal/pkg/openai/allowed_client_test.go b/backend/internal/pkg/openai/allowed_client_test.go new file mode 100644 index 00000000..c42aa4d5 --- /dev/null +++ b/backend/internal/pkg/openai/allowed_client_test.go @@ -0,0 +1,95 @@ +package openai + +import "testing" + +// 真实的 Claude Code codex 插件请求头:originator 与 UA 前缀同源于 clientInfo.name="Claude Code"。 +const ( + testClaudeCodeOriginator = "Claude Code" + testClaudeCodeUserAgent = "Claude Code/0.5.0 (Macos 15.5; arm64) iTerm2.app (Claude Code; 1.0.4)" +) + +func TestIsAllowedClientMatch(t *testing.T) { + entry := AllowedClientEntry{Originator: "Claude Code", UAContains: []string{"Claude Code/"}} + + tests := []struct { + name string + ua string + originator string + want bool + }{ + {name: "真实签名命中", ua: testClaudeCodeUserAgent, originator: testClaudeCodeOriginator, want: true}, + {name: "大小写不敏感", ua: "claude code/0.5.0 (macos)", originator: "claude code", want: true}, + {name: "originator 两侧空白被裁剪", ua: testClaudeCodeUserAgent, originator: " Claude Code ", want: true}, + {name: "originator 非精确(带后缀)不命中", ua: testClaudeCodeUserAgent, originator: "Claude Code Extra", want: false}, + {name: "originator 为空不命中", ua: testClaudeCodeUserAgent, originator: "", want: false}, + {name: "originator 是官方 codex 不命中", ua: testClaudeCodeUserAgent, originator: "codex_cli_rs", want: false}, + {name: "UA 缺少 Claude Code/ 标记不命中", ua: "curl/8.0", originator: testClaudeCodeOriginator, want: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := IsAllowedClientMatch(tt.ua, tt.originator, entry); got != tt.want { + t.Fatalf("IsAllowedClientMatch(%q, %q) = %v, want %v", tt.ua, tt.originator, got, tt.want) + } + }) + } +} + +func TestIsAllowedClientMatch_EmptyOriginatorEntryNeverMatches(t *testing.T) { + // registry 条目若没有配置 Originator,绝不放行,避免成为宽松后门。 + entry := AllowedClientEntry{Originator: "", UAContains: []string{"Claude Code/"}} + if IsAllowedClientMatch(testClaudeCodeUserAgent, "", entry) { + t.Fatal("空 Originator 的条目不应匹配任何请求") + } +} + +func TestIsAllowedClientMatch_EmptyUAContainsNeverMatches(t *testing.T) { + // 预设必须声明 UA 特征,否则退化为仅凭可伪造的 originator 单因子匹配,绝不放行。 + entry := AllowedClientEntry{Originator: "Claude Code", UAContains: nil} + if IsAllowedClientMatch(testClaudeCodeUserAgent, testClaudeCodeOriginator, entry) { + t.Fatal("未声明 UA 特征的预设不应匹配,避免退化为单因子 originator 匹配") + } +} + +func TestIsAllowedClientMatch_WhitespaceUAMarkerNeverMatches(t *testing.T) { + // 全空白 marker 归一化后为空,若被跳过则退化为仅 originator 单因子; + // 任何空白 marker 视为非法预设配置,必须安全失败。 + entry := AllowedClientEntry{Originator: "Claude Code", UAContains: []string{" "}} + if IsAllowedClientMatch(testClaudeCodeUserAgent, testClaudeCodeOriginator, entry) { + t.Fatal("UAContains 含全空白 marker 不应匹配,避免退化为单因子 originator 匹配") + } +} + +func TestIsAllowedClientMatch_MixedEmptyUAMarkerNeverMatches(t *testing.T) { + // 即便 UAContains 含一个真实 marker,只要其中混入任何空白 marker 也视为非法配置; + // 防止维护者只为对齐凑数而插入空字符串。 + entry := AllowedClientEntry{Originator: "Claude Code", UAContains: []string{"", "Claude Code/"}} + if IsAllowedClientMatch(testClaudeCodeUserAgent, testClaudeCodeOriginator, entry) { + t.Fatal("UAContains 混入空白 marker 不应匹配") + } +} + +func TestMatchAllowedClients(t *testing.T) { + tests := []struct { + name string + ua string + originator string + clientIDs []string + want bool + }{ + {name: "claude_code 预设命中真实签名", ua: testClaudeCodeUserAgent, originator: testClaudeCodeOriginator, clientIDs: []string{AllowedClientClaudeCode}, want: true}, + {name: "claude_code 预设 + 伪造 originator 不命中", ua: testClaudeCodeUserAgent, originator: "my_client", clientIDs: []string{AllowedClientClaudeCode}, want: false}, + {name: "空列表不放行", ua: testClaudeCodeUserAgent, originator: testClaudeCodeOriginator, clientIDs: nil, want: false}, + {name: "未知预设 ID 不放行", ua: testClaudeCodeUserAgent, originator: testClaudeCodeOriginator, clientIDs: []string{"unknown_client"}, want: false}, + {name: "ID 大小写/空白容错", ua: testClaudeCodeUserAgent, originator: testClaudeCodeOriginator, clientIDs: []string{" Claude_Code "}, want: true}, + {name: "多预设任一命中即放行", ua: testClaudeCodeUserAgent, originator: testClaudeCodeOriginator, clientIDs: []string{"unknown_client", AllowedClientClaudeCode}, want: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := MatchAllowedClients(tt.ua, tt.originator, tt.clientIDs); got != tt.want { + t.Fatalf("MatchAllowedClients(%q, %q, %v) = %v, want %v", tt.ua, tt.originator, tt.clientIDs, got, tt.want) + } + }) + } +} diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index 662daed1..9eea0924 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -843,6 +843,7 @@ func TestAPIContracts(t *testing.T) { "payment_visible_method_wxpay_enabled": false, "openai_advanced_scheduler_enabled": true, "openai_codex_user_agent": "", + "openai_allow_claude_code_codex_plugin": false, "openai_fast_policy_settings": { "rules": [] }, @@ -1079,6 +1080,7 @@ func TestAPIContracts(t *testing.T) { "payment_visible_method_wxpay_enabled": false, "openai_advanced_scheduler_enabled": false, "openai_codex_user_agent": "", + "openai_allow_claude_code_codex_plugin": false, "openai_fast_policy_settings": { "rules": [] }, diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go index d488aa75..f51f0325 100644 --- a/backend/internal/service/account.go +++ b/backend/internal/service/account.go @@ -1442,6 +1442,38 @@ func (a *Account) IsCodexCLIOnlyEnabled() bool { return ok && enabled } +// GetCodexCLIOnlyAllowedClients 返回 codex_cli_only 之上额外放行的命名客户端预设 ID 列表。 +// 仅 OpenAI OAuth 账号生效;缺失或类型不符时返回空。预设 ID 的具体匹配规则由 +// openai 包的 registry 固化,配置只能引用预设键、不能自定义规则。 +func (a *Account) GetCodexCLIOnlyAllowedClients() []string { + if a == nil || !a.IsOpenAIOAuth() || a.Extra == nil { + return nil + } + raw, ok := a.Extra["codex_cli_only_allowed_clients"] + if !ok || raw == nil { + return nil + } + switch v := raw.(type) { + case []string: + result := make([]string, 0, len(v)) + for _, s := range v { + if strings.TrimSpace(s) != "" { + result = append(result, s) + } + } + return result + case []any: + result := make([]string, 0, len(v)) + for _, item := range v { + if s, ok := item.(string); ok && strings.TrimSpace(s) != "" { + result = append(result, s) + } + } + return result + } + return nil +} + // WindowCostSchedulability 窗口费用调度状态 type WindowCostSchedulability int diff --git a/backend/internal/service/account_codex_cli_only_allowed_clients_test.go b/backend/internal/service/account_codex_cli_only_allowed_clients_test.go new file mode 100644 index 00000000..c835ea27 --- /dev/null +++ b/backend/internal/service/account_codex_cli_only_allowed_clients_test.go @@ -0,0 +1,68 @@ +package service + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestAccount_GetCodexCLIOnlyAllowedClients(t *testing.T) { + t.Run("OAuth 账号读取 []any 字符串列表", func(t *testing.T) { + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{"codex_cli_only_allowed_clients": []any{"claude_code"}}, + } + require.Equal(t, []string{"claude_code"}, account.GetCodexCLIOnlyAllowedClients()) + }) + + t.Run("OAuth 账号读取 []string 列表", func(t *testing.T) { + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{"codex_cli_only_allowed_clients": []string{"claude_code"}}, + } + require.Equal(t, []string{"claude_code"}, account.GetCodexCLIOnlyAllowedClients()) + }) + + t.Run("[]string 跳过空白元素", func(t *testing.T) { + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{"codex_cli_only_allowed_clients": []string{"claude_code", "", " "}}, + } + require.Equal(t, []string{"claude_code"}, account.GetCodexCLIOnlyAllowedClients()) + }) + + t.Run("跳过非字符串与空白元素", func(t *testing.T) { + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{"codex_cli_only_allowed_clients": []any{"claude_code", 123, "", " "}}, + } + require.Equal(t, []string{"claude_code"}, account.GetCodexCLIOnlyAllowedClients()) + }) + + t.Run("非 OAuth 账号返回空", func(t *testing.T) { + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Extra: map[string]any{"codex_cli_only_allowed_clients": []any{"claude_code"}}, + } + require.Empty(t, account.GetCodexCLIOnlyAllowedClients()) + }) + + t.Run("Extra 为空返回空", func(t *testing.T) { + account := &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth} + require.Empty(t, account.GetCodexCLIOnlyAllowedClients()) + }) + + t.Run("字段缺失返回空", func(t *testing.T) { + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{}, + } + require.Empty(t, account.GetCodexCLIOnlyAllowedClients()) + }) +} diff --git a/backend/internal/service/domain_constants.go b/backend/internal/service/domain_constants.go index 59c34eaa..b6441238 100644 --- a/backend/internal/service/domain_constants.go +++ b/backend/internal/service/domain_constants.go @@ -431,6 +431,9 @@ const ( // 当客户端 UA 被识别为浏览器(Chrome/Firefox/Safari/Edge 等)时,转发给 OpenAI 上游前会替换为此值, // 用于避免 Cloudflare 对浏览器型 UA 的质询拦截。 SettingKeyOpenAICodexUserAgent = "openai_codex_user_agent" + // SettingKeyOpenAIAllowClaudeCodeCodexPlugin 全局开关:是否额外放行 Claude Code 的 Codex 插件(默认 false)。 + // 仅在账号 codex_cli_only 开启时生效;开启后无需逐账号配置 codex_cli_only_allowed_clients。 + SettingKeyOpenAIAllowClaudeCodeCodexPlugin = "openai_allow_claude_code_codex_plugin" // 余额不足提醒 SettingKeyBalanceLowNotifyEnabled = "balance_low_notify_enabled" // 全局开关 diff --git a/backend/internal/service/openai_client_restriction_detector.go b/backend/internal/service/openai_client_restriction_detector.go index d1784e11..8589737a 100644 --- a/backend/internal/service/openai_client_restriction_detector.go +++ b/backend/internal/service/openai_client_restriction_detector.go @@ -13,6 +13,10 @@ const ( CodexClientRestrictionReasonMatchedUA = "official_client_user_agent_matched" // CodexClientRestrictionReasonMatchedOriginator 表示请求命中官方客户端 originator 白名单。 CodexClientRestrictionReasonMatchedOriginator = "official_client_originator_matched" + // CodexClientRestrictionReasonMatchedAllowedClient 表示请求命中账号级额外放行的命名客户端预设。 + CodexClientRestrictionReasonMatchedAllowedClient = "allowed_client_matched" + // CodexClientRestrictionReasonMatchedGlobalAllowedClient 表示请求命中全局额外放行的命名客户端预设。 + CodexClientRestrictionReasonMatchedGlobalAllowedClient = "global_allowed_client_matched" // CodexClientRestrictionReasonNotMatchedUA 表示请求未命中官方客户端 UA 白名单。 CodexClientRestrictionReasonNotMatchedUA = "official_client_user_agent_not_matched" // CodexClientRestrictionReasonForceCodexCLI 表示通过 ForceCodexCLI 配置兜底放行。 @@ -28,7 +32,7 @@ type CodexClientRestrictionDetectionResult struct { // CodexClientRestrictionDetector 定义 codex_cli_only 统一检测入口。 type CodexClientRestrictionDetector interface { - Detect(c *gin.Context, account *Account) CodexClientRestrictionDetectionResult + Detect(c *gin.Context, account *Account, globalAllowedClients []string) CodexClientRestrictionDetectionResult } // OpenAICodexClientRestrictionDetector 为 OpenAI OAuth codex_cli_only 的默认实现。 @@ -40,7 +44,7 @@ func NewOpenAICodexClientRestrictionDetector(cfg *config.Config) *OpenAICodexCli return &OpenAICodexClientRestrictionDetector{cfg: cfg} } -func (d *OpenAICodexClientRestrictionDetector) Detect(c *gin.Context, account *Account) CodexClientRestrictionDetectionResult { +func (d *OpenAICodexClientRestrictionDetector) Detect(c *gin.Context, account *Account, globalAllowedClients []string) CodexClientRestrictionDetectionResult { if account == nil || !account.IsCodexCLIOnlyEnabled() { return CodexClientRestrictionDetectionResult{ Enabled: false, @@ -78,6 +82,26 @@ func (d *OpenAICodexClientRestrictionDetector) Detect(c *gin.Context, account *A } } + // 官方客户端白名单未命中时,先尝试账号级额外放行的命名客户端预设(如 Claude Code codex 插件)。 + if allowed := account.GetCodexCLIOnlyAllowedClients(); len(allowed) > 0 && + openai.MatchAllowedClients(userAgent, originator, allowed) { + return CodexClientRestrictionDetectionResult{ + Enabled: true, + Matched: true, + Reason: CodexClientRestrictionReasonMatchedAllowedClient, + } + } + + // 再尝试由更高作用域(全局设置)注入的额外放行客户端列表。 + if len(globalAllowedClients) > 0 && + openai.MatchAllowedClients(userAgent, originator, globalAllowedClients) { + return CodexClientRestrictionDetectionResult{ + Enabled: true, + Matched: true, + Reason: CodexClientRestrictionReasonMatchedGlobalAllowedClient, + } + } + return CodexClientRestrictionDetectionResult{ Enabled: true, Matched: false, diff --git a/backend/internal/service/openai_client_restriction_detector_test.go b/backend/internal/service/openai_client_restriction_detector_test.go index 984b4ff6..fc115128 100644 --- a/backend/internal/service/openai_client_restriction_detector_test.go +++ b/backend/internal/service/openai_client_restriction_detector_test.go @@ -30,7 +30,7 @@ func TestOpenAICodexClientRestrictionDetector_Detect(t *testing.T) { detector := NewOpenAICodexClientRestrictionDetector(nil) account := &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth, Extra: map[string]any{}} - result := detector.Detect(newCodexDetectorTestContext("curl/8.0", ""), account) + result := detector.Detect(newCodexDetectorTestContext("curl/8.0", ""), account, nil) require.False(t, result.Enabled) require.False(t, result.Matched) require.Equal(t, CodexClientRestrictionReasonDisabled, result.Reason) @@ -44,7 +44,7 @@ func TestOpenAICodexClientRestrictionDetector_Detect(t *testing.T) { Extra: map[string]any{"codex_cli_only": true}, } - result := detector.Detect(newCodexDetectorTestContext("codex_cli_rs/0.99.0", ""), account) + result := detector.Detect(newCodexDetectorTestContext("codex_cli_rs/0.99.0", ""), account, nil) require.True(t, result.Enabled) require.True(t, result.Matched) require.Equal(t, CodexClientRestrictionReasonMatchedUA, result.Reason) @@ -58,7 +58,7 @@ func TestOpenAICodexClientRestrictionDetector_Detect(t *testing.T) { Extra: map[string]any{"codex_cli_only": true}, } - result := detector.Detect(newCodexDetectorTestContext("codex_vscode/1.0.0", ""), account) + result := detector.Detect(newCodexDetectorTestContext("codex_vscode/1.0.0", ""), account, nil) require.True(t, result.Enabled) require.True(t, result.Matched) require.Equal(t, CodexClientRestrictionReasonMatchedUA, result.Reason) @@ -72,7 +72,7 @@ func TestOpenAICodexClientRestrictionDetector_Detect(t *testing.T) { Extra: map[string]any{"codex_cli_only": true}, } - result := detector.Detect(newCodexDetectorTestContext("codex_app/2.1.0", ""), account) + result := detector.Detect(newCodexDetectorTestContext("codex_app/2.1.0", ""), account, nil) require.True(t, result.Enabled) require.True(t, result.Matched) require.Equal(t, CodexClientRestrictionReasonMatchedUA, result.Reason) @@ -86,7 +86,7 @@ func TestOpenAICodexClientRestrictionDetector_Detect(t *testing.T) { Extra: map[string]any{"codex_cli_only": true}, } - result := detector.Detect(newCodexDetectorTestContext("curl/8.0", "codex_chatgpt_desktop"), account) + result := detector.Detect(newCodexDetectorTestContext("curl/8.0", "codex_chatgpt_desktop"), account, nil) require.True(t, result.Enabled) require.True(t, result.Matched) require.Equal(t, CodexClientRestrictionReasonMatchedOriginator, result.Reason) @@ -100,7 +100,7 @@ func TestOpenAICodexClientRestrictionDetector_Detect(t *testing.T) { Extra: map[string]any{"codex_cli_only": true}, } - result := detector.Detect(newCodexDetectorTestContext("curl/8.0", "my_client"), account) + result := detector.Detect(newCodexDetectorTestContext("curl/8.0", "my_client"), account, nil) require.True(t, result.Enabled) require.False(t, result.Matched) require.Equal(t, CodexClientRestrictionReasonNotMatchedUA, result.Reason) @@ -116,9 +116,146 @@ func TestOpenAICodexClientRestrictionDetector_Detect(t *testing.T) { Extra: map[string]any{"codex_cli_only": true}, } - result := detector.Detect(newCodexDetectorTestContext("curl/8.0", "my_client"), account) + result := detector.Detect(newCodexDetectorTestContext("curl/8.0", "my_client"), account, nil) require.True(t, result.Enabled) require.True(t, result.Matched) require.Equal(t, CodexClientRestrictionReasonForceCodexCLI, result.Reason) }) } + +func TestOpenAICodexClientRestrictionDetector_Detect_AllowedClients(t *testing.T) { + gin.SetMode(gin.TestMode) + + const ( + claudeCodeUA = "Claude Code/0.5.0 (Macos 15.5; arm64) iTerm2.app (Claude Code; 1.0.4)" + claudeCodeOriginator = "Claude Code" + ) + + t.Run("配置 claude_code 白名单且命中真实签名时放行", func(t *testing.T) { + detector := NewOpenAICodexClientRestrictionDetector(nil) + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{ + "codex_cli_only": true, + "codex_cli_only_allowed_clients": []any{"claude_code"}, + }, + } + + result := detector.Detect(newCodexDetectorTestContext(claudeCodeUA, claudeCodeOriginator), account, nil) + require.True(t, result.Enabled) + require.True(t, result.Matched) + require.Equal(t, CodexClientRestrictionReasonMatchedAllowedClient, result.Reason) + }) + + t.Run("配置白名单但伪造 originator 仍拒绝", func(t *testing.T) { + detector := NewOpenAICodexClientRestrictionDetector(nil) + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{ + "codex_cli_only": true, + "codex_cli_only_allowed_clients": []any{"claude_code"}, + }, + } + + result := detector.Detect(newCodexDetectorTestContext(claudeCodeUA, "my_client"), account, nil) + require.True(t, result.Enabled) + require.False(t, result.Matched) + require.Equal(t, CodexClientRestrictionReasonNotMatchedUA, result.Reason) + }) + + t.Run("未配置白名单时 Claude Code 签名仍拒绝", func(t *testing.T) { + detector := NewOpenAICodexClientRestrictionDetector(nil) + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{"codex_cli_only": true}, + } + + result := detector.Detect(newCodexDetectorTestContext(claudeCodeUA, claudeCodeOriginator), account, nil) + require.True(t, result.Enabled) + require.False(t, result.Matched) + require.Equal(t, CodexClientRestrictionReasonNotMatchedUA, result.Reason) + }) + + t.Run("未开启 codex_cli_only 时白名单不参与,直接绕过", func(t *testing.T) { + detector := NewOpenAICodexClientRestrictionDetector(nil) + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{"codex_cli_only_allowed_clients": []any{"claude_code"}}, + } + + result := detector.Detect(newCodexDetectorTestContext(claudeCodeUA, claudeCodeOriginator), account, nil) + require.False(t, result.Enabled) + require.False(t, result.Matched) + require.Equal(t, CodexClientRestrictionReasonDisabled, result.Reason) + }) + + t.Run("全局列表含 claude_code + 命中签名 → 放行(global)", func(t *testing.T) { + detector := NewOpenAICodexClientRestrictionDetector(nil) + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{"codex_cli_only": true}, + } + result := detector.Detect( + newCodexDetectorTestContext("Claude Code/0.5.0 (Macos 15.5; arm64) iTerm2.app (Claude Code; 1.0.4)", "Claude Code"), + account, + []string{"claude_code"}, + ) + require.True(t, result.Enabled) + require.True(t, result.Matched) + require.Equal(t, CodexClientRestrictionReasonMatchedGlobalAllowedClient, result.Reason) + }) + + t.Run("全局列表含 claude_code + 非签名 → 403", func(t *testing.T) { + detector := NewOpenAICodexClientRestrictionDetector(nil) + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{"codex_cli_only": true}, + } + result := detector.Detect(newCodexDetectorTestContext("curl/8.0", "my_client"), account, []string{"claude_code"}) + require.True(t, result.Enabled) + require.False(t, result.Matched) + require.Equal(t, CodexClientRestrictionReasonNotMatchedUA, result.Reason) + }) + + t.Run("全局列表为空 + 账号未配 → 403", func(t *testing.T) { + detector := NewOpenAICodexClientRestrictionDetector(nil) + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{"codex_cli_only": true}, + } + result := detector.Detect( + newCodexDetectorTestContext("Claude Code/0.5.0 (Macos) (Claude Code; 1.0.4)", "Claude Code"), + account, + nil, + ) + require.True(t, result.Enabled) + require.False(t, result.Matched) + require.Equal(t, CodexClientRestrictionReasonNotMatchedUA, result.Reason) + }) + + t.Run("账号白名单优先于全局列表(reason=account)", func(t *testing.T) { + detector := NewOpenAICodexClientRestrictionDetector(nil) + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{ + "codex_cli_only": true, + "codex_cli_only_allowed_clients": []any{"claude_code"}, + }, + } + result := detector.Detect( + newCodexDetectorTestContext("Claude Code/0.5.0 (Macos) (Claude Code; 1.0.4)", "Claude Code"), + account, + []string{"claude_code"}, + ) + require.True(t, result.Matched) + require.Equal(t, CodexClientRestrictionReasonMatchedAllowedClient, result.Reason) + }) +} diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index f93cc221..997423b7 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -901,7 +901,17 @@ func SnapshotOpenAICompatibilityFallbackMetrics() OpenAICompatibilityFallbackMet } func (s *OpenAIGatewayService) detectCodexClientRestriction(c *gin.Context, account *Account) CodexClientRestrictionDetectionResult { - return s.getCodexClientRestrictionDetector().Detect(c, account) + var globalAllowedClients []string + if account != nil && account.IsCodexCLIOnlyEnabled() && s != nil && s.settingService != nil { + ctx := context.Background() + if c != nil && c.Request != nil { + ctx = c.Request.Context() + } + if s.settingService.IsOpenAIAllowClaudeCodeCodexPluginEnabled(ctx) { + globalAllowedClients = []string{openai.AllowedClientClaudeCode} + } + } + return s.getCodexClientRestrictionDetector().Detect(c, account, globalAllowedClients) } func getAPIKeyIDFromContext(c *gin.Context) int64 { @@ -959,6 +969,7 @@ func logCodexCLIOnlyDetection(ctx context.Context, c *gin.Context, account *Acco } log := logger.FromContext(ctx).With(fields...) if result.Matched { + log.Info("OpenAI codex_cli_only 放行请求") return } log.Warn("OpenAI codex_cli_only 拒绝非官方客户端请求") diff --git a/backend/internal/service/openai_gateway_service_codex_cli_only_test.go b/backend/internal/service/openai_gateway_service_codex_cli_only_test.go index 17a874ea..10d58654 100644 --- a/backend/internal/service/openai_gateway_service_codex_cli_only_test.go +++ b/backend/internal/service/openai_gateway_service_codex_cli_only_test.go @@ -18,7 +18,7 @@ type stubCodexRestrictionDetector struct { result CodexClientRestrictionDetectionResult } -func (s *stubCodexRestrictionDetector) Detect(_ *gin.Context, _ *Account) CodexClientRestrictionDetectionResult { +func (s *stubCodexRestrictionDetector) Detect(_ *gin.Context, _ *Account, _ []string) CodexClientRestrictionDetectionResult { return s.result } @@ -52,7 +52,7 @@ func TestOpenAIGatewayService_GetCodexClientRestrictionDetector(t *testing.T) { c.Request.Header.Set("User-Agent", "curl/8.0") account := &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth, Extra: map[string]any{"codex_cli_only": true}} - result := got.Detect(c, account) + result := got.Detect(c, account, nil) require.True(t, result.Enabled) require.True(t, result.Matched) require.Equal(t, CodexClientRestrictionReasonForceCodexCLI, result.Reason) diff --git a/backend/internal/service/setting_service.go b/backend/internal/service/setting_service.go index e6f0f2bc..08c0d045 100644 --- a/backend/internal/service/setting_service.go +++ b/backend/internal/service/setting_service.go @@ -141,6 +141,17 @@ const openAICodexUserAgentCacheTTL = 60 * time.Second const openAICodexUserAgentErrorTTL = 5 * time.Second const openAICodexUserAgentDBTimeout = 5 * time.Second +// cachedOpenAIAllowCodexPlugin Codex 插件放行开关缓存(进程内缓存,60s TTL)。 +// IsOpenAIAllowClaudeCodeCodexPluginEnabled 在每个 codex_cli_only 账号的网关请求热路径上被调用,避免每次访问 DB。 +type cachedOpenAIAllowCodexPlugin struct { + value bool + expiresAt int64 // unix nano +} + +const openAIAllowCodexPluginCacheTTL = 60 * time.Second +const openAIAllowCodexPluginErrorTTL = 5 * time.Second +const openAIAllowCodexPluginDBTimeout = 5 * time.Second + // DefaultSubscriptionGroupReader validates group references used by default subscriptions. type DefaultSubscriptionGroupReader interface { GetByID(ctx context.Context, id int64) (*Group, error) @@ -152,17 +163,19 @@ type WebSearchManagerBuilder func(cfg *WebSearchEmulationConfig, proxyURLs map[i // SettingService 系统设置服务 type SettingService struct { - settingRepo SettingRepository - defaultSubGroupReader DefaultSubscriptionGroupReader - proxyRepo ProxyRepository // for resolving websearch provider proxy URLs - cfg *config.Config - onUpdate func() // Callback when settings are updated (for cache invalidation) - version string // Application version - webSearchManagerBuilder WebSearchManagerBuilder - antigravityUAVersionCache atomic.Value // *cachedAntigravityUserAgentVersion - antigravityUAVersionSF singleflight.Group - openAICodexUACache atomic.Value // *cachedOpenAICodexUserAgent - openAICodexUASF singleflight.Group + settingRepo SettingRepository + defaultSubGroupReader DefaultSubscriptionGroupReader + proxyRepo ProxyRepository // for resolving websearch provider proxy URLs + cfg *config.Config + onUpdate func() // Callback when settings are updated (for cache invalidation) + version string // Application version + webSearchManagerBuilder WebSearchManagerBuilder + antigravityUAVersionCache atomic.Value // *cachedAntigravityUserAgentVersion + antigravityUAVersionSF singleflight.Group + openAICodexUACache atomic.Value // *cachedOpenAICodexUserAgent + openAICodexUASF singleflight.Group + openAIAllowCodexPluginCache atomic.Value // *cachedOpenAIAllowCodexPlugin + openAIAllowCodexPluginSF singleflight.Group } // DefaultPlatformQuotaSetting 单 platform 三档限额(nil = 沿用上层;0 = 显式禁用;>0 = 上限) @@ -1015,6 +1028,54 @@ func (s *SettingService) GetOpenAICodexUserAgent(ctx context.Context) string { return fallback } +// IsOpenAIAllowClaudeCodeCodexPluginEnabled 全局开关:是否额外放行 Claude Code 的 Codex 插件(默认关闭)。 +// 仅在调用方已确认账号 codex_cli_only 开启时读取,避免对非受限账号产生无谓查询。 +// 使用进程内 atomic.Value 缓存(60s TTL),避免在每个网关请求热路径上访问 DB。 +func (s *SettingService) IsOpenAIAllowClaudeCodeCodexPluginEnabled(ctx context.Context) bool { + if cached, ok := s.openAIAllowCodexPluginCache.Load().(*cachedOpenAIAllowCodexPlugin); ok && cached != nil { + if time.Now().UnixNano() < cached.expiresAt { + return cached.value + } + } + result, _, _ := s.openAIAllowCodexPluginSF.Do("openai_allow_codex_plugin_enabled", func() (any, error) { + if cached, ok := s.openAIAllowCodexPluginCache.Load().(*cachedOpenAIAllowCodexPlugin); ok && cached != nil { + if time.Now().UnixNano() < cached.expiresAt { + return cached.value, nil + } + } + dbCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), openAIAllowCodexPluginDBTimeout) + defer cancel() + value, err := s.settingRepo.GetValue(dbCtx, SettingKeyOpenAIAllowClaudeCodeCodexPlugin) + if err != nil { + if errors.Is(err, ErrSettingNotFound) { + // 设置不存在 → 默认关闭,正常 TTL 缓存 + s.openAIAllowCodexPluginCache.Store(&cachedOpenAIAllowCodexPlugin{ + value: false, + expiresAt: time.Now().Add(openAIAllowCodexPluginCacheTTL).UnixNano(), + }) + return false, nil + } + slog.Warn("failed to get openai_allow_claude_code_codex_plugin setting", "error", err) + // DB 错误 → 安全默认关闭,短 TTL 快速重试 + s.openAIAllowCodexPluginCache.Store(&cachedOpenAIAllowCodexPlugin{ + value: false, + expiresAt: time.Now().Add(openAIAllowCodexPluginErrorTTL).UnixNano(), + }) + return false, nil + } + enabled := value == "true" + s.openAIAllowCodexPluginCache.Store(&cachedOpenAIAllowCodexPlugin{ + value: enabled, + expiresAt: time.Now().Add(openAIAllowCodexPluginCacheTTL).UnixNano(), + }) + return enabled, nil + }) + if val, ok := result.(bool); ok { + return val + } + return false +} + // SetOnUpdateCallback sets a callback function to be called when settings are updated // This is used for cache invalidation (e.g., HTML cache in frontend server) func (s *SettingService) SetOnUpdateCallback(callback func()) { @@ -1816,6 +1877,7 @@ func (s *SettingService) buildSystemSettingsUpdates(ctx context.Context, setting updates[SettingKeyRewriteMessageCacheControl] = strconv.FormatBool(settings.RewriteMessageCacheControl) updates[SettingKeyAntigravityUserAgentVersion] = antigravity.NormalizeUserAgentVersion(settings.AntigravityUserAgentVersion) updates[SettingKeyOpenAICodexUserAgent] = strings.TrimSpace(settings.OpenAICodexUserAgent) + updates[SettingKeyOpenAIAllowClaudeCodeCodexPlugin] = strconv.FormatBool(settings.OpenAIAllowClaudeCodeCodexPlugin) updates[SettingPaymentVisibleMethodAlipaySource] = settings.PaymentVisibleMethodAlipaySource updates[SettingPaymentVisibleMethodWxpaySource] = settings.PaymentVisibleMethodWxpaySource updates[SettingPaymentVisibleMethodAlipayEnabled] = strconv.FormatBool(settings.PaymentVisibleMethodAlipayEnabled) @@ -1968,6 +2030,11 @@ func (s *SettingService) refreshCachedSettings(settings *SystemSettings) { if s.cfg != nil { s.cfg.SetTrustForwardedIPForAPIKeyACL(settings.APIKeyACLTrustForwardedIP) } + s.openAIAllowCodexPluginSF.Forget("openai_allow_codex_plugin_enabled") + s.openAIAllowCodexPluginCache.Store(&cachedOpenAIAllowCodexPlugin{ + value: settings.OpenAIAllowClaudeCodeCodexPlugin, + expiresAt: time.Now().Add(openAIAllowCodexPluginCacheTTL).UnixNano(), + }) if s.onUpdate != nil { s.onUpdate() // Invalidate cache after settings update } @@ -3233,6 +3300,7 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin } result.AntigravityUserAgentVersion = antigravity.NormalizeUserAgentVersion(settings[SettingKeyAntigravityUserAgentVersion]) result.OpenAICodexUserAgent = strings.TrimSpace(settings[SettingKeyOpenAICodexUserAgent]) + result.OpenAIAllowClaudeCodeCodexPlugin = settings[SettingKeyOpenAIAllowClaudeCodeCodexPlugin] == "true" // Web search emulation: quick enabled check from the JSON config if raw := settings[SettingKeyWebSearchEmulationConfig]; raw != "" { diff --git a/backend/internal/service/setting_service_openai_allow_claude_code_test.go b/backend/internal/service/setting_service_openai_allow_claude_code_test.go new file mode 100644 index 00000000..22059f07 --- /dev/null +++ b/backend/internal/service/setting_service_openai_allow_claude_code_test.go @@ -0,0 +1,55 @@ +package service + +import ( + "context" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/stretchr/testify/require" +) + +type allowClaudeCodeSettingRepoStub struct{ values map[string]string } + +func (s *allowClaudeCodeSettingRepoStub) Get(ctx context.Context, key string) (*Setting, error) { + panic("unused") +} +func (s *allowClaudeCodeSettingRepoStub) GetValue(ctx context.Context, key string) (string, error) { + if v, ok := s.values[key]; ok { + return v, nil + } + return "", ErrSettingNotFound +} +func (s *allowClaudeCodeSettingRepoStub) Set(ctx context.Context, key, value string) error { + panic("unused") +} +func (s *allowClaudeCodeSettingRepoStub) GetMultiple(ctx context.Context, keys []string) (map[string]string, error) { + panic("unused") +} +func (s *allowClaudeCodeSettingRepoStub) SetMultiple(ctx context.Context, settings map[string]string) error { + panic("unused") +} +func (s *allowClaudeCodeSettingRepoStub) GetAll(ctx context.Context) (map[string]string, error) { + panic("unused") +} +func (s *allowClaudeCodeSettingRepoStub) Delete(ctx context.Context, key string) error { + panic("unused") +} + +func TestSettingService_IsOpenAIAllowClaudeCodeCodexPluginEnabled(t *testing.T) { + t.Run("默认关闭(设置缺失)", func(t *testing.T) { + svc := NewSettingService(&allowClaudeCodeSettingRepoStub{values: map[string]string{}}, &config.Config{}) + require.False(t, svc.IsOpenAIAllowClaudeCodeCodexPluginEnabled(context.Background())) + }) + t.Run("值为 true 时开启", func(t *testing.T) { + svc := NewSettingService(&allowClaudeCodeSettingRepoStub{values: map[string]string{ + SettingKeyOpenAIAllowClaudeCodeCodexPlugin: "true", + }}, &config.Config{}) + require.True(t, svc.IsOpenAIAllowClaudeCodeCodexPluginEnabled(context.Background())) + }) + t.Run("值非 true 时关闭", func(t *testing.T) { + svc := NewSettingService(&allowClaudeCodeSettingRepoStub{values: map[string]string{ + SettingKeyOpenAIAllowClaudeCodeCodexPlugin: "false", + }}, &config.Config{}) + require.False(t, svc.IsOpenAIAllowClaudeCodeCodexPluginEnabled(context.Background())) + }) +} diff --git a/backend/internal/service/settings_view.go b/backend/internal/service/settings_view.go index 3f961ab2..7b45ef1a 100644 --- a/backend/internal/service/settings_view.go +++ b/backend/internal/service/settings_view.go @@ -195,6 +195,7 @@ type SystemSettings struct { RewriteMessageCacheControl bool // 是否改写 messages[*].content[*].cache_control(默认 false) AntigravityUserAgentVersion string // Antigravity 上游 User-Agent 版本号;空值使用配置/默认值 OpenAICodexUserAgent string // OpenAI Codex 上游完整 User-Agent;空值使用内置默认 + OpenAIAllowClaudeCodeCodexPlugin bool // 全局开关:是否额外放行 Claude Code 的 Codex 插件(默认 false) // Web Search Emulation WebSearchEmulationEnabled bool // 是否启用 web search 模拟 diff --git a/frontend/src/api/admin/settings.ts b/frontend/src/api/admin/settings.ts index d2b878cc..6d8e6cee 100644 --- a/frontend/src/api/admin/settings.ts +++ b/frontend/src/api/admin/settings.ts @@ -560,6 +560,7 @@ export interface SystemSettings { rewrite_message_cache_control: boolean; antigravity_user_agent_version: string; openai_codex_user_agent: string; + openai_allow_claude_code_codex_plugin: boolean; web_search_emulation_enabled?: boolean; // Payment configuration @@ -792,6 +793,7 @@ export interface UpdateSettingsRequest { rewrite_message_cache_control?: boolean; antigravity_user_agent_version?: string; openai_codex_user_agent?: string; + openai_allow_claude_code_codex_plugin?: boolean; // Payment configuration payment_enabled?: boolean; risk_control_enabled?: boolean; diff --git a/frontend/src/components/account/BulkEditAccountModal.vue b/frontend/src/components/account/BulkEditAccountModal.vue index c8d53220..6e71fe4b 100644 --- a/frontend/src/components/account/BulkEditAccountModal.vue +++ b/frontend/src/components/account/BulkEditAccountModal.vue @@ -742,6 +742,50 @@
+ +
+
+ + +
+
+

+ {{ t('admin.accounts.openai.codexCLIOnlyAllowClaudeCodeDesc') }} +

+ +
+
+
@@ -1219,6 +1263,7 @@ const enableOpenAIPassthrough = ref(false) const enableOpenAIWSMode = ref(false) const enableOpenAIAPIKeyWSMode = ref(false) const enableCodexCLIOnly = ref(false) +const enableCodexCLIOnlyAllowClaudeCode = ref(false) const enableOpenAICompactMode = ref(false) const enableOpenAICompactModelMapping = ref(false) const enableRpmLimit = ref(false) @@ -1246,6 +1291,7 @@ const openaiPassthroughEnabled = ref(false) const openaiOAuthResponsesWebSocketV2Mode = ref(OPENAI_WS_MODE_OFF) const openaiAPIKeyResponsesWebSocketV2Mode = ref(OPENAI_WS_MODE_OFF) const codexCLIOnlyEnabled = ref(false) +const codexCLIOnlyAllowClaudeCodeEnabled = ref(false) const openAICompactMode = ref('auto') const openAICompactModelMappings = ref([]) const rpmLimitEnabled = ref(false) @@ -1496,6 +1542,11 @@ const buildUpdatePayload = (): Record | null => { extra.codex_cli_only = codexCLIOnlyEnabled.value } + if (enableCodexCLIOnlyAllowClaudeCode.value) { + const extra = ensureExtra() + extra.codex_cli_only_allowed_clients = codexCLIOnlyAllowClaudeCodeEnabled.value ? ['claude_code'] : [] + } + if (enableOpenAICompactMode.value) { const extra = ensureExtra() extra.openai_compact_mode = openAICompactMode.value @@ -1602,6 +1653,7 @@ const handleSubmit = async () => { enableOpenAIWSMode.value || enableOpenAIAPIKeyWSMode.value || enableCodexCLIOnly.value || + enableCodexCLIOnlyAllowClaudeCode.value || enableOpenAICompactMode.value || enableOpenAICompactModelMapping.value || enableRpmLimit.value || @@ -1704,6 +1756,7 @@ watch( enableOpenAIWSMode.value = false enableOpenAIAPIKeyWSMode.value = false enableCodexCLIOnly.value = false + enableCodexCLIOnlyAllowClaudeCode.value = false enableOpenAICompactMode.value = false enableOpenAICompactModelMapping.value = false enableRpmLimit.value = false @@ -1727,6 +1780,7 @@ watch( openaiOAuthResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF openaiAPIKeyResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF codexCLIOnlyEnabled.value = false + codexCLIOnlyAllowClaudeCodeEnabled.value = false openAICompactMode.value = 'auto' openAICompactModelMappings.value = [] rpmLimitEnabled.value = false diff --git a/frontend/src/components/account/CreateAccountModal.vue b/frontend/src/components/account/CreateAccountModal.vue index 331295f7..c699df37 100644 --- a/frontend/src/components/account/CreateAccountModal.vue +++ b/frontend/src/components/account/CreateAccountModal.vue @@ -2635,6 +2635,32 @@ />
+
+
+ +

+ {{ t('admin.accounts.openai.codexCLIOnlyAllowClaudeCodeDesc') }} +

+
+ +
@@ -3353,6 +3379,7 @@ const openAIResponsesMode = ref('auto') const openaiOAuthResponsesWebSocketV2Mode = ref(OPENAI_WS_MODE_OFF) const openaiAPIKeyResponsesWebSocketV2Mode = ref(OPENAI_WS_MODE_OFF) const codexCLIOnlyEnabled = ref(false) +const codexCLIOnlyAllowClaudeCodeEnabled = ref(false) const anthropicPassthroughEnabled = ref(false) const webSearchEmulationMode = ref('default') const webSearchGlobalEnabled = ref(false) @@ -3724,6 +3751,7 @@ watch( openaiOAuthResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF openaiAPIKeyResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF codexCLIOnlyEnabled.value = false + codexCLIOnlyAllowClaudeCodeEnabled.value = false } if (newPlatform !== 'anthropic') { anthropicPassthroughEnabled.value = false @@ -3744,6 +3772,7 @@ watch( ([category, platform]) => { if (platform === 'openai' && category !== 'oauth-based') { codexCLIOnlyEnabled.value = false + codexCLIOnlyAllowClaudeCodeEnabled.value = false } if (platform !== 'anthropic' || category !== 'apikey') { anthropicPassthroughEnabled.value = false @@ -4123,6 +4152,7 @@ const resetForm = () => { openaiOAuthResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF openaiAPIKeyResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF codexCLIOnlyEnabled.value = false + codexCLIOnlyAllowClaudeCodeEnabled.value = false anthropicPassthroughEnabled.value = false webSearchEmulationMode.value = 'default' // Reset quota control state @@ -4201,6 +4231,15 @@ const buildOpenAIExtra = (base?: Record): Record +
+
+ +

+ {{ t('admin.accounts.openai.codexCLIOnlyAllowClaudeCodeDesc') }} +

+
+ +
('auto') const openaiOAuthResponsesWebSocketV2Mode = ref(OPENAI_WS_MODE_OFF) const openaiAPIKeyResponsesWebSocketV2Mode = ref(OPENAI_WS_MODE_OFF) const codexCLIOnlyEnabled = ref(false) +const codexCLIOnlyAllowClaudeCodeEnabled = ref(false) type CodexImageGenerationBridgeMode = 'inherit' | 'enabled' | 'disabled' const codexImageGenerationBridgeMode = ref('inherit') const anthropicPassthroughEnabled = ref(false) @@ -2728,6 +2755,7 @@ const syncFormFromAccount = (newAccount: Account | null) => { openaiOAuthResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF openaiAPIKeyResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF codexCLIOnlyEnabled.value = false + codexCLIOnlyAllowClaudeCodeEnabled.value = false codexImageGenerationBridgeMode.value = 'inherit' anthropicPassthroughEnabled.value = false webSearchEmulationMode.value = 'default' @@ -2759,6 +2787,9 @@ const syncFormFromAccount = (newAccount: Account | null) => { }) if (newAccount.type === 'oauth') { codexCLIOnlyEnabled.value = extra?.codex_cli_only === true + codexCLIOnlyAllowClaudeCodeEnabled.value = + Array.isArray(extra?.codex_cli_only_allowed_clients) && + (extra.codex_cli_only_allowed_clients as unknown[]).includes('claude_code') } const credentials = newAccount.credentials as Record | undefined const compactMappings = credentials?.compact_model_mapping as Record | undefined @@ -3877,6 +3908,12 @@ const handleSubmit = async () => { } else { delete newExtra.codex_cli_only } + // 仅当 codex_cli_only 开启且子开关开启时写入 Claude Code 插件白名单,否则清除避免孤立字段 + if (codexCLIOnlyEnabled.value && codexCLIOnlyAllowClaudeCodeEnabled.value) { + newExtra.codex_cli_only_allowed_clients = ['claude_code'] + } else { + delete newExtra.codex_cli_only_allowed_clients + } } updatePayload.extra = newExtra diff --git a/frontend/src/components/account/__tests__/BulkEditAccountModal.spec.ts b/frontend/src/components/account/__tests__/BulkEditAccountModal.spec.ts index caa307fc..3ae75ee9 100644 --- a/frontend/src/components/account/__tests__/BulkEditAccountModal.spec.ts +++ b/frontend/src/components/account/__tests__/BulkEditAccountModal.spec.ts @@ -197,6 +197,25 @@ describe('BulkEditAccountModal', () => { }) }) + it('OpenAI OAuth 批量编辑应提交 codex_cli_only_allowed_clients 字段', async () => { + const wrapper = mountModal({ + selectedPlatforms: ['openai'], + selectedTypes: ['oauth'] + }) + + await wrapper.get('#bulk-edit-openai-codex-allow-claude-code-enabled').setValue(true) + await wrapper.get('#bulk-edit-openai-codex-allow-claude-code-toggle').trigger('click') + await wrapper.get('#bulk-edit-account-form').trigger('submit.prevent') + await flushPromises() + + expect(adminAPI.accounts.bulkUpdate).toHaveBeenCalledTimes(1) + expect(adminAPI.accounts.bulkUpdate).toHaveBeenCalledWith([1, 2], { + extra: { + codex_cli_only_allowed_clients: ['claude_code'] + } + }) + }) + it('OpenAI API Key 批量编辑应提交 API Key 专属 WS mode 字段', async () => { const wrapper = mountModal({ selectedPlatforms: ['openai'], diff --git a/frontend/src/i18n/locales/en.ts b/frontend/src/i18n/locales/en.ts index ff5ea651..956b5e7a 100644 --- a/frontend/src/i18n/locales/en.ts +++ b/frontend/src/i18n/locales/en.ts @@ -3338,6 +3338,9 @@ export default { codexCLIOnly: 'Codex official clients only', codexCLIOnlyDesc: 'Only applies to OpenAI OAuth. When enabled, only Codex official client families are allowed; when disabled, the gateway bypasses this restriction and keeps existing behavior.', + codexCLIOnlyAllowClaudeCode: "Also allow Claude Code's Codex plugin", + codexCLIOnlyAllowClaudeCodeDesc: + 'Only takes effect when the switch above is on. Additionally allows requests from the Claude Code Codex plugin (exact match on originator=Claude Code) without weakening blocking of other non-official clients.', codexImageGenerationBridge: 'Codex image-generation bridge', codexImageGenerationBridgeDesc: 'Account policy takes precedence over channel and global settings. Only controls whether Codex requests through the /responses text endpoint receive the image_generation tool; standalone image-generation endpoints are unaffected.', @@ -5577,6 +5580,9 @@ export default { openaiCodexUserAgent: 'OpenAI Codex UA', openaiCodexUserAgentPlaceholder: 'codex-tui/0.125.0 (Ubuntu 22.4.0; x86_64) xterm-256color (codex-tui; 0.125.0)', openaiCodexUserAgentHint: 'Used to bypass Cloudflare browser-UA challenges on the OpenAI upstream. Only applies when the client User-Agent is detected as a browser (Mozilla/...). Leave empty to use the built-in default.', + openaiAllowClaudeCodeCodexPlugin: "Allow using the Codex plugin in Claude Code", + openaiAllowClaudeCodeCodexPluginDesc: + "Global switch; only affects OpenAI OAuth accounts that have 'Codex official clients only' enabled. When on, all such accounts additionally allow requests from the Claude Code Codex plugin (exact match on originator=Claude Code) without per-account config; upstream requests remain pass-through.", }, webSearchEmulation: { title: 'Web Search Emulation', diff --git a/frontend/src/i18n/locales/zh.ts b/frontend/src/i18n/locales/zh.ts index b8ac7d2c..2bdebf06 100644 --- a/frontend/src/i18n/locales/zh.ts +++ b/frontend/src/i18n/locales/zh.ts @@ -3483,6 +3483,8 @@ export default { responsesStatusForcedChatCompletions: '已强制 Chat Completions', codexCLIOnly: '仅允许 Codex 官方客户端', codexCLIOnlyDesc: '仅对 OpenAI OAuth 生效。开启后仅允许 Codex 官方客户端家族访问;关闭后完全绕过并保持原逻辑。', + codexCLIOnlyAllowClaudeCode: '额外放行 Claude Code 的 Codex 插件', + codexCLIOnlyAllowClaudeCodeDesc: '仅在上方开关开启时生效。额外放行通过 Claude Code 的 Codex 插件发起的请求(精确匹配 originator=Claude Code),不影响对其他非官方客户端的拦截。', codexImageGenerationBridge: 'Codex 图片生成桥接', codexImageGenerationBridgeDesc: '账号级策略优先于渠道和全局配置。仅控制 Codex 走 /responses 文本端点时是否注入 image_generation 工具;不影响独立图片生成接口。', @@ -5733,6 +5735,9 @@ export default { openaiCodexUserAgent: 'OpenAI Codex UA', openaiCodexUserAgentPlaceholder: 'codex-tui/0.125.0 (Ubuntu 22.4.0; x86_64) xterm-256color (codex-tui; 0.125.0)', openaiCodexUserAgentHint: '用于规避 OpenAI 上游 Cloudflare 对浏览器 UA 的访问质询。仅在检测到客户端 User-Agent 为浏览器(Mozilla/...)时生效,其他客户端原样透传。留空使用内置默认值。', + openaiAllowClaudeCodeCodexPlugin: '允许在 Claude Code 中使用 Codex 插件', + openaiAllowClaudeCodeCodexPluginDesc: + '全局开关,仅对已开启「仅允许 Codex 官方客户端」的 OpenAI OAuth 账号生效。开启后,所有此类账号都额外放行通过 Claude Code 的 Codex 插件发起的请求(精确匹配 originator=Claude Code),无需逐账号配置;上游请求仍保持透传。', }, webSearchEmulation: { title: 'Web Search 模拟', diff --git a/frontend/src/views/admin/SettingsView.vue b/frontend/src/views/admin/SettingsView.vue index 68eb4849..239ce2d7 100644 --- a/frontend/src/views/admin/SettingsView.vue +++ b/frontend/src/views/admin/SettingsView.vue @@ -3948,6 +3948,19 @@ }}

+ + +
+
+ +

+ {{ t("admin.settings.gatewayForwarding.openaiAllowClaudeCodeCodexPluginDesc") }} +

+
+ +
@@ -7162,6 +7175,7 @@ const form = reactive({ rewrite_message_cache_control: false, antigravity_user_agent_version: "", openai_codex_user_agent: "", + openai_allow_claude_code_codex_plugin: false, // 余额、订阅到期与账号限额通知 balance_low_notify_enabled: false, balance_low_notify_threshold: 0, @@ -8267,6 +8281,7 @@ async function saveSettings() { form.antigravity_user_agent_version?.trim() || "", openai_codex_user_agent: form.openai_codex_user_agent?.trim() || "", + openai_allow_claude_code_codex_plugin: form.openai_allow_claude_code_codex_plugin, // Payment configuration payment_enabled: form.payment_enabled, risk_control_enabled: form.risk_control_enabled, From ddf91e9a7f7f7c72160d5097ffe2211ee6c20930 Mon Sep 17 00:00:00 2001 From: alfadb Date: Wed, 27 May 2026 15:46:05 +0800 Subject: [PATCH 16/79] =?UTF-8?q?fix(gateway):=20=E6=8C=89=E6=9C=80?= =?UTF-8?q?=E7=BB=88=20anthropic-beta=20header=20=E5=AF=B9=20body.context?= =?UTF-8?q?=5Fmanagement=20=E5=81=9A=E8=83=BD=E5=8A=9B=E7=BB=B4=E5=BA=A6?= =?UTF-8?q?=20sanitize?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 上游 Anthropic 在 body 含 `context_management` 但最终发出去的 `anthropic-beta` header 不含 `context-management-2025-06-27` 时会拒收: { "type": "invalid_request_error", "message": "context_management: Extra inputs are not permitted" } (HTTP 400, request_id 形如 req_011C...) 该 400 在 haiku 路径上触发,因为三个 beta header 构造器有意排除了 context-management beta: - HaikuBetaHeader (messages, OAuth / mimic CC) - APIKeyHaikuBetaHeader (messages, API-key) - CountTokensBetaHeader (count_tokens, 所有认证类型) 但 body 中仍然带着 `context_management` 字段,原因有二: 1. normalizeClaudeOAuthRequestBody 在 thinking_enabled / thinking_adaptive 打开时为 `clear_thinking_20251015` 主动注入; 2. 客户端 (Claude Code CLI >= 2.1.87) 原样发送, 网关透传时一并转发。 修复方案: 能力维度对称约束 ========================== 对齐已有的 Bedrock 模式 (`backend/internal/service/bedrock_request.go` 中的 `sanitizeBedrockFieldsForBetaTokens`): 根据 **最终** 发出的 `anthropic-beta` header 决定是否保留 `body.context_management`, 而不是按 model 名或路由分类来决定。 新增纯函数: sanitizeAnthropicBodyForBetaTokens(body, betaHeader) (body, changed) 如果 `betaHeader` 不含 `context-management-2025-06-27`, 用 sjson 把 body 字段 strip 掉; 否则原样返回。 在所有 Anthropic / Anthropic-兼容 上游出口都接入: | 路径 | sanitize 接入点 | |--------------------------------------------|-------------------------------------------------------| | /v1/messages OAuth mimic CC | buildUpstreamRequest | | /v1/messages OAuth 真 CC 透传 | buildUpstreamRequest | | /v1/messages API-key | buildUpstreamRequest | | /v1/messages API-key passthrough | buildUpstreamRequestAnthropicAPIKeyPassthrough | | /v1/messages Vertex / service-account | buildUpstreamRequestAnthropicVertex | | /v1/messages/count_tokens (全部 4 条路径) | buildCountTokensRequest, | | | buildCountTokensRequestAnthropicAPIKeyPassthrough | | Antigravity Anthropic-兼容 上游 | AntigravityGatewayService.ForwardUpstream | | Bedrock | (已由 sanitizeBedrockFieldsForBetaTokens 处理) | 为什么要重排 (而不是加一行调用) ================================ sanitize 必须 **在** `signBillingHeaderCCH` 之前运行。CCH 对整个 body 取 xxHash64 摘要后写入 billing header 里 5 位十六进制的 `cch` 字段; 如果先签名再 strip, 上游对发出去的 body 重算 hash 会和 `cch` 不一致, 请求被判为 third-party。这就要求在 `http.NewRequest` 之前算出最终的 `anthropic-beta` header, 所以把原本内联在 builder 里的 beta 计算逻辑 抽成了两个纯函数: - computeFinalAnthropicBeta (messages 路径: mimic 不透传 客户端 beta) - computeFinalCountTokensAnthropicBeta (count_tokens 路径: mimic 不 跳过白名单透传) 两者逐位保留原行为: - mimic 路径在 messages 上跳过客户端 beta, 在 count_tokens 上合并 - API-key 路径尊重 `InjectBetaForAPIKey` 开关 - dropSet (`defaultDroppedBetasSet` + BetaPolicy filter) 应用在主路径, passthrough / Vertex 路径有意不应用 —— 这条原有的不对称行为本 PR 不动。 一条语义测试 (`TestSanitizeMustBeBeforeCCHSigning_HashConsistency`) 把 顺序约束文档化并强制守住: 它证明 `sanitize -> signBillingHeaderCCH` 产生的 `cch` 与最终 body 一致, 而 `signBillingHeaderCCH -> sanitize` 产生的 `cch` 会被上游 hash 重算判失败。 为什么是能力维度 (而不是 haiku 模型名匹配) ========================================== 最朴素的"按 model 名 strip"方案 (`strings.Contains(modelID, "haiku") -> DeleteBytes "context_management"`) 有四个真实失败模式: 1. 过度删除。CLI >= 2.1.87 的真 Claude Code 客户端在 haiku 上同时 发送 body 字段 **和** `anthropic-beta: context-management-2025-06-27`。 一律 strip 会让该用户的 `clear_thinking_20251015` 静默失效。 2. 别名漂移。未来的 haiku 别名 (`claude-3-haiku-...`, `claude-haiku-...` 等) 改变匹配面; 任何新别名都会悄悄绕过 strip。 3. count_tokens 漏覆盖。count_tokens 有自己的 builder 和不同的 beta header 集合; 在一个地方做 model 名检查会漏掉这条路径。 4. API-key passthrough 早退。passthrough builder 在 model 名 strip 之前就 return 了, strip 根本不执行。 能力维度沿着 header 端到端走, 上述 4 个 case 都由构造方式保证正确, 不依赖任何 modelID 匹配。 防御项 ====== - 当 `sjson.DeleteBytes` 在 `gjson` 刚验证过字段存在的 body 上失败时, `sanitizeAnthropicBodyForBetaTokens` 会记 warning 日志 —— 这种情况 现实中仅在请求中途被破坏时发生, 日志把此前会静默发生的 body / header 不一致暴露出来。 - `header_util.go` 新增 `deleteHeaderAllForms`: 在白名单透传已经写入 canonical 大小写的 `Anthropic-Beta` 之后再覆盖, 否则会同时留下两条。 测试 ==== `backend/internal/service` 下新增 44 个测试: - 纯函数: anthropicBetaTokensContains x 5, sanitize keep/strip x 6, computeFinal{Anthropic,CountTokens}AnthropicBeta x 12 - normalize 回归 x 5 - buildUpstreamRequest 端到端 x 4 (OAuth mimic haiku strip / mimic sonnet preserve / 真 CC haiku 带客户端 beta preserve / API-key haiku strip) - buildCountTokensRequest 端到端 x 2 - buildUpstreamRequestAnthropicAPIKeyPassthrough x 2 (strip / preserve) - buildCountTokensRequestAnthropicAPIKeyPassthrough x 2 (strip / preserve) - buildUpstreamRequestAnthropicVertex x 2 (strip / preserve, 含 outgoing `anthropic-beta` header 对称断言) - CCH 顺序语义测试 x 1 unit 套件全过 (本机 88s), `golangci-lint` 0 issues。 已知局限 (本 PR 范围外) ======================== - Vertex 路径用透传过来的客户端 `anthropic-beta` header 作为 sanitize 依据, 而不是 Vertex 侧的能力矩阵。最坏情况是过度 strip (= 当前 main 的行为, 主路径本来什么都不 strip); 不是 regression。完整的 Vertex 能力模型属于单独的 PR。 - Vertex builder 仍然不应用 BetaPolicy filter / dropSet。这是该 builder 早 return 的既有架构决策, 本 PR 不动。 - count_tokens mimic 在 haiku 上仍然注入 `context-management-2025-06-27` (因为原 count_tokens mimic 逻辑并不像 messages mimic 那样排除它)。 本 PR 逐位保留 main 的行为; 是否要让它与 messages mimic 的排除策略 统一是另一个问题。 - `sanitizeAnthropicBodyForBetaTokens` 目前只处理 `context_management <-> context-management-2025-06-27` 这一对。如果 Anthropic 后续推出更多 beta-gated body 字段, 可以在后续 PR 重构为 `{body 路径 -> required beta token}` 注册表的形式。 --- .../service/antigravity_gateway_service.go | 10 +- ...y_anthropic_vertex_service_account_test.go | 64 ++ .../gateway_context_management_test.go | 667 ++++++++++++++++++ backend/internal/service/gateway_request.go | 64 ++ backend/internal/service/gateway_service.go | 279 ++++++-- backend/internal/service/header_util.go | 14 + 6 files changed, 1024 insertions(+), 74 deletions(-) create mode 100644 backend/internal/service/gateway_context_management_test.go diff --git a/backend/internal/service/antigravity_gateway_service.go b/backend/internal/service/antigravity_gateway_service.go index 9882b010..2b849bdd 100644 --- a/backend/internal/service/antigravity_gateway_service.go +++ b/backend/internal/service/antigravity_gateway_service.go @@ -4209,6 +4209,14 @@ func (s *AntigravityGatewayService) ForwardUpstream(ctx context.Context, c *gin. // 构建上游请求 URL upstreamURL := baseURL + "/v1/messages" + // 能力维度 sanitize:Anthropic-compatible 上游透传路径也需要保证 body↔beta header + // 对称。客户端 anthropic-beta header 不含 context-management-2025-06-27 但 body 带 + // context_management 时 strip,与 Anthropic 直连 / Bedrock / Vertex 路径保持一致。 + clientBeta := c.GetHeader("anthropic-beta") + if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, clientBeta); changed { + body = sanitized + } + // 创建请求 req, err := http.NewRequestWithContext(ctx, http.MethodPost, upstreamURL, bytes.NewReader(body)) if err != nil { @@ -4224,7 +4232,7 @@ func (s *AntigravityGatewayService) ForwardUpstream(ctx context.Context, c *gin. if v := c.GetHeader("anthropic-version"); v != "" { req.Header.Set("anthropic-version", v) } - if v := c.GetHeader("anthropic-beta"); v != "" { + if v := clientBeta; v != "" { req.Header.Set("anthropic-beta", v) } 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 aa779805..2f42b0ab 100644 --- a/backend/internal/service/gateway_anthropic_vertex_service_account_test.go +++ b/backend/internal/service/gateway_anthropic_vertex_service_account_test.go @@ -66,3 +66,67 @@ func readRequestBodyForTest(t *testing.T, req *http.Request) []byte { require.NoError(t, err) return body } + +// Vertex 路径回归保护:同样需要 +// body↔beta header 能力维度对称。客户端 header 不带 context-management beta +// 但 body 带 context_management 字段 → Vertex builder 必须 strip 字段,与 Anthropic +// 直连 / Bedrock 路径保持一致。 +func TestGatewayService_BuildAnthropicVertexServiceAccount_StripsContextManagementWhenBetaMissing(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil) + // 客户端 header 只带 interleaved-thinking,不带 context-management-2025-06-27 + c.Request.Header.Set("Anthropic-Beta", "interleaved-thinking-2025-05-14") + + account := &Account{ + ID: 302, Platform: PlatformAnthropic, Type: AccountTypeServiceAccount, + Credentials: map[string]any{"project_id": "vertex-proj", "location": "us-east5"}, + } + // body 带了 context_management 字段(客户端透传 / normalize 补齐 / mimicry 注入等场景都可能导致) + 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( + context.Background(), c, account, body, + "vertex-token", "service_account", "claude-haiku-4-5@20251001", false, false, + ) + require.NoError(t, err) + + got := readRequestBodyForTest(t, req) + require.False(t, gjson.GetBytes(got, "context_management").Exists(), + "Vertex 路径下客户端 header 缺 context-management beta 时,必须 strip body 同名字段") + // header 对称断言:覆盖未来某人在 Vertex builder 里加入与 sanitize 不一致的 header 处理。 + outBeta := getHeaderRaw(req.Header, "anthropic-beta") + require.False(t, anthropicBetaTokensContains(outBeta, "context-management-2025-06-27"), + "与 body 对称:outgoing anthropic-beta header 也不含 context-management beta") +} + +// Vertex 路径反面:客户端 header 含 context-management beta 时保留字段。 +func TestGatewayService_BuildAnthropicVertexServiceAccount_PreservesContextManagementWhenBetaPresent(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil) + c.Request.Header.Set("Anthropic-Beta", "interleaved-thinking-2025-05-14,context-management-2025-06-27") + + account := &Account{ + ID: 303, Platform: PlatformAnthropic, Type: AccountTypeServiceAccount, + Credentials: map[string]any{"project_id": "vertex-proj", "location": "us-east5"}, + } + body := []byte(`{"model":"claude-sonnet-4-6","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`) + + svc := &GatewayService{} + req, err := svc.buildUpstreamRequest( + context.Background(), c, account, body, + "vertex-token", "service_account", "claude-sonnet-4-6@20260218", false, false, + ) + require.NoError(t, err) + + got := readRequestBodyForTest(t, req) + require.True(t, gjson.GetBytes(got, "context_management").Exists(), + "Vertex + 客户端 header 包含 context-management beta 时字段必须保留") + outBeta := getHeaderRaw(req.Header, "anthropic-beta") + require.True(t, anthropicBetaTokensContains(outBeta, "context-management-2025-06-27"), + "与 body 对称:outgoing anthropic-beta header 同步含 context-management beta") +} diff --git a/backend/internal/service/gateway_context_management_test.go b/backend/internal/service/gateway_context_management_test.go new file mode 100644 index 00000000..c2263bdc --- /dev/null +++ b/backend/internal/service/gateway_context_management_test.go @@ -0,0 +1,667 @@ +//go:build unit + +package service + +import ( + "context" + "io" + "net/http" + "net/http/httptest" + "regexp" + "strings" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/claude" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +// ============================================================================ +// 背景 +// ============================================================================ +// +// Anthropic 上游对 body.context_management 字段实施 Pydantic schema 校验: +// 当且仅当 anthropic-beta header 含 context-management-2025-06-27 时接受。 +// 否则报: +// "context_management: Extra inputs are not permitted" +// +// 本仓采用能力维度对称约束(与 Bedrock 路径的 sanitizeBedrockFieldsForBetaTokens +// 对称):在所有 Anthropic 直连出口,按最终 anthropic-beta header 是否含上述 token +// 决定 body 是否保留同名字段。 +// +// 本文件覆盖: +// 1) sanitizeAnthropicBodyForBetaTokens 纯函数 +// 2) anthropicBetaTokensContains 解析辅助函数 +// 3) computeFinalAnthropicBeta / computeFinalCountTokensAnthropicBeta 各路径 +// 4) normalizeClaudeOAuthRequestBody 的 context_management 补齐行为(不再按 model 短路) + +// ============================================================================ +// anthropicBetaTokensContains +// ============================================================================ + +func TestAnthropicBetaTokensContains_EmptyInputs(t *testing.T) { + require.False(t, anthropicBetaTokensContains("", "context-management-2025-06-27")) + require.False(t, anthropicBetaTokensContains("oauth-2025-04-20", "")) +} + +func TestAnthropicBetaTokensContains_SingleToken(t *testing.T) { + require.True(t, anthropicBetaTokensContains("context-management-2025-06-27", "context-management-2025-06-27")) +} + +func TestAnthropicBetaTokensContains_MultiTokenComma(t *testing.T) { + header := "oauth-2025-04-20,context-management-2025-06-27,interleaved-thinking-2025-05-14" + require.True(t, anthropicBetaTokensContains(header, "context-management-2025-06-27")) + require.True(t, anthropicBetaTokensContains(header, "oauth-2025-04-20")) + require.False(t, anthropicBetaTokensContains(header, "fast-mode-2026-02-01")) +} + +func TestAnthropicBetaTokensContains_ToleratesWhitespace(t *testing.T) { + header := "oauth-2025-04-20 , context-management-2025-06-27 , interleaved-thinking-2025-05-14" + require.True(t, anthropicBetaTokensContains(header, "context-management-2025-06-27")) +} + +func TestAnthropicBetaTokensContains_SubstringNotMatched(t *testing.T) { + // 严格 token 比较,不应被子串误匹配 + require.False(t, anthropicBetaTokensContains("context-management-2025-06-27-rev2", "context-management-2025-06-27"), + "必须按 token 边界匹配,不允许 prefix 子串误命中") +} + +// ============================================================================ +// sanitizeAnthropicBodyForBetaTokens +// ============================================================================ + +func TestSanitizeAnthropicBodyForBetaTokens_NoFieldNoChange(t *testing.T) { + body := []byte(`{"model":"claude-haiku-4-5","messages":[]}`) + out, changed := sanitizeAnthropicBodyForBetaTokens(body, "oauth-2025-04-20") + require.False(t, changed) + require.Equal(t, string(body), string(out)) +} + +func TestSanitizeAnthropicBodyForBetaTokens_FieldKeptWhenBetaPresent(t *testing.T) { + body := []byte(`{"model":"claude-opus-4-7","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`) + out, changed := sanitizeAnthropicBodyForBetaTokens(body, + "oauth-2025-04-20,context-management-2025-06-27,interleaved-thinking-2025-05-14") + require.False(t, changed) + require.True(t, gjson.GetBytes(out, "context_management").Exists()) + require.Equal(t, "clear_thinking_20251015", + gjson.GetBytes(out, "context_management.edits.0.type").String()) +} + +func TestSanitizeAnthropicBodyForBetaTokens_FieldStrippedWhenBetaMissing(t *testing.T) { + body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`) + out, changed := sanitizeAnthropicBodyForBetaTokens(body, "oauth-2025-04-20,interleaved-thinking-2025-05-14") + require.True(t, changed) + require.False(t, gjson.GetBytes(out, "context_management").Exists(), + "header 不含 context-management beta 时必须 strip 同名字段") +} + +func TestSanitizeAnthropicBodyForBetaTokens_FieldStrippedWhenBetaEmpty(t *testing.T) { + body := []byte(`{"context_management":{"edits":[]},"messages":[]}`) + out, changed := sanitizeAnthropicBodyForBetaTokens(body, "") + require.True(t, changed) + require.False(t, gjson.GetBytes(out, "context_management").Exists()) +} + +func TestSanitizeAnthropicBodyForBetaTokens_EmptyBody(t *testing.T) { + out, changed := sanitizeAnthropicBodyForBetaTokens([]byte{}, "") + require.False(t, changed) + require.Empty(t, out) + + out, changed = sanitizeAnthropicBodyForBetaTokens(nil, "") + require.False(t, changed) + require.Empty(t, out) +} + +// ★ 关键回归断言:能力维度 sanitize 解决了 "真 CC + haiku" 路径的过度删除问题。 +// 真实 Claude Code CLI 2.1.87+ 客户端 header 含 context-management beta; +// 即使 model 是 haiku,sanitize 也不应剥离功能字段。 +func TestSanitizeAnthropicBodyForBetaTokens_HaikuRealCCClientPreservesField(t *testing.T) { + body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015","keep":"all"}]},"messages":[]}`) + // 真 Claude Code CLI 2.1.87+ 客户端 header 含 context-management beta + clientBeta := "claude-code-20250219,oauth-2025-04-20,interleaved-thinking-2025-05-14,context-management-2025-06-27" + out, changed := sanitizeAnthropicBodyForBetaTokens(body, clientBeta) + require.False(t, changed, + "真 CC 客户端 header 含 context-management beta 时,haiku body 字段必须保留(功能不丢)") + require.True(t, gjson.GetBytes(out, "context_management").Exists()) +} + +// ============================================================================ +// computeFinalAnthropicBeta — 关键路径 +// ============================================================================ + +func newTestGatewayServiceForBeta(injectBetaForAPIKey bool) *GatewayService { + cfg := &config.Config{} + cfg.Gateway.InjectBetaForAPIKey = injectBetaForAPIKey + return &GatewayService{cfg: cfg} +} + +func TestComputeFinalAnthropicBeta_OAuthMimic_NonHaiku_IncludesContextManagement(t *testing.T) { + s := newTestGatewayServiceForBeta(false) + final, ok := s.computeFinalAnthropicBeta("oauth", true, "claude-sonnet-4-6", http.Header{}, []byte(`{}`), nil) + require.True(t, ok) + require.True(t, anthropicBetaTokensContains(final, claude.BetaContextManagement), + "OAuth mimic non-haiku 必须注入完整 CC mimicry beta,含 context-management-2025-06-27") + require.True(t, anthropicBetaTokensContains(final, claude.BetaOAuth)) + require.True(t, anthropicBetaTokensContains(final, claude.BetaClaudeCode)) +} + +func TestComputeFinalAnthropicBeta_OAuthMimic_Haiku_ExcludesContextManagement(t *testing.T) { + s := newTestGatewayServiceForBeta(false) + final, ok := s.computeFinalAnthropicBeta("oauth", true, "claude-haiku-4-5", http.Header{}, []byte(`{}`), nil) + require.True(t, ok) + require.False(t, anthropicBetaTokensContains(final, claude.BetaContextManagement), + "OAuth mimic haiku 仅注入 oauth + interleaved-thinking,不含 context-management") + require.True(t, anthropicBetaTokensContains(final, claude.BetaOAuth)) + require.True(t, anthropicBetaTokensContains(final, claude.BetaInterleavedThinking)) +} + +func TestComputeFinalAnthropicBeta_OAuthMimic_IgnoresClientBeta(t *testing.T) { + // mimic 路径下原代码白名单透传被跳过,client beta 应被忽略 + s := newTestGatewayServiceForBeta(false) + hdr := http.Header{} + hdr.Set("anthropic-beta", "custom-experimental-beta") + final, ok := s.computeFinalAnthropicBeta("oauth", true, "claude-sonnet-4-6", hdr, []byte(`{}`), nil) + require.True(t, ok) + require.False(t, strings.Contains(final, "custom-experimental-beta"), + "mimic 路径必须忽略客户端 anthropic-beta header") +} + +func TestComputeFinalAnthropicBeta_OAuthTransparent_NonHaiku_PreservesClientContextManagement(t *testing.T) { + // 真 CC 客户端透传:客户端 header 中的 context-management beta 必须保留 + s := newTestGatewayServiceForBeta(false) + hdr := http.Header{} + hdr.Set("anthropic-beta", "claude-code-20250219,oauth-2025-04-20,context-management-2025-06-27") + final, ok := s.computeFinalAnthropicBeta("oauth", false, "claude-sonnet-4-6", hdr, []byte(`{}`), nil) + require.True(t, ok) + require.True(t, anthropicBetaTokensContains(final, claude.BetaContextManagement)) +} + +func TestComputeFinalAnthropicBeta_OAuthTransparent_Haiku_RealCCPreservesContextManagement(t *testing.T) { + // haiku 透传 + 客户端带 context-management beta → 必须保留 + // (能力维度核心场景:避免 model-name 误删客户端透传的功能 beta) + s := newTestGatewayServiceForBeta(false) + hdr := http.Header{} + hdr.Set("anthropic-beta", "claude-code-20250219,oauth-2025-04-20,context-management-2025-06-27,interleaved-thinking-2025-05-14") + final, ok := s.computeFinalAnthropicBeta("oauth", false, "claude-haiku-4-5", hdr, []byte(`{}`), nil) + require.True(t, ok) + require.True(t, anthropicBetaTokensContains(final, claude.BetaContextManagement), + "真 CC + haiku + 客户端带 context-management beta → 透传必须保留") +} + +func TestComputeFinalAnthropicBeta_APIKey_PassesClientBetaThroughDropSet(t *testing.T) { + s := newTestGatewayServiceForBeta(false) + hdr := http.Header{} + hdr.Set("anthropic-beta", "oauth-2025-04-20,custom-beta") + final, ok := s.computeFinalAnthropicBeta("apikey", false, "claude-sonnet-4-6", hdr, []byte(`{}`), nil) + require.True(t, ok) + require.True(t, anthropicBetaTokensContains(final, "oauth-2025-04-20")) + require.True(t, anthropicBetaTokensContains(final, "custom-beta")) +} + +func TestComputeFinalAnthropicBeta_APIKey_NoClientBetaInjectOff_ShouldNotSet(t *testing.T) { + s := newTestGatewayServiceForBeta(false) + final, ok := s.computeFinalAnthropicBeta("apikey", false, "claude-sonnet-4-6", http.Header{}, []byte(`{}`), nil) + require.False(t, ok, "API-key + 客户端未传 + InjectBetaForAPIKey 关 → 不应主动设置 anthropic-beta") + require.Equal(t, "", final) +} + +// ============================================================================ +// computeFinalCountTokensAnthropicBeta +// ============================================================================ + +func TestComputeFinalCountTokensAnthropicBeta_OAuthMimic_AlwaysIncludesContextManagement(t *testing.T) { + // count_tokens 路径下 mimic 不按 haiku 排除:始终注入完整 mimicry beta + s := newTestGatewayServiceForBeta(false) + final, ok := s.computeFinalCountTokensAnthropicBeta("oauth", true, "claude-haiku-4-5", http.Header{}, []byte(`{}`), nil) + require.True(t, ok) + require.True(t, anthropicBetaTokensContains(final, claude.BetaContextManagement), + "count_tokens + mimic 即使 haiku 也注入 context-management beta(与 messages 不同)") + require.True(t, anthropicBetaTokensContains(final, claude.BetaTokenCounting), + "count_tokens 路径必须含 token-counting beta") +} + +// 重构等价性回归: +// 原 main buildCountTokensRequest 在 count_tokens mimic 分支上不跳过白名单透传 +// (与 messages mimic 不同),incomingBeta 取自客户端透传。重构后必须从 clientHeaders +// 拿同一个值并 merge,否则会丢失客户端 beta。 +func TestComputeFinalCountTokensAnthropicBeta_OAuthMimic_PreservesClientBeta(t *testing.T) { + s := newTestGatewayServiceForBeta(false) + hdr := http.Header{} + hdr.Set("anthropic-beta", "custom-experimental-beta,context-1m-2025-08-07") + final, ok := s.computeFinalCountTokensAnthropicBeta("oauth", true, "claude-haiku-4-5", hdr, []byte(`{}`), nil) + require.True(t, ok) + require.True(t, anthropicBetaTokensContains(final, "custom-experimental-beta"), + "count_tokens mimic 不同于 messages mimic:原代码会保留客户端透传的 beta") + require.True(t, anthropicBetaTokensContains(final, "context-1m-2025-08-07"), + "客户端透传的其他 beta token 同样需要保留") + require.True(t, anthropicBetaTokensContains(final, claude.BetaContextManagement), + "同时 FullClaudeCodeMimicryBetas 不打折扣") + require.True(t, anthropicBetaTokensContains(final, claude.BetaTokenCounting), + "同时补齐 token-counting beta") +} + +// messages mimic 路径反向验证:原代码会跳过白名单透传, +// 客户端 beta 不会进入 mimic 计算。重构后 messages computeFinalAnthropicBeta +// mimic 分支依然不该使用 clientBeta。 +func TestComputeFinalAnthropicBeta_OAuthMimic_IgnoresClientBetaExplicit(t *testing.T) { + s := newTestGatewayServiceForBeta(false) + hdr := http.Header{} + hdr.Set("anthropic-beta", "custom-experimental-beta") + final, ok := s.computeFinalAnthropicBeta("oauth", true, "claude-sonnet-4-6", hdr, []byte(`{}`), nil) + require.True(t, ok) + require.False(t, anthropicBetaTokensContains(final, "custom-experimental-beta"), + "messages mimic 原代码跳过白名单透传 → 客户端 beta 不进入计算。"+ + "与 count_tokens mimic 是不同的设计,不能合并为同一函数。") +} + +func TestComputeFinalCountTokensAnthropicBeta_OAuthTransparent_NoClientBetaInjectsDefault(t *testing.T) { + // 真 CC 客户端透传 + 客户端未传 anthropic-beta → 用 CountTokensBetaHeader 兜底 + s := newTestGatewayServiceForBeta(false) + final, ok := s.computeFinalCountTokensAnthropicBeta("oauth", false, "claude-haiku-4-5", http.Header{}, []byte(`{}`), nil) + require.True(t, ok) + require.Equal(t, claude.CountTokensBetaHeader, final) + // CountTokensBetaHeader 不含 context-management beta + require.False(t, anthropicBetaTokensContains(final, claude.BetaContextManagement)) +} + +func TestComputeFinalCountTokensAnthropicBeta_OAuthTransparent_AppendsBetaTokenCounting(t *testing.T) { + s := newTestGatewayServiceForBeta(false) + hdr := http.Header{} + hdr.Set("anthropic-beta", "oauth-2025-04-20,context-management-2025-06-27") + final, ok := s.computeFinalCountTokensAnthropicBeta("oauth", false, "claude-sonnet-4-6", hdr, []byte(`{}`), nil) + require.True(t, ok) + require.True(t, anthropicBetaTokensContains(final, claude.BetaTokenCounting), + "客户端未带 token-counting beta 时必须补齐") + require.True(t, anthropicBetaTokensContains(final, claude.BetaContextManagement), + "客户端带的 context-management beta 必须保留") +} + +// ============================================================================ +// normalizeClaudeOAuthRequestBody — 回归:context_management 补齐恢复原行为 +// ============================================================================ +// +// 重构后该函数不再按 model 名短路:thinking=enabled/adaptive 时补齐 context_management, +// 与 model 无关。strip 责任移交 sanitizeAnthropicBodyForBetaTokens(在 +// buildUpstreamRequest 层按最终 beta header 执行)。 + +func TestNormalizeClaudeOAuthRequestBody_InjectsContextManagement_ThinkingEnabled(t *testing.T) { + body := []byte(`{"model":"claude-sonnet-4-6","thinking":{"type":"enabled","budget_tokens":1000},"messages":[]}`) + out, _ := normalizeClaudeOAuthRequestBody(body, "claude-sonnet-4-6", claudeOAuthNormalizeOptions{}) + require.True(t, gjson.GetBytes(out, "context_management").Exists()) + require.Equal(t, "clear_thinking_20251015", + gjson.GetBytes(out, "context_management.edits.0.type").String()) +} + +func TestNormalizeClaudeOAuthRequestBody_InjectsContextManagement_ThinkingAdaptive(t *testing.T) { + body := []byte(`{"model":"claude-opus-4-7","thinking":{"type":"adaptive"},"messages":[]}`) + out, _ := normalizeClaudeOAuthRequestBody(body, "claude-opus-4-7", claudeOAuthNormalizeOptions{}) + require.True(t, gjson.GetBytes(out, "context_management").Exists()) +} + +func TestNormalizeClaudeOAuthRequestBody_HaikuStillInjects_StripDeferredToSanitize(t *testing.T) { + // haiku + thinking=enabled:normalize 阶段仍按 CLI mimicry 行为补齐字段; + // strip 由 buildUpstreamRequest 层的 sanitize 兜底(如果 final beta 不含 token)。 + body := []byte(`{"model":"claude-haiku-4-5","thinking":{"type":"enabled","budget_tokens":1000},"messages":[]}`) + out, _ := normalizeClaudeOAuthRequestBody(body, "claude-haiku-4-5", claudeOAuthNormalizeOptions{}) + require.True(t, gjson.GetBytes(out, "context_management").Exists(), + "normalize 不再按 model 名短路;strip 责任移交 sanitize 层") +} + +func TestNormalizeClaudeOAuthRequestBody_PreservesClientContextManagement(t *testing.T) { + body := []byte(`{"model":"claude-opus-4-7","context_management":{"edits":[{"type":"custom_strategy"}]},"thinking":{"type":"enabled","budget_tokens":1000},"messages":[]}`) + out, _ := normalizeClaudeOAuthRequestBody(body, "claude-opus-4-7", claudeOAuthNormalizeOptions{}) + require.Equal(t, "custom_strategy", + gjson.GetBytes(out, "context_management.edits.0.type").String(), + "客户端透传的 context_management 内容必须原样保留") +} + +func TestNormalizeClaudeOAuthRequestBody_NoThinking_NoInject(t *testing.T) { + body := []byte(`{"model":"claude-sonnet-4-6","messages":[]}`) + out, _ := normalizeClaudeOAuthRequestBody(body, "claude-sonnet-4-6", claudeOAuthNormalizeOptions{}) + require.False(t, gjson.GetBytes(out, "context_management").Exists()) +} + +// ============================================================================ +// passthrough 集成测试:buildUpstreamRequest- +// AnthropicAPIKeyPassthrough 与 buildCountTokensRequestAnthropicAPIKeyPassthrough +// 路径上 sanitize 是否生效。 +// ============================================================================ + +// passthrough 集成测试不设 base_url,避开 validateUpstreamBaseURL 对 cfg.Security 的依赖。 +// targetURL 会走默认 claudeAPIURL,sanitize 逻辑与 baseURL 是否存在无关。 +func newAnthropicAPIKeyPassthroughAccountForBetaTest() *Account { + return &Account{ + ID: 501, + Name: "anthropic-apikey-passthrough-ctxmgmt-test", + Platform: PlatformAnthropic, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_key": "upstream-key", + }, + Extra: map[string]any{"anthropic_passthrough": true}, + Status: StatusActive, + Schedulable: true, + } +} + +func readUpstreamBodyForTest(t *testing.T, req *http.Request) []byte { + t.Helper() + require.NotNil(t, req.Body) + b, err := io.ReadAll(req.Body) + require.NoError(t, err) + return b +} + +func TestBuildUpstreamRequestAnthropicAPIKeyPassthrough_StripsContextManagementWhenClientHeaderMissingBeta(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil) + // 客户端仅带 oauth beta,不带 context-management-2025-06-27 + c.Request.Header.Set("Anthropic-Beta", "oauth-2025-04-20") + + body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`) + svc := &GatewayService{cfg: &config.Config{}} + req, err := svc.buildUpstreamRequestAnthropicAPIKeyPassthrough( + context.Background(), c, newAnthropicAPIKeyPassthroughAccountForBetaTest(), body, "token", + ) + require.NoError(t, err) + require.False(t, gjson.GetBytes(readUpstreamBodyForTest(t, req), "context_management").Exists(), + "API-key passthrough + 客户端未带 context-management beta → strip body 字段") +} + +func TestBuildUpstreamRequestAnthropicAPIKeyPassthrough_PreservesContextManagementWhenClientHeaderHasBeta(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil) + c.Request.Header.Set("Anthropic-Beta", "oauth-2025-04-20,context-management-2025-06-27") + + body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`) + svc := &GatewayService{cfg: &config.Config{}} + req, err := svc.buildUpstreamRequestAnthropicAPIKeyPassthrough( + context.Background(), c, newAnthropicAPIKeyPassthroughAccountForBetaTest(), body, "token", + ) + require.NoError(t, err) + require.True(t, gjson.GetBytes(readUpstreamBodyForTest(t, req), "context_management").Exists(), + "API-key passthrough + 客户端带 context-management beta → 字段保留(不过度删除)") +} + +func TestBuildCountTokensRequestAnthropicAPIKeyPassthrough_StripsContextManagementWhenClientHeaderMissingBeta(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", nil) + c.Request.Header.Set("Anthropic-Beta", "oauth-2025-04-20,token-counting-2024-11-01") + + body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[]},"messages":[]}`) + svc := &GatewayService{cfg: &config.Config{}} + req, err := svc.buildCountTokensRequestAnthropicAPIKeyPassthrough( + context.Background(), c, newAnthropicAPIKeyPassthroughAccountForBetaTest(), body, "token", + ) + require.NoError(t, err) + require.False(t, gjson.GetBytes(readUpstreamBodyForTest(t, req), "context_management").Exists(), + "count_tokens passthrough + 客户端未带 context-management beta → strip") +} + +// ============================================================================ +// 集成测试:buildUpstreamRequest +// 全路径验证上游 outgoing body 与 anthropic-beta header 严格对称。 +// 这个测试能挡住未来某人忘调 sanitize / 将 sanitize 挪到 CCH 之后 等 regression。 +// ============================================================================ + +func TestBuildUpstreamRequest_OAuthMimicHaiku_StripsContextManagementEndToEnd(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil) + + account := &Account{ID: 401, Platform: PlatformAnthropic, Type: AccountTypeOAuth, + Credentials: map[string]any{"access_token": "oauth-tok"}, + Status: StatusActive, + Schedulable: true, + } + // haiku + mimic CC → final beta = HaikuBetaHeader(不含 context-management)→ + // 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( + context.Background(), c, account, body, + "oauth-tok", "oauth", "claude-haiku-4-5", false, true, // mimicClaudeCode=true + ) + require.NoError(t, err) + + outBody := readUpstreamBodyForTest(t, req) + outBeta := getHeaderRaw(req.Header, "anthropic-beta") + + require.False(t, gjson.GetBytes(outBody, "context_management").Exists(), + "OAuth mimic + haiku 端到端:outgoing body 不应含 context_management") + require.False(t, anthropicBetaTokensContains(outBeta, claude.BetaContextManagement), + "对称约束:outgoing anthropic-beta header 也不带 context-management beta") +} + +func TestBuildUpstreamRequest_OAuthMimicNonHaiku_PreservesContextManagementEndToEnd(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil) + + account := &Account{ID: 402, Platform: PlatformAnthropic, Type: AccountTypeOAuth, + Credentials: map[string]any{"access_token": "oauth-tok"}, + Status: StatusActive, + Schedulable: true, + } + // sonnet + mimic CC → final beta = FullClaudeCodeMimicryBetas(含 context-management)→ + // 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( + context.Background(), c, account, body, + "oauth-tok", "oauth", "claude-sonnet-4-6", false, true, + ) + require.NoError(t, err) + + outBody := readUpstreamBodyForTest(t, req) + outBeta := getHeaderRaw(req.Header, "anthropic-beta") + + require.True(t, gjson.GetBytes(outBody, "context_management").Exists(), + "OAuth mimic + non-haiku:outgoing body 必须保留 context_management。") + require.True(t, anthropicBetaTokensContains(outBeta, claude.BetaContextManagement), + "对称约束:outgoing anthropic-beta header 同时含 context-management beta") +} + +func TestBuildUpstreamRequest_OAuthTransparentHaikuWithRealCCBeta_PreservesField(t *testing.T) { + // 端到端验证:真 CC 客户端 + haiku + 客户端 header 带 context-management beta + // → final beta 透传 → 不应该过度删除 body 字段 + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil) + c.Request.Header.Set("Anthropic-Beta", + "claude-code-20250219,oauth-2025-04-20,interleaved-thinking-2025-05-14,context-management-2025-06-27") + + account := &Account{ID: 403, Platform: PlatformAnthropic, Type: AccountTypeOAuth, + Credentials: map[string]any{"access_token": "oauth-tok"}, + Status: StatusActive, Schedulable: true, + } + 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( + context.Background(), c, account, body, + "oauth-tok", "oauth", "claude-haiku-4-5", false, false, // mimicClaudeCode=false(真 CC) + ) + require.NoError(t, err) + + outBody := readUpstreamBodyForTest(t, req) + outBeta := getHeaderRaw(req.Header, "anthropic-beta") + + require.True(t, anthropicBetaTokensContains(outBeta, claude.BetaContextManagement), + "真 CC 透传路径:客户端 header 中的 context-management beta 必须保留") + require.True(t, gjson.GetBytes(outBody, "context_management").Exists(), + "回归保护:真 CC + haiku + 客户端带 beta token 时,clear_thinking_20251015 功能不能静默失效") +} + +// CCH 顺序语义测试:sanitize 必须在 signBillingHeaderCCH 之前, +// 否则签名的 hash 与最终发送的 body 不一致,被 Anthropic 判 third-party。 +// +// 该测试不走 buildUpstreamRequest 完整路径(需要 mock SettingService 成本高), +// 而是直接验证两个顺序产生的 cch 不同,证明二者不可交换。 +// 测试名本身是语义约束的文档化 marker。 +func TestSanitizeMustBeBeforeCCHSigning_HashConsistency(t *testing.T) { + // 构造 body:含 context_management + cch=00000 占位符 + body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"system":[{"type":"text","text":"x-anthropic-billing-header: cc_version=2.1.92; cch=00000;"}],"messages":[]}`) + + // 最终发送场景:final beta 不含 context-management beta → sanitize 会 strip + finalBeta := "oauth-2025-04-20,interleaved-thinking-2025-05-14" + + extractCCH := func(t *testing.T, b []byte) string { + t.Helper() + m := regexp.MustCompile(`\bcch=([0-9a-fA-F]{5})\b`).FindSubmatch(b) + require.NotNil(t, m, "body 里找不到 cch=<5hex> :%s", string(b)) + return string(m[1]) + } + + // === 正确顺序:sanitize → signBillingHeaderCCH === + // 1. strip context_management + sanitizedFirst, changed := sanitizeAnthropicBodyForBetaTokens(body, finalBeta) + require.True(t, changed) + require.False(t, gjson.GetBytes(sanitizedFirst, "context_management").Exists()) + // 2. 基于“strip 后的 body”算 hash + correctFinal := signBillingHeaderCCH(sanitizedFirst) + correctCCH := extractCCH(t, correctFinal) + require.NotEqual(t, "00000", correctCCH, "placeholder 应被替换") + + // === 错误顺序:signBillingHeaderCCH → sanitize(未来 regression 场景)=== + // 1. 先基于“含 context_management 的 body”算 hash → cch=H_with + signedFirst := signBillingHeaderCCH(body) + wrongCCH := extractCCH(t, signedFirst) + require.NotEqual(t, "00000", wrongCCH) + // 2. 后 strip context_management → body 变化但 cch 仍是 H_with + wrongFinal, _ := sanitizeAnthropicBodyForBetaTokens(signedFirst, finalBeta) + wrongFinalCCH := extractCCH(t, wrongFinal) + + // === 关键断言 === + // 上游验证逻辑:将 outgoing body 的 cch 还原为 00000、重算 hash、与 cch 字段比较。 + // 模拟上游验证:用发送 body 算出“期望的 cch”,与发送 body 里的 cch 字段比。 + recomputeExpected := func(b []byte, currentCCH string) string { + t.Helper() + // 把 cch= 还原为 cch=00000 + re := regexp.MustCompile(`(\bcch=)` + currentCCH + `(\b)`) + restored := re.ReplaceAll(b, []byte("${1}00000${2}")) + return extractCCH(t, signBillingHeaderCCH(restored)) + } + + // 正确顺序:发送 body 的 cch == 重算 hash → 上游验证过 + require.Equal(t, correctCCH, recomputeExpected(correctFinal, correctCCH), + "正确顺序:final body 里的 cch 与重算 hash 一致 → 上游验证通过") + + // 错误顺序:发送 body 的 cch 是“含 ctx 算的”,但最终 body 不含 ctx → 重算 hash 不同 + require.NotEqual(t, wrongFinalCCH, recomputeExpected(wrongFinal, wrongFinalCCH), + "错误顺序:final body 里的 cch 是基于含 ctx 的 body 算的,"+ + "但发送 body 已 strip ctx → 上游重算 hash 与 cch 不一致 → 被判 third-party。"+ + "这是 buildUpstreamRequest / buildCountTokensRequest 里 sanitize 必须在 "+ + "signBillingHeaderCCH 之前的原因。") +} + +// count_tokens 主路径 E2E 集成测试 +func TestBuildCountTokensRequest_OAuthMimicHaiku_PreservesContextManagementEndToEnd(t *testing.T) { + // count_tokens 路径下 mimic 不按 haiku 排除,始终注入 BetaContextManagement + // → sanitize 看到最终 beta header 含 context-management beta → 字段保留。 + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", nil) + + account := &Account{ID: 411, Platform: PlatformAnthropic, Type: AccountTypeOAuth, + Credentials: map[string]any{"access_token": "oauth-tok"}, + Status: StatusActive, Schedulable: true, + } + body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`) + svc := &GatewayService{cfg: &config.Config{}} + req, err := svc.buildCountTokensRequest( + context.Background(), c, account, body, + "oauth-tok", "oauth", "claude-haiku-4-5", true, // mimicClaudeCode=true + ) + require.NoError(t, err) + + outBody := readUpstreamBodyForTest(t, req) + outBeta := getHeaderRaw(req.Header, "anthropic-beta") + + require.True(t, anthropicBetaTokensContains(outBeta, claude.BetaContextManagement), + "count_tokens mimic 始终注入 context-management beta") + require.True(t, gjson.GetBytes(outBody, "context_management").Exists(), + "对称约束:final beta 含 token 时 body 字段保留") + require.True(t, anthropicBetaTokensContains(outBeta, claude.BetaTokenCounting), + "count_tokens 路径必须含 token-counting beta") +} + +func TestBuildCountTokensRequest_APIKeyHaiku_StripsContextManagementEndToEnd(t *testing.T) { + // API-key + haiku + 客户端 header 不带 context-management beta → final beta 不含 → strip + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", nil) + c.Request.Header.Set("Anthropic-Beta", "interleaved-thinking-2025-05-14") + + account := &Account{ID: 412, Platform: PlatformAnthropic, Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-ant-xxx"}, + Status: StatusActive, Schedulable: true, + } + body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[]},"messages":[]}`) + svc := &GatewayService{cfg: &config.Config{}} + req, err := svc.buildCountTokensRequest( + context.Background(), c, account, body, + "sk-ant-xxx", "apikey", "claude-haiku-4-5", false, + ) + require.NoError(t, err) + + outBody := readUpstreamBodyForTest(t, req) + require.False(t, gjson.GetBytes(outBody, "context_management").Exists(), + "count_tokens API-key + 客户端未带 beta token → body strip") +} + +// count_tokens passthrough preserve 测试 +func TestBuildCountTokensRequestAnthropicAPIKeyPassthrough_PreservesContextManagementWhenClientHeaderHasBeta(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", nil) + c.Request.Header.Set("Anthropic-Beta", "oauth-2025-04-20,context-management-2025-06-27,token-counting-2024-11-01") + + body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`) + svc := &GatewayService{cfg: &config.Config{}} + req, err := svc.buildCountTokensRequestAnthropicAPIKeyPassthrough( + context.Background(), c, newAnthropicAPIKeyPassthroughAccountForBetaTest(), body, "token", + ) + require.NoError(t, err) + require.True(t, gjson.GetBytes(readUpstreamBodyForTest(t, req), "context_management").Exists(), + "count_tokens passthrough + 客户端带 context-management beta → 字段保留") +} + +func TestBuildUpstreamRequest_APIKeyHaikuWithContextManagement_StripsField(t *testing.T) { + // API-key + haiku + body 带 context_management + 客户端 header 未带 context-management beta + // → final beta 不含 → body 字段被 strip + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil) + c.Request.Header.Set("Anthropic-Beta", "interleaved-thinking-2025-05-14") + + account := &Account{ID: 404, Platform: PlatformAnthropic, Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-ant-xxx"}, + Status: StatusActive, Schedulable: true, + } + body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[]},"messages":[]}`) + svc := &GatewayService{cfg: &config.Config{}} + req, err := svc.buildUpstreamRequest( + context.Background(), c, account, body, + "sk-ant-xxx", "apikey", "claude-haiku-4-5", false, false, + ) + require.NoError(t, err) + + outBody := readUpstreamBodyForTest(t, req) + require.False(t, gjson.GetBytes(outBody, "context_management").Exists(), + "API-key + haiku + 客户端未带 beta token → body 字段必须被 strip") +} diff --git a/backend/internal/service/gateway_request.go b/backend/internal/service/gateway_request.go index 498336a4..91f7601c 100644 --- a/backend/internal/service/gateway_request.go +++ b/backend/internal/service/gateway_request.go @@ -12,6 +12,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/domain" "github.com/Wei-Shaw/sub2api/internal/pkg/antigravity" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/tidwall/gjson" "github.com/tidwall/sjson" ) @@ -665,6 +666,69 @@ func removeThinkingDependentContextStrategies(body []byte) []byte { return body } +// anthropicBetaContextManagementToken 是 context_management 字段受的 beta token。 +// 与 claude.BetaContextManagement 保持一致;在本文件本地定义以避免震荡 +// claude package 的该常量含义。 +const anthropicBetaContextManagementToken = "context-management-2025-06-27" + +// sanitizeAnthropicBodyForBetaTokens 是对 Anthropic 直连路径上 body↔beta header +// **能力维度**对称约束的统一实现,与 Bedrock 路径的 +// `sanitizeBedrockFieldsForBetaTokens` 对称。 +// +// 问题场景: +// - context_management 是 Claude Code CLI 2.1.87+ 默认携带的 beta 字段 +// (含 clear_thinking_20251015 等清理策略) +// - 其被 Anthropic 上游接受的前提是 anthropic-beta header 含 +// `context-management-2025-06-27` +// - 若两侧不一致上游 Pydantic schema 拒收: +// "context_management: Extra inputs are not permitted" +// +// 本函数按最终发送的 anthropic-beta header 决定是否保留 body 中的 +// context_management 字段:缺 beta token → strip。这将限制完全建立在 +// "能力维度" 上,与 model 名 / token type / mimicry 子路径无关。 +// +// 调用约束:必须在 CCH 签名之前调用,否则签名 hash 与最终 body +// 不一致,上游会以 third-party 拒收。 +// +// 返回 (sanitized, changed):changed 表示是否发生实际删除,供调用方决定 +// 是否重用原 body 引用。 +func sanitizeAnthropicBodyForBetaTokens(body []byte, anthropicBetaHeader string) ([]byte, bool) { + if len(body) == 0 { + return body, false + } + if !gjson.GetBytes(body, "context_management").Exists() { + return body, false + } + if anthropicBetaTokensContains(anthropicBetaHeader, anthropicBetaContextManagementToken) { + return body, false + } + if b, err := sjson.DeleteBytes(body, "context_management"); err == nil { + return b, true + } else { + // 不应发生:gjson 刚验证过字段存在 + body 是合法 JSON。如果 sjson 仍报错, + // 调用方会拿到 (body, false),但此前 computeFinalAnthropicBeta 已按“strip 后” + // 计算了 finalBeta——两侧会不一致。记录 warning 最小限度提醒运维。 + logger.LegacyPrintf("service.gateway", + "[CtxMgmtSanitize] sjson.DeleteBytes failed unexpectedly: %v (body len=%d). "+ + "body and final anthropic-beta header may be out of sync.", err, len(body)) + } + return body, false +} + +// anthropicBetaTokensContains 检测逗号分隔的 anthropic-beta header 是否含指定 token。 +// 宋体空格宽容;区分大小写(Anthropic beta token 始终是小写)。 +func anthropicBetaTokensContains(header, token string) bool { + if header == "" || token == "" { + return false + } + for _, part := range strings.Split(header, ",") { + if strings.TrimSpace(part) == token { + return true + } + } + return false +} + // FilterSignatureSensitiveBlocksForRetry is a stronger retry filter for cases where upstream errors indicate // signature/thought_signature validation issues involving tool blocks. // diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index 4a8175a4..7c48f243 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -1155,6 +1155,12 @@ func normalizeClaudeOAuthRequestBody(body []byte, modelID string, opts claudeOAu // context_management:thinking.type 为 enabled/adaptive 时,真实 CLI 会自动 // 附带 {"edits":[{"type":"clear_thinking_20251015","keep":"all"}]}。 // 客户端显式传了就透传;否则按 CLI 行为补齐。 + // + // 注:本函数不按 model 名决定是否保留 context_management。“最终 beta + // header 不含 context-management-2025-06-27 时 strip 字段”的能力维度 + // 对称约束由 sanitizeAnthropicBodyForBetaTokens 在 buildUpstreamRequest / + // buildCountTokensRequest 层统一执行,与 Bedrock 路径的 + // sanitizeBedrockFieldsForBetaTokens 对称。 if !gjson.GetBytes(out, "context_management").Exists() { thinkingType := gjson.GetBytes(out, "thinking.type").String() if thinkingType == "enabled" || thinkingType == "adaptive" { @@ -5248,6 +5254,17 @@ func (s *GatewayService) buildUpstreamRequestAnthropicAPIKeyPassthrough( targetURL = validatedURL + "/v1/messages?beta=true" } + // 能力维度 body sanitize:透传路径上 anthropic-beta header 原样透传客户端值, + // 依此决定是否保留 body 中的 context_management。避免“客户端 body 带字段但 + // header 忘记带 beta token”的客户端 bug 在透传场景下让上游 400。 + clientBeta := "" + if c != nil && c.Request != nil { + clientBeta = getHeaderRaw(c.Request.Header, "anthropic-beta") + } + if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, clientBeta); changed { + body = sanitized + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body)) if err != nil { return nil, err @@ -6106,6 +6123,29 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex if fingerprint != nil { body = syncBillingHeaderVersion(body, fingerprint.UserAgent) } + + // === 计算最终 anthropic-beta header(先于 body sanitize 与 CCH 签名)=== + // + // 顺序约束: + // 1) 算 finalBeta(纯函数,不依赖 req.Header;mimicry 路径会忽略客户端 beta, + // 与原“OAuth + mimicClaudeCode 跳过白名单透传”行为对齐) + // 2) 按 finalBeta 做能力维度 body sanitize(如 context-management beta 缺失 → + // strip body.context_management,与 Bedrock 路径对称) + // 3) CCH 签名(必须使用 strip 后的 body,否则 hash 与最终 body 不一致 → + // 被 Anthropic 判 third-party) + // 4) NewRequest(body 至此最终敲定) + // 5) 透传白名单 / fingerprint / mimic header / 写入 finalBeta + policyFilterSet := s.getBetaPolicyFilterSet(ctx, c, account, modelID) + effectiveDropSet := mergeDropSets(policyFilterSet) + finalBetaHeader, finalBetaShouldSet := s.computeFinalAnthropicBeta( + tokenType, mimicClaudeCode, modelID, clientHeaders, body, effectiveDropSet, + ) + + // 能力维度 body sanitize:与最终 anthropic-beta header 对称 + if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, finalBetaHeader); changed { + body = sanitized + } + // CCH 签名:将 cch=00000 占位符替换为 xxHash64 签名(需在所有 body 修改之后) if enableCCH { body = signBillingHeaderCCH(body) @@ -6156,46 +6196,18 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex applyClaudeOAuthHeaderDefaults(req) } - // Build effective drop set: merge static defaults with dynamic beta policy filter rules - policyFilterSet := s.getBetaPolicyFilterSet(ctx, c, account, modelID) - effectiveDropSet := mergeDropSets(policyFilterSet) + // OAuth + mimic Claude Code:强制注入 CLI 指纹相关 header + // (user-agent/x-stainless-*/x-app/Accept/x-stainless-helper-method/x-client-request-id) + if tokenType == "oauth" && mimicClaudeCode { + applyClaudeCodeMimicHeaders(req, reqStream) + } - // 处理 anthropic-beta header(OAuth 账号需要包含 oauth beta) - if tokenType == "oauth" { - if mimicClaudeCode { - // 非 Claude Code 客户端:按 opencode 的策略处理: - // - 强制 Claude Code 指纹相关请求头(尤其是 user-agent/x-stainless/x-app) - // - 保留 incoming beta 的同时,确保 OAuth 所需 beta 存在 - applyClaudeCodeMimicHeaders(req, reqStream) - - incomingBeta := getHeaderRaw(req.Header, "anthropic-beta") - // Claude Code OAuth credentials are scoped to Claude Code. - // Non-haiku models MUST include claude-code beta for Anthropic to recognize - // this as a legitimate Claude Code request; without it, the request is - // rejected as third-party ("out of extra usage"). - // Haiku models are exempt from third-party detection and don't need it. - requiredBetas := []string{claude.BetaOAuth, claude.BetaInterleavedThinking} - if !strings.Contains(strings.ToLower(modelID), "haiku") { - requiredBetas = claude.FullClaudeCodeMimicryBetas() - } - setHeaderRaw(req.Header, "anthropic-beta", mergeAnthropicBetaDropping(requiredBetas, incomingBeta, effectiveDropSet)) - } else { - // Claude Code 客户端:尽量透传原始 header,仅补齐 oauth beta - clientBetaHeader := getHeaderRaw(req.Header, "anthropic-beta") - setHeaderRaw(req.Header, "anthropic-beta", stripBetaTokensWithSet(s.getBetaHeader(modelID, clientBetaHeader), effectiveDropSet)) - } - } else { - // API-key accounts: apply beta policy filter to strip controlled tokens - if existingBeta := getHeaderRaw(req.Header, "anthropic-beta"); existingBeta != "" { - setHeaderRaw(req.Header, "anthropic-beta", stripBetaTokensWithSet(existingBeta, effectiveDropSet)) - } else if s.cfg != nil && s.cfg.Gateway.InjectBetaForAPIKey { - // API-key:仅在请求显式使用 beta 特性且客户端未提供时,按需补齐(默认关闭) - if requestNeedsBetaFeatures(body) { - if beta := defaultAPIKeyBetaHeader(body); beta != "" { - setHeaderRaw(req.Header, "anthropic-beta", beta) - } - } - } + // 写入最终 anthropic-beta header + // 注:透传分支白名单可能写入了客户端 anthropic-beta,无条件 Del 一次再按 finalBeta + // 决定是否 set,确保 dropSet 过滤后的结果一定覆盖客户端原始值。 + deleteHeaderAllForms(req.Header, "anthropic-beta") + if finalBetaShouldSet { + setHeaderRaw(req.Header, "anthropic-beta", finalBetaHeader) } // 同步 X-Claude-Code-Session-Id 头:取 body 中已处理的 metadata.user_id 的 session_id 覆盖 @@ -6242,6 +6254,16 @@ func (s *GatewayService) buildUpstreamRequestAnthropicVertex( if err != nil { return nil, err } + + // 能力维度 sanitize:Vertex 路径上 anthropic-beta header 原样透传客户端值 + // (下面白名单跳过 anthropic-version 但保留 anthropic-beta),依此决定是否 + // 保留 body 中的 context_management,与 Anthropic 直连 / Bedrock 路径对称。 + if c != nil && c.Request != nil { + clientBeta := getHeaderRaw(c.Request.Header, "anthropic-beta") + if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(vertexBody, clientBeta); changed { + vertexBody = sanitized + } + } fullURL, err := buildVertexAnthropicURL(account.VertexProjectID(), account.VertexLocation(modelID), modelID, reqStream) if err != nil { return nil, err @@ -6410,6 +6432,121 @@ func mergeAnthropicBetaDropping(required []string, incoming string, drop map[str return strings.Join(out, ",") } +// computeFinalAnthropicBeta 计算发往上游的最终 anthropic-beta header 值。 +// +// 设计动机:将原本在 buildUpstreamRequest 内联在一起、依赖 req.Header 的 +// anthropic-beta 计算逻辑抽成纯函数。这样调用方可以在 NewRequest 之前 +// 就提前拿到最终 beta header,进而能按它对 body 做能力维度 sanitize 后再做 +// CCH 签名——一举修复了以下之前由顺序依赖导致的能力维度 sanitize +// 无法部署的问题(签名与最终 body 不一致可以被判 third-party)。 +// +// 返回 (value, shouldSet): +// - shouldSet=false 意为“不主动设置 anthropic-beta header”,与原代码“ +// API-key 账号 + 客户端未传 anthropic-beta + InjectBetaForAPIKey 未开启或 +// requestNeedsBetaFeatures=false”的行为对齐。 +// - shouldSet=true 时 value 可能为空字符串(例如客户端透传的 beta 被 dropSet +// 全部过滤掉),这与原代码中 setHeaderRaw 的结果一致。 +// +// clientHeaders 是客户端原始 HTTP header(通常为 c.Request.Header);nil 时按“客户端 +// 未传”处理。body 是已经 metadata 重写 / billing version sync 之后但未 sanitize 上游 +// 不兼容字段之前的版本。 +func (s *GatewayService) computeFinalAnthropicBeta( + tokenType string, + mimicClaudeCode bool, + modelID string, + clientHeaders http.Header, + body []byte, + effectiveDropSet map[string]struct{}, +) (string, bool) { + clientBeta := "" + if clientHeaders != nil { + clientBeta = getHeaderRaw(clientHeaders, "anthropic-beta") + } + + if tokenType == "oauth" { + if mimicClaudeCode { + // mimic 路径:原代码跳过白名单透传,incomingBeta 总是空字符串。 + // 这里传空 string 以严格对齐原行为。 + requiredBetas := []string{claude.BetaOAuth, claude.BetaInterleavedThinking} + if !strings.Contains(strings.ToLower(modelID), "haiku") { + requiredBetas = claude.FullClaudeCodeMimicryBetas() + } + return mergeAnthropicBetaDropping(requiredBetas, "", effectiveDropSet), true + } + // 真 Claude Code 客户端透传路径 + return stripBetaTokensWithSet(s.getBetaHeader(modelID, clientBeta), effectiveDropSet), true + } + + // API-key accounts + if clientBeta != "" { + return stripBetaTokensWithSet(clientBeta, effectiveDropSet), true + } + if s.cfg != nil && s.cfg.Gateway.InjectBetaForAPIKey { + if requestNeedsBetaFeatures(body) { + if beta := defaultAPIKeyBetaHeader(body); beta != "" { + return beta, true + } + } + } + return "", false +} + +// computeFinalCountTokensAnthropicBeta 是 count_tokens 路径上 anthropic-beta header 的 +// 计算纯函数。语义与 computeFinalAnthropicBeta 对齐,但备份了 count_tokens 独有的 +// 两条特殊规则: +// +// - OAuth mimic:requiredBetas 为 FullClaudeCodeMimicryBetas + BetaTokenCounting +// (与 messages 不同的是:不按 haiku 排除;count_tokens 始终携带 token-counting beta) +// - OAuth 透传 + 客户端未传 anthropic-beta:补齐 CountTokensBetaHeader +// - OAuth 透传 + 客户端传了:补齐 BetaTokenCounting(如果未含) +// +// 返回语义同 computeFinalAnthropicBeta。 +func (s *GatewayService) computeFinalCountTokensAnthropicBeta( + tokenType string, + mimicClaudeCode bool, + modelID string, + clientHeaders http.Header, + body []byte, + effectiveDropSet map[string]struct{}, +) (string, bool) { + clientBeta := "" + if clientHeaders != nil { + clientBeta = getHeaderRaw(clientHeaders, "anthropic-beta") + } + + if tokenType == "oauth" { + if mimicClaudeCode { + // 与原代码严格等价:original buildCountTokensRequest 在 count_tokens mimic + // 分支上**不**会跳过白名单透传(与 messages mimic 路径不同),所以 + // incomingBeta = req.Header[anthropic-beta] = 客户端透传过来的 client beta。 + // 重构后直接从 clientHeaders 拿同一个值,保持行为一致。 + requiredBetas := append(claude.FullClaudeCodeMimicryBetas(), claude.BetaTokenCounting) + return mergeAnthropicBetaDropping(requiredBetas, clientBeta, effectiveDropSet), true + } + if clientBeta == "" { + return claude.CountTokensBetaHeader, true + } + beta := s.getBetaHeader(modelID, clientBeta) + if !strings.Contains(beta, claude.BetaTokenCounting) { + beta = beta + "," + claude.BetaTokenCounting + } + return stripBetaTokensWithSet(beta, effectiveDropSet), true + } + + // API-key accounts + if clientBeta != "" { + return stripBetaTokensWithSet(clientBeta, effectiveDropSet), true + } + if s.cfg != nil && s.cfg.Gateway.InjectBetaForAPIKey { + if requestNeedsBetaFeatures(body) { + if beta := defaultAPIKeyBetaHeader(body); beta != "" { + return beta, true + } + } + } + return "", false +} + // stripBetaTokens removes the given beta tokens from a comma-separated header value. func stripBetaTokens(header string, tokens []string) string { if header == "" || len(tokens) == 0 { @@ -9312,6 +9449,15 @@ func (s *GatewayService) buildCountTokensRequestAnthropicAPIKeyPassthrough( targetURL = validatedURL + "/v1/messages/count_tokens?beta=true" } + // 同 buildUpstreamRequestAnthropicAPIKeyPassthrough:能力维度 sanitize。 + clientBeta := "" + if c != nil && c.Request != nil { + clientBeta = getHeaderRaw(c.Request.Header, "anthropic-beta") + } + if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, clientBeta); changed { + body = sanitized + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body)) if err != nil { return nil, err @@ -9402,6 +9548,19 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con if ctFingerprint != nil && ctEnableFP { body = syncBillingHeaderVersion(body, ctFingerprint.UserAgent) } + + // === 计算最终 anthropic-beta header(先于 body sanitize 与 CCH 签名)=== + // 顺序约束同 buildUpstreamRequest。 + ctEffectiveDropSet := mergeDropSets(s.getBetaPolicyFilterSet(ctx, c, account, modelID)) + finalBetaHeader, finalBetaShouldSet := s.computeFinalCountTokensAnthropicBeta( + tokenType, mimicClaudeCode, modelID, clientHeaders, body, ctEffectiveDropSet, + ) + + // 能力维度 body sanitize:与最终 anthropic-beta header 对称 + if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, finalBetaHeader); changed { + body = sanitized + } + if ctEnableCCH { body = signBillingHeaderCCH(body) } @@ -9445,41 +9604,15 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con applyClaudeOAuthHeaderDefaults(req) } - // Build effective drop set for count_tokens: merge static defaults with dynamic beta policy filter rules - ctEffectiveDropSet := mergeDropSets(s.getBetaPolicyFilterSet(ctx, c, account, modelID)) + // OAuth + mimic Claude Code:强制注入 CLI 指纹 header + if tokenType == "oauth" && mimicClaudeCode { + applyClaudeCodeMimicHeaders(req, false) + } - // OAuth 账号:处理 anthropic-beta header - if tokenType == "oauth" { - if mimicClaudeCode { - applyClaudeCodeMimicHeaders(req, false) - - incomingBeta := getHeaderRaw(req.Header, "anthropic-beta") - requiredBetas := append(claude.FullClaudeCodeMimicryBetas(), claude.BetaTokenCounting) - setHeaderRaw(req.Header, "anthropic-beta", mergeAnthropicBetaDropping(requiredBetas, incomingBeta, ctEffectiveDropSet)) - } else { - clientBetaHeader := getHeaderRaw(req.Header, "anthropic-beta") - if clientBetaHeader == "" { - setHeaderRaw(req.Header, "anthropic-beta", claude.CountTokensBetaHeader) - } else { - beta := s.getBetaHeader(modelID, clientBetaHeader) - if !strings.Contains(beta, claude.BetaTokenCounting) { - beta = beta + "," + claude.BetaTokenCounting - } - setHeaderRaw(req.Header, "anthropic-beta", stripBetaTokensWithSet(beta, ctEffectiveDropSet)) - } - } - } else { - // API-key accounts: apply beta policy filter to strip controlled tokens - if existingBeta := getHeaderRaw(req.Header, "anthropic-beta"); existingBeta != "" { - setHeaderRaw(req.Header, "anthropic-beta", stripBetaTokensWithSet(existingBeta, ctEffectiveDropSet)) - } else if s.cfg != nil && s.cfg.Gateway.InjectBetaForAPIKey { - // API-key:与 messages 同步的按需 beta 注入(默认关闭) - if requestNeedsBetaFeatures(body) { - if beta := defaultAPIKeyBetaHeader(body); beta != "" { - setHeaderRaw(req.Header, "anthropic-beta", beta) - } - } - } + // 写入最终 anthropic-beta header(Del 一次避免白名单透传值残留) + deleteHeaderAllForms(req.Header, "anthropic-beta") + if finalBetaShouldSet { + setHeaderRaw(req.Header, "anthropic-beta", finalBetaHeader) } // 同步 X-Claude-Code-Session-Id 头:取 body 中已处理的 metadata.user_id 的 session_id 覆盖 diff --git a/backend/internal/service/header_util.go b/backend/internal/service/header_util.go index 1091070d..f8da068d 100644 --- a/backend/internal/service/header_util.go +++ b/backend/internal/service/header_util.go @@ -109,6 +109,20 @@ func addHeaderRaw(h http.Header, key, value string) { h[key] = append(h[key], value) } +// deleteHeaderAllForms removes a header in all common key forms (raw, wire casing, +// canonical) so subsequent setHeaderRaw will not coexist with a passthrough value +// written under a different casing. +func deleteHeaderAllForms(h http.Header, key string) { + if h == nil || key == "" { + return + } + h.Del(key) // canonical + delete(h, key) + if wk := resolveWireCasing(key); wk != key { + delete(h, wk) + } +} + // getHeaderRaw reads a header value, trying multiple key forms to handle the mismatch // between Go canonical keys, wire casing keys, and raw keys: // 1. exact key as provided From 20f5340784484d8be82c2139bb40ff34f1a7b715 Mon Sep 17 00:00:00 2001 From: JIA-ss <627723154@qq.com> Date: Thu, 28 May 2026 00:38:25 +0800 Subject: [PATCH 17/79] =?UTF-8?q?fix(apicompat):=20Responses=E2=86=92Chat?= =?UTF-8?q?=20=E8=BD=AC=E6=8D=A2=E8=A1=A5=E9=BD=90=20completion=5Ftokens?= =?UTF-8?q?=5Fdetails=20=E9=80=8F=E4=BC=A0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit OpenAI Responses API 在 gpt-5.x 等 reasoning 模型上会返回 output_tokens_details.reasoning_tokens, 但 ResponsesToChatCompletions 只映射了 input_tokens_details.cached_tokens, 导致客户端拿到的 chat.completion.usage 中 completion_tokens 出现无法解释的波动 (短 prompt 也可能 30+ token), 且缺失 reasoning_tokens 细分字段, 难以与 OpenAI 原生 Chat Completions 响应对账。 按 OpenAI 官方 CompletionUsage schema (openai/openai-go SDK completion.go) 补齐所有 token-details 字段, 全部 omitempty: prompt_tokens_details: - cached_tokens (原已支持) - audio_tokens (新增) completion_tokens_details: - reasoning_tokens (新增) - audio_tokens (新增) - accepted_prediction_tokens (新增) - rejected_prediction_tokens (新增) 实现细节: - 抽出 promptDetailsFromResponses / completionDetailsFromResponses 两个 helper, 全零字段返回 nil - 非流路径 ResponsesToChatCompletions 复用已存在的 chatUsageFromResponsesUsage helper, 消除两条路径间的重复 - 非 reasoning / 非 audio 上游 (Anthropic, Gemini, gpt-4o) 不填这些 字段, helper 返回 nil → CompletionTokensDetails 不输出, 对现有响应 字节级兼容 新增单测: - TestResponsesToChatCompletions_ReasoningTokens - TestResponsesToChatCompletions_AllTokenDetailsPassThrough - TestResponsesToChatCompletions_NoReasoningTokensWhenZero - TestResponsesEventToChatChunks_CompletedWithReasoningTokens --- .../chatcompletions_responses_test.go | 135 ++++++++++++++++++ .../apicompat/responses_to_chatcompletions.go | 58 +++++--- backend/internal/pkg/apicompat/types.go | 30 +++- 3 files changed, 198 insertions(+), 25 deletions(-) diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_test.go b/backend/internal/pkg/apicompat/chatcompletions_responses_test.go index 016c2415..b03b012f 100644 --- a/backend/internal/pkg/apicompat/chatcompletions_responses_test.go +++ b/backend/internal/pkg/apicompat/chatcompletions_responses_test.go @@ -663,6 +663,115 @@ func TestResponsesToChatCompletions_CachedTokens(t *testing.T) { assert.Equal(t, 80, chat.Usage.PromptTokensDetails.CachedTokens) } +func TestResponsesToChatCompletions_ReasoningTokens(t *testing.T) { + resp := &ResponsesResponse{ + ID: "resp_reasoning", + Status: "completed", + Output: []ResponsesOutput{ + { + Type: "message", + Content: []ResponsesContentPart{{Type: "output_text", Text: "ping"}}, + }, + }, + Usage: &ResponsesUsage{ + InputTokens: 24, + OutputTokens: 33, + TotalTokens: 57, + OutputTokensDetails: &ResponsesOutputTokensDetails{ + ReasoningTokens: 32, + }, + }, + } + + chat := ResponsesToChatCompletions(resp, "gpt-5.5") + require.NotNil(t, chat.Usage) + assert.Equal(t, 33, chat.Usage.CompletionTokens) + require.NotNil(t, chat.Usage.CompletionTokensDetails) + assert.Equal(t, 32, chat.Usage.CompletionTokensDetails.ReasoningTokens) +} + +func TestResponsesToChatCompletions_AllTokenDetailsPassThrough(t *testing.T) { + // Covers the full OpenAI CompletionUsage detail field set so future audio + // and prediction-outputs responses propagate without further changes. + resp := &ResponsesResponse{ + ID: "resp_full_details", + Status: "completed", + Output: []ResponsesOutput{ + { + Type: "message", + Content: []ResponsesContentPart{{Type: "output_text", Text: "x"}}, + }, + }, + Usage: &ResponsesUsage{ + InputTokens: 100, + OutputTokens: 50, + TotalTokens: 150, + InputTokensDetails: &ResponsesInputTokensDetails{ + CachedTokens: 60, + AudioTokens: 4, + }, + OutputTokensDetails: &ResponsesOutputTokensDetails{ + ReasoningTokens: 30, + AudioTokens: 2, + AcceptedPredictionTokens: 10, + RejectedPredictionTokens: 3, + }, + }, + } + + chat := ResponsesToChatCompletions(resp, "gpt-5.5") + require.NotNil(t, chat.Usage) + require.NotNil(t, chat.Usage.PromptTokensDetails) + assert.Equal(t, 60, chat.Usage.PromptTokensDetails.CachedTokens) + assert.Equal(t, 4, chat.Usage.PromptTokensDetails.AudioTokens) + + require.NotNil(t, chat.Usage.CompletionTokensDetails) + assert.Equal(t, 30, chat.Usage.CompletionTokensDetails.ReasoningTokens) + assert.Equal(t, 2, chat.Usage.CompletionTokensDetails.AudioTokens) + assert.Equal(t, 10, chat.Usage.CompletionTokensDetails.AcceptedPredictionTokens) + assert.Equal(t, 3, chat.Usage.CompletionTokensDetails.RejectedPredictionTokens) + + raw, err := json.Marshal(chat.Usage) + require.NoError(t, err) + assert.Contains(t, string(raw), `"prompt_tokens_details"`) + assert.Contains(t, string(raw), `"completion_tokens_details"`) + assert.Contains(t, string(raw), `"reasoning_tokens":30`) + assert.Contains(t, string(raw), `"accepted_prediction_tokens":10`) +} + +func TestResponsesToChatCompletions_NoReasoningTokensWhenZero(t *testing.T) { + // Non-reasoning models do not return reasoning_tokens. The mapping must + // omit completion_tokens_details entirely rather than emitting a zero-valued + // field, so non-reasoning responses stay clean. + resp := &ResponsesResponse{ + ID: "resp_no_reasoning", + Status: "completed", + Output: []ResponsesOutput{ + { + Type: "message", + Content: []ResponsesContentPart{{Type: "output_text", Text: "hi"}}, + }, + }, + Usage: &ResponsesUsage{ + InputTokens: 10, + OutputTokens: 5, + TotalTokens: 15, + OutputTokensDetails: &ResponsesOutputTokensDetails{ + ReasoningTokens: 0, + }, + }, + } + + chat := ResponsesToChatCompletions(resp, "gpt-4o") + require.NotNil(t, chat.Usage) + assert.Nil(t, chat.Usage.CompletionTokensDetails) + + raw, err := json.Marshal(chat.Usage) + require.NoError(t, err) + assert.NotContains(t, string(raw), "completion_tokens_details") + assert.NotContains(t, string(raw), "reasoning_tokens") +} + func TestResponsesToChatCompletions_WebSearch(t *testing.T) { resp := &ResponsesResponse{ ID: "resp_ws", @@ -825,6 +934,32 @@ func TestResponsesEventToChatChunks_Completed(t *testing.T) { assert.Equal(t, 30, chunks[1].Usage.PromptTokensDetails.CachedTokens) } +func TestResponsesEventToChatChunks_CompletedWithReasoningTokens(t *testing.T) { + state := NewResponsesEventToChatState() + state.Model = "gpt-5.5" + state.IncludeUsage = true + + chunks := ResponsesEventToChatChunks(&ResponsesStreamEvent{ + Type: "response.completed", + Response: &ResponsesResponse{ + Status: "completed", + Usage: &ResponsesUsage{ + InputTokens: 24, + OutputTokens: 33, + TotalTokens: 57, + OutputTokensDetails: &ResponsesOutputTokensDetails{ + ReasoningTokens: 32, + }, + }, + }, + }, state) + require.Len(t, chunks, 2) + + require.NotNil(t, chunks[1].Usage) + require.NotNil(t, chunks[1].Usage.CompletionTokensDetails) + assert.Equal(t, 32, chunks[1].Usage.CompletionTokensDetails.ReasoningTokens) +} + func TestResponsesEventToChatChunks_ResponseDone(t *testing.T) { state := NewResponsesEventToChatState() state.Model = "gpt-4o" diff --git a/backend/internal/pkg/apicompat/responses_to_chatcompletions.go b/backend/internal/pkg/apicompat/responses_to_chatcompletions.go index 7e8354ee..8809b4fc 100644 --- a/backend/internal/pkg/apicompat/responses_to_chatcompletions.go +++ b/backend/internal/pkg/apicompat/responses_to_chatcompletions.go @@ -81,19 +81,7 @@ func ResponsesToChatCompletions(resp *ResponsesResponse, model string) *ChatComp FinishReason: finishReason, }} - if resp.Usage != nil { - usage := &ChatUsage{ - PromptTokens: resp.Usage.InputTokens, - CompletionTokens: resp.Usage.OutputTokens, - TotalTokens: resp.Usage.InputTokens + resp.Usage.OutputTokens, - } - if resp.Usage.InputTokensDetails != nil && resp.Usage.InputTokensDetails.CachedTokens > 0 { - usage.PromptTokensDetails = &ChatTokenDetails{ - CachedTokens: resp.Usage.InputTokensDetails.CachedTokens, - } - } - out.Usage = usage - } + out.Usage = chatUsageFromResponsesUsage(resp.Usage) return out } @@ -341,14 +329,48 @@ func chatUsageFromResponsesUsage(u *ResponsesUsage) *ChatUsage { CompletionTokens: u.OutputTokens, TotalTokens: u.InputTokens + u.OutputTokens, } - if u.InputTokensDetails != nil && u.InputTokensDetails.CachedTokens > 0 { - usage.PromptTokensDetails = &ChatTokenDetails{ - CachedTokens: u.InputTokensDetails.CachedTokens, - } - } + usage.PromptTokensDetails = promptDetailsFromResponses(u.InputTokensDetails) + usage.CompletionTokensDetails = completionDetailsFromResponses(u.OutputTokensDetails) return usage } +// promptDetailsFromResponses maps Responses-API input_tokens_details into a +// Chat-Completions prompt_tokens_details. Returns nil when nothing would be +// emitted, so upstreams that do not break down prompt usage stay clean. +func promptDetailsFromResponses(src *ResponsesInputTokensDetails) *ChatTokenDetails { + if src == nil { + return nil + } + if src.CachedTokens == 0 && src.AudioTokens == 0 { + return nil + } + return &ChatTokenDetails{ + CachedTokens: src.CachedTokens, + AudioTokens: src.AudioTokens, + } +} + +// completionDetailsFromResponses maps Responses-API output_tokens_details +// into a Chat-Completions completion_tokens_details. Mirrors the OpenAI +// official CompletionUsage schema: reasoning_tokens, audio_tokens, and +// the predicted-outputs accepted/rejected counts. Returns nil when nothing +// would be emitted so non-reasoning, non-audio responses stay clean. +func completionDetailsFromResponses(src *ResponsesOutputTokensDetails) *ChatTokenDetails { + if src == nil { + return nil + } + if src.ReasoningTokens == 0 && src.AudioTokens == 0 && + src.AcceptedPredictionTokens == 0 && src.RejectedPredictionTokens == 0 { + return nil + } + return &ChatTokenDetails{ + ReasoningTokens: src.ReasoningTokens, + AudioTokens: src.AudioTokens, + AcceptedPredictionTokens: src.AcceptedPredictionTokens, + RejectedPredictionTokens: src.RejectedPredictionTokens, + } +} + func makeChatDeltaChunk(state *ResponsesEventToChatState, delta ChatDelta) ChatCompletionsChunk { return ChatCompletionsChunk{ ID: state.ID, diff --git a/backend/internal/pkg/apicompat/types.go b/backend/internal/pkg/apicompat/types.go index 8b576647..b4451f23 100644 --- a/backend/internal/pkg/apicompat/types.go +++ b/backend/internal/pkg/apicompat/types.go @@ -362,11 +362,15 @@ func (u *ResponsesUsage) UnmarshalJSON(data []byte) error { // ResponsesInputTokensDetails breaks down input token usage. type ResponsesInputTokensDetails struct { CachedTokens int `json:"cached_tokens,omitempty"` + AudioTokens int `json:"audio_tokens,omitempty"` } // ResponsesOutputTokensDetails breaks down output token usage. type ResponsesOutputTokensDetails struct { - ReasoningTokens int `json:"reasoning_tokens,omitempty"` + ReasoningTokens int `json:"reasoning_tokens,omitempty"` + AudioTokens int `json:"audio_tokens,omitempty"` + AcceptedPredictionTokens int `json:"accepted_prediction_tokens,omitempty"` + RejectedPredictionTokens int `json:"rejected_prediction_tokens,omitempty"` } // --------------------------------------------------------------------------- @@ -517,15 +521,27 @@ type ChatChoice struct { // ChatUsage holds token counts in Chat Completions format. type ChatUsage struct { - PromptTokens int `json:"prompt_tokens"` - CompletionTokens int `json:"completion_tokens"` - TotalTokens int `json:"total_tokens"` - PromptTokensDetails *ChatTokenDetails `json:"prompt_tokens_details,omitempty"` + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + TotalTokens int `json:"total_tokens"` + PromptTokensDetails *ChatTokenDetails `json:"prompt_tokens_details,omitempty"` + CompletionTokensDetails *ChatTokenDetails `json:"completion_tokens_details,omitempty"` } -// ChatTokenDetails provides a breakdown of token usage. +// ChatTokenDetails provides a breakdown of token usage. The same type is +// reused for both prompt_tokens_details and completion_tokens_details; +// unset fields are omitted so each side only emits the fields that apply. +// +// Field set mirrors OpenAI's official CompletionUsage schema: +// - prompt_tokens_details: cached_tokens, audio_tokens +// - completion_tokens_details: reasoning_tokens, audio_tokens, +// accepted_prediction_tokens, rejected_prediction_tokens type ChatTokenDetails struct { - CachedTokens int `json:"cached_tokens,omitempty"` + CachedTokens int `json:"cached_tokens,omitempty"` + AudioTokens int `json:"audio_tokens,omitempty"` + ReasoningTokens int `json:"reasoning_tokens,omitempty"` + AcceptedPredictionTokens int `json:"accepted_prediction_tokens,omitempty"` + RejectedPredictionTokens int `json:"rejected_prediction_tokens,omitempty"` } // ChatCompletionsChunk is a single streaming chunk from POST /v1/chat/completions. From d7bed40dda4a5bcf7cb969d73c221829de02a91d Mon Sep 17 00:00:00 2001 From: siyuan <740665504@qq.com> Date: Thu, 28 May 2026 01:27:11 +0800 Subject: [PATCH 18/79] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=20OpenAI=20WS=20?= =?UTF-8?q?=E5=85=BC=E5=AE=B9=E6=80=A7=E4=B8=8E=20usage=20=E7=BB=9F?= =?UTF-8?q?=E8=AE=A1?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 对齐 WS 与流式终态 usage 解析,补齐 failed/done/incomplete/cancelled 等事件 - 兼容后续 WS response.create 省略 model,保持模型映射与权限判断一致 - 补齐 passthrough header 透传和图片 usage 字段映射 --- .../service/openai_gateway_service.go | 2 +- .../service/openai_gateway_service_test.go | 6 + .../internal/service/openai_ws_forwarder.go | 39 ++- ..._ws_forwarder_hotpath_optimization_test.go | 18 ++ ...penai_ws_forwarder_ingress_session_test.go | 252 ++++++++++++++++++ .../openai_ws_forwarder_success_test.go | 64 +++++ .../service/openai_ws_v2/passthrough_relay.go | 26 +- .../passthrough_relay_internal_test.go | 40 ++- .../openai_ws_v2_passthrough_adapter.go | 11 +- 9 files changed, 439 insertions(+), 19 deletions(-) diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index f93cc221..7a04f78e 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -4861,7 +4861,7 @@ func (s *OpenAIGatewayService) parseSSEUsageBytes(data []byte, usage *OpenAIUsag return } eventType := gjson.GetBytes(data, "type").String() - if eventType != "response.completed" && eventType != "response.done" && + if eventType != "response.completed" && eventType != "response.done" && eventType != "response.failed" && eventType != "response.incomplete" && eventType != "response.cancelled" && eventType != "response.canceled" { return } diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index 8bed920d..8aad2fa6 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -2218,6 +2218,12 @@ func TestParseSSEUsage_SelectiveParsing(t *testing.T) { require.Equal(t, 15, usage.OutputTokens) require.Equal(t, 4, usage.CacheReadInputTokens) + // failed 事件在部分上游路径也会携带已消耗 usage,应与 WS/passthrough 保持一致 + svc.parseSSEUsage(`{"type":"response.failed","response":{"usage":{"input_tokens":17,"output_tokens":19,"input_tokens_details":{"cached_tokens":6}}}}`, usage) + require.Equal(t, 17, usage.InputTokens) + require.Equal(t, 19, usage.OutputTokens) + require.Equal(t, 6, usage.CacheReadInputTokens) + svc.parseSSEUsage(`{"type":"response.completed","response":{"usage":{"prompt_tokens":21,"completion_tokens":8,"prompt_tokens_details":{"cached_tokens":6}}}}`, usage) require.Equal(t, 21, usage.InputTokens) require.Equal(t, 8, usage.OutputTokens) diff --git a/backend/internal/service/openai_ws_forwarder.go b/backend/internal/service/openai_ws_forwarder.go index b8e558ae..75f6559b 100644 --- a/backend/internal/service/openai_ws_forwarder.go +++ b/backend/internal/service/openai_ws_forwarder.go @@ -369,7 +369,12 @@ func openAIWSEventMayContainToolCalls(eventType string) bool { } func openAIWSEventShouldParseUsage(eventType string) bool { - return eventType == "response.completed" || strings.TrimSpace(eventType) == "response.completed" + switch strings.TrimSpace(eventType) { + case "response.completed", "response.done", "response.failed", "response.incomplete", "response.cancelled", "response.canceled": + return true + default: + return false + } } func parseOpenAIWSEventEnvelope(message []byte) (eventType string, responseID string, response gjson.Result) { @@ -2484,6 +2489,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( imageInputSize string payloadBytes int } + ingressSessionOriginalModel := "" applyPayloadMutation := func(current []byte, path string, value any) ([]byte, error) { next, err := sjson.SetBytes(current, path, value) @@ -2547,12 +2553,21 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( } originalModel := strings.TrimSpace(values[1].String()) + modelMissing := originalModel == "" if originalModel == "" { - return openAIWSClientPayload{}, NewOpenAIWSClientCloseError( - coderws.StatusPolicyViolation, - "model is required in response.create payload", - nil, - ) + // 入站 WS 长会话里,部分客户端只在第一轮 response.create 上声明 + // model,后续 turn 复用同一 session-level model。为避免因省略 + // model 直接断开用户连接,这里回落到上一轮已通过校验的客户端模型, + // 并在下方写回上游 payload,保证账号模型映射/fast policy/图片权限 + // 仍按同一模型执行。 + originalModel = ingressSessionOriginalModel + if originalModel == "" { + return openAIWSClientPayload{}, NewOpenAIWSClientCloseError( + coderws.StatusPolicyViolation, + "model is required in response.create payload", + nil, + ) + } } promptCacheKey := strings.TrimSpace(values[2].String()) previousResponseID := strings.TrimSpace(values[3].String()) @@ -2572,7 +2587,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( normalized = next } upstreamModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel)) - if upstreamModel != originalModel { + if modelMissing || upstreamModel != originalModel { next, setErr := applyPayloadMutation(normalized, "model", upstreamModel) if setErr != nil { return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", setErr) @@ -2602,11 +2617,10 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( // single integration point for all WS ingress turns (first + follow-up // frames flow through here). // - // Model fallback: parseClientPayload above rejects any frame whose - // "model" field is missing (line ~2493-2500), so by the time we - // reach this point upstreamModel is always derived from a non-empty - // per-frame model. The capturedSessionModel fallback used in the - // passthrough adapter is therefore not needed in this path. + // Model fallback: first turn still requires model at the handler layer; + // follow-up response.create frames may omit it and then reuse + // ingressSessionOriginalModel. We always write a concrete upstream model + // before evaluating policy, so whitelist / filter behavior remains stable. policyApplied, blocked, policyErr := s.applyOpenAIFastPolicyToWSResponseCreate(ctx, account, upstreamModel, normalized) if policyErr != nil { return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", policyErr) @@ -2635,6 +2649,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( ) } normalized = policyApplied + ingressSessionOriginalModel = originalModel return openAIWSClientPayload{ payloadRaw: normalized, diff --git a/backend/internal/service/openai_ws_forwarder_hotpath_optimization_test.go b/backend/internal/service/openai_ws_forwarder_hotpath_optimization_test.go index 0350bde9..2622f7f2 100644 --- a/backend/internal/service/openai_ws_forwarder_hotpath_optimization_test.go +++ b/backend/internal/service/openai_ws_forwarder_hotpath_optimization_test.go @@ -39,6 +39,24 @@ func TestParseOpenAIWSResponseUsageFromCompletedEvent(t *testing.T) { require.Equal(t, 4, usage.CacheReadInputTokens) } +func TestOpenAIWSEventShouldParseUsageTerminalEvents(t *testing.T) { + t.Parallel() + + for _, eventType := range []string{ + "response.completed", + "response.done", + "response.failed", + "response.incomplete", + "response.cancelled", + "response.canceled", + } { + require.True(t, openAIWSEventShouldParseUsage(eventType), eventType) + require.True(t, openAIWSEventShouldParseUsage(" "+eventType+" "), eventType) + } + require.False(t, openAIWSEventShouldParseUsage("response.output_text.delta")) + require.False(t, openAIWSEventShouldParseUsage("")) +} + func TestOpenAIWSErrorEventHelpers_ConsistentWithWrapper(t *testing.T) { message := []byte(`{"type":"error","error":{"type":"invalid_request_error","code":"invalid_request","message":"invalid input"}}`) codeRaw, errTypeRaw, errMsgRaw := parseOpenAIWSErrorEventFields(message) 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 edb6fbcd..b7f1bc4f 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress_session_test.go +++ b/backend/internal/service/openai_ws_forwarder_ingress_session_test.go @@ -164,6 +164,140 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_KeepLeaseAcrossT require.Len(t, captureConn.writes, 2, "应向同一上游连接发送两轮 response.create") } +func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_FollowupCreateCanOmitModel(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_omit_model_1","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`), + []byte(`{"type":"response.completed","response":{"id":"resp_omit_model_2","model":"gpt-5.1","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, + } + account := &Account{ + ID: 115, + Name: "openai-ingress-omit-model", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "model_mapping": map[string]any{ + "client-model": "gpt-5.1", + }, + }, + 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() + }() + + writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second) + err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"client-model","stream":false}`)) + cancelWrite() + require.NoError(t, err) + + readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second) + _, firstEvent, readErr := clientConn.Read(readCtx) + cancelRead() + require.NoError(t, readErr) + require.Equal(t, "resp_omit_model_1", gjson.GetBytes(firstEvent, "response.id").String()) + + writeCtx, cancelWrite = context.WithTimeout(context.Background(), 3*time.Second) + err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","stream":false,"previous_response_id":"resp_omit_model_1"}`)) + cancelWrite() + require.NoError(t, err) + + readCtx, cancelRead = context.WithTimeout(context.Background(), 3*time.Second) + _, secondEvent, readErr := clientConn.Read(readCtx) + cancelRead() + require.NoError(t, readErr) + require.Equal(t, "resp_omit_model_2", gjson.GetBytes(secondEvent, "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, 2) + require.Equal(t, "gpt-5.1", gjson.Get(requestToJSONString(captureConn.writes[0]), "model").String()) + require.Equal(t, "gpt-5.1", gjson.Get(requestToJSONString(captureConn.writes[1]), "model").String()) + require.Equal(t, "resp_omit_model_1", gjson.Get(requestToJSONString(captureConn.writes[1]), "previous_response_id").String()) +} + func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_DedicatedModeDoesNotReuseConnAcrossSessions(t *testing.T) { gin.SetMode(gin.TestMode) @@ -441,6 +575,124 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_PassthroughModeR require.Len(t, upstreamConn.writes, 1, "passthrough 模式应透传首条 response.create") } +func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_PassthroughHeadersUsePromptCacheAndTurnState(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.ModeRouterV2Enabled = true + cfg.Gateway.OpenAIWS.IngressModeDefault = OpenAIWSIngressModeCtxPool + cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3 + cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3 + cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3 + + upstreamConn := &openAIWSCaptureConn{ + events: [][]byte{ + []byte(`{"type":"response.completed","response":{"id":"resp_passthrough_headers","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`), + }, + } + captureDialer := &openAIWSCaptureDialer{conn: upstreamConn} + svc := &OpenAIGatewayService{ + cfg: cfg, + httpUpstream: &httpUpstreamRecorder{}, + cache: &stubGatewayCache{}, + openaiWSResolver: NewOpenAIWSProtocolResolver(cfg), + toolCorrector: NewCodexToolCorrector(), + openaiWSPassthroughDialer: captureDialer, + } + account := &Account{ + ID: 453, + Name: "openai-ingress-passthrough-headers", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "oauth-token", + }, + Extra: map[string]any{ + "openai_oauth_responses_websockets_v2_mode": OpenAIWSIngressModePassthrough, + }, + } + + 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") + req.Header.Set(openAIWSTurnStateHeader, "turn-state-1") + req.Header.Set(openAIWSTurnMetadataHeader, "turn-meta-1") + 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, "oauth-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.1","stream":false,"prompt_cache_key":"pcache_passthrough"}`)) + cancelWrite() + require.NoError(t, err) + + readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second) + _, event, readErr := clientConn.Read(readCtx) + cancelRead() + require.NoError(t, readErr) + require.Equal(t, "resp_passthrough_headers", gjson.GetBytes(event, "response.id").String()) + _ = clientConn.Close(coderws.StatusNormalClosure, "done") + + select { + case serverErr := <-serverErrCh: + if serverErr != nil { + require.Contains(t, serverErr.Error(), "StatusNormalClosure") + } + case <-time.After(5 * time.Second): + t.Fatal("等待 passthrough websocket 结束超时") + } + + require.Equal(t, isolateOpenAISessionID(0, "pcache_passthrough"), captureDialer.lastHeaders.Get("session_id")) + require.Equal(t, "turn-state-1", captureDialer.lastHeaders.Get(openAIWSTurnStateHeader)) + require.Equal(t, "turn-meta-1", captureDialer.lastHeaders.Get(openAIWSTurnMetadataHeader)) +} + func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_ModeOffReturnsPolicyViolation(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/openai_ws_forwarder_success_test.go b/backend/internal/service/openai_ws_forwarder_success_test.go index cd816533..e949560f 100644 --- a/backend/internal/service/openai_ws_forwarder_success_test.go +++ b/backend/internal/service/openai_ws_forwarder_success_test.go @@ -727,6 +727,70 @@ func TestOpenAIGatewayService_Forward_WSv2_HeaderSessionFallbackFromPromptCacheK require.True(t, gjson.Get(requestToJSONString(captureConn.lastWrite), "stream").Exists()) } +func TestOpenAIGatewayService_Forward_WSv2_ResponseDoneUsageParsed(t *testing.T) { + gin.SetMode(gin.TestMode) + + 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.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 + + captureConn := &openAIWSCaptureConn{ + events: [][]byte{ + []byte(`{"type":"response.done","response":{"id":"resp_done_usage","model":"gpt-5.1","usage":{"input_tokens":13,"output_tokens":8,"input_tokens_details":{"cached_tokens":5},"cache_creation_input_tokens":2,"output_tokens_details":{"image_tokens":4}}}}`), + }, + } + 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, + } + account := &Account{ + ID: 32, + Name: "openai-ws-done", + 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, + }, + } + + body := []byte(`{"model":"gpt-5.1","stream":false,"input":[{"type":"input_text","text":"hi"}]}`) + result, err := svc.Forward(context.Background(), c, account, body) + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "resp_done_usage", result.RequestID) + require.Equal(t, 13, result.Usage.InputTokens) + require.Equal(t, 8, result.Usage.OutputTokens) + require.Equal(t, 5, result.Usage.CacheReadInputTokens) + require.Equal(t, 2, result.Usage.CacheCreationInputTokens) + require.Equal(t, 4, result.Usage.ImageOutputTokens) +} + func TestOpenAIGatewayService_Forward_WSv1_Unsupported(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/openai_ws_v2/passthrough_relay.go b/backend/internal/service/openai_ws_v2/passthrough_relay.go index 35c7569d..6aba3b7d 100644 --- a/backend/internal/service/openai_ws_v2/passthrough_relay.go +++ b/backend/internal/service/openai_ws_v2/passthrough_relay.go @@ -25,6 +25,7 @@ type Usage struct { OutputTokens int CacheCreationInputTokens int CacheReadInputTokens int + ImageOutputTokens int } type RelayResult struct { @@ -756,8 +757,21 @@ func parseUsageAndAccumulate( } inputResult := gjson.GetBytes(message, "response.usage.input_tokens") + if !inputResult.Exists() { + inputResult = gjson.GetBytes(message, "response.usage.prompt_tokens") + } outputResult := gjson.GetBytes(message, "response.usage.output_tokens") + if !outputResult.Exists() { + outputResult = gjson.GetBytes(message, "response.usage.completion_tokens") + } cachedResult := gjson.GetBytes(message, "response.usage.input_tokens_details.cached_tokens") + if !cachedResult.Exists() { + cachedResult = gjson.GetBytes(message, "response.usage.prompt_tokens_details.cached_tokens") + } + imageTokens := usageResult.Get("output_tokens_details.image_tokens").Int() + if imageTokens == 0 { + imageTokens = usageResult.Get("completion_tokens_details.image_tokens").Int() + } inputTokens, inputOK := parseUsageIntField(inputResult, true) outputTokens, outputOK := parseUsageIntField(outputResult, true) @@ -771,14 +785,18 @@ func parseUsageAndAccumulate( return Usage{} } parsedUsage := Usage{ - InputTokens: inputTokens, - OutputTokens: outputTokens, - CacheReadInputTokens: cachedTokens, + InputTokens: inputTokens, + OutputTokens: outputTokens, + CacheCreationInputTokens: int(usageResult.Get("cache_creation_input_tokens").Int()), + CacheReadInputTokens: cachedTokens, + ImageOutputTokens: int(imageTokens), } state.usage.InputTokens += parsedUsage.InputTokens state.usage.OutputTokens += parsedUsage.OutputTokens + state.usage.CacheCreationInputTokens += parsedUsage.CacheCreationInputTokens state.usage.CacheReadInputTokens += parsedUsage.CacheReadInputTokens + state.usage.ImageOutputTokens += parsedUsage.ImageOutputTokens return parsedUsage } @@ -840,7 +858,7 @@ func isTerminalEvent(eventType string) bool { func shouldParseUsage(eventType string) bool { switch eventType { - case "response.completed", "response.done", "response.failed": + case "response.completed", "response.done", "response.failed", "response.incomplete", "response.cancelled", "response.canceled": return true default: return false diff --git a/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go b/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go index 52104482..13c51f66 100644 --- a/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go +++ b/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go @@ -300,20 +300,41 @@ func TestParseUsageAndEnrichCoverage(t *testing.T) { require.Equal(t, 0, state.usage.OutputTokens) require.Equal(t, 0, state.usage.CacheReadInputTokens) - parseUsageAndAccumulate(state, []byte(`{"type":"response.completed","response":{"usage":{"input_tokens":2,"output_tokens":1,"input_tokens_details":{"cached_tokens":1}}}}`), "response.completed", nil) + parseUsageAndAccumulate(state, []byte(`{"type":"response.completed","response":{"usage":{"input_tokens":2,"output_tokens":1,"input_tokens_details":{"cached_tokens":1},"cache_creation_input_tokens":4,"output_tokens_details":{"image_tokens":3}}}}`), "response.completed", nil) require.Equal(t, 2, state.usage.InputTokens) require.Equal(t, 1, state.usage.OutputTokens) require.Equal(t, 1, state.usage.CacheReadInputTokens) + require.Equal(t, 4, state.usage.CacheCreationInputTokens) + require.Equal(t, 3, state.usage.ImageOutputTokens) result := &RelayResult{} enrichResult(result, state, 5*time.Millisecond) require.Equal(t, state.usage.InputTokens, result.Usage.InputTokens) + require.Equal(t, state.usage.CacheCreationInputTokens, result.Usage.CacheCreationInputTokens) + require.Equal(t, state.usage.ImageOutputTokens, result.Usage.ImageOutputTokens) require.Equal(t, 5*time.Millisecond, result.Duration) parseUsageAndAccumulate(state, []byte(`{"type":"response.in_progress","response":{"usage":{"input_tokens":9}}}`), "response.in_progress", nil) require.Equal(t, 2, state.usage.InputTokens) enrichResult(nil, state, 0) } +func TestParseUsageAndAccumulateAcceptsChatUsageAliases(t *testing.T) { + t.Parallel() + + state := &relayState{} + got := parseUsageAndAccumulate( + state, + []byte(`{"type":"response.done","response":{"usage":{"prompt_tokens":12,"completion_tokens":6,"prompt_tokens_details":{"cached_tokens":4},"completion_tokens_details":{"image_tokens":2}}}}`), + "response.done", + nil, + ) + require.Equal(t, 12, got.InputTokens) + require.Equal(t, 6, got.OutputTokens) + require.Equal(t, 4, got.CacheReadInputTokens) + require.Equal(t, 2, got.ImageOutputTokens) + require.Equal(t, got, state.usage) +} + func TestEmitTurnCompleteCoverage(t *testing.T) { t.Parallel() @@ -377,6 +398,23 @@ func TestIsTokenEventCoverageBranches(t *testing.T) { require.True(t, isTokenEvent("response.done")) } +func TestShouldParseUsageTerminalEvents(t *testing.T) { + t.Parallel() + + for _, eventType := range []string{ + "response.completed", + "response.done", + "response.failed", + "response.incomplete", + "response.cancelled", + "response.canceled", + } { + require.True(t, shouldParseUsage(eventType), eventType) + } + require.False(t, shouldParseUsage("response.output_text.delta")) + require.False(t, shouldParseUsage("")) +} + func TestRelayTurnTimingHelpersCoverage(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/openai_ws_v2_passthrough_adapter.go b/backend/internal/service/openai_ws_v2_passthrough_adapter.go index 17543dc0..c93d0981 100644 --- a/backend/internal/service/openai_ws_v2_passthrough_adapter.go +++ b/backend/internal/service/openai_ws_v2_passthrough_adapter.go @@ -312,6 +312,7 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( // goroutine)和 OnTurnComplete / final result(runUpstreamToClient // goroutine)之间同步当前 turn 的 usage metadata。 usageMeta.initFromFirstFrame(firstClientMessage) + promptCacheKey := strings.TrimSpace(gjson.GetBytes(firstClientMessage, "prompt_cache_key").String()) wsURL, err := s.buildOpenAIResponsesWSURL(account) if err != nil { @@ -338,7 +339,13 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( if s.cfg != nil && s.cfg.Gateway.ForceCodexCLI { isCodexCLI = true } - headers, _ := s.buildOpenAIWSHeaders(c, account, token, wsDecision, isCodexCLI, "", "", "") + turnState := "" + turnMetadata := "" + if c != nil { + turnState = strings.TrimSpace(c.GetHeader(openAIWSTurnStateHeader)) + turnMetadata = strings.TrimSpace(c.GetHeader(openAIWSTurnMetadataHeader)) + } + headers, _ := s.buildOpenAIWSHeaders(c, account, token, wsDecision, isCodexCLI, turnState, turnMetadata, promptCacheKey) proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { proxyURL = account.Proxy.URL() @@ -519,6 +526,7 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( OutputTokens: turn.Usage.OutputTokens, CacheCreationInputTokens: turn.Usage.CacheCreationInputTokens, CacheReadInputTokens: turn.Usage.CacheReadInputTokens, + ImageOutputTokens: turn.Usage.ImageOutputTokens, }, Model: turn.RequestModel, ServiceTier: usageMeta.serviceTier.Load(), @@ -593,6 +601,7 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( OutputTokens: relayResult.Usage.OutputTokens, CacheCreationInputTokens: relayResult.Usage.CacheCreationInputTokens, CacheReadInputTokens: relayResult.Usage.CacheReadInputTokens, + ImageOutputTokens: relayResult.Usage.ImageOutputTokens, }, Model: relayResult.RequestModel, ServiceTier: usageMeta.serviceTier.Load(), From 27600b1d2c9579f83abbb1ee469bfc45b908f840 Mon Sep 17 00:00:00 2001 From: Pluviobyte Date: Thu, 28 May 2026 05:40:50 +0000 Subject: [PATCH 19/79] fix(gateway): filter count_tokens generation fields Anthropic count_tokens rejects generation-only fields such as temperature, top_p, top_k, stream, and stop sequences. Passing the original messages payload through unchanged can turn otherwise valid requests into upstream 400 errors. Sanitize only the count_tokens upstream payload after the gateway's existing request normalization, preserving fields that existing compatibility paths rely on while removing parameters the count_tokens endpoint does not accept. Fixes #2764 Co-authored-by: Cursor --- ...teway_anthropic_apikey_passthrough_test.go | 60 +++++++++++++++++++ backend/internal/service/gateway_service.go | 21 +++++++ 2 files changed, 81 insertions(+) diff --git a/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go b/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go index 5cb03f30..9062c517 100644 --- a/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go +++ b/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go @@ -476,6 +476,66 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_ModelMappingPreservesOtherFie require.Equal(t, int64(1024), gjson.GetBytes(sentBody, "max_tokens").Int(), "max_tokens 不应被修改") } +func TestGatewayService_AnthropicAPIKeyPassthrough_CountTokensFiltersGenerationFields(t *testing.T) { + gin.SetMode(gin.TestMode) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", nil) + + 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, + Model: "claude-sonnet-4-20250514", + } + + upstreamRespBody := `{"input_tokens":42}` + upstream := &anthropicHTTPUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(upstreamRespBody)), + }, + } + + svc := &GatewayService{ + cfg: &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}}, + httpUpstream: upstream, + rateLimitService: &RateLimitService{}, + } + + account := &Account{ + ID: 302, + Name: "count-token-filter-test", + Platform: PlatformAnthropic, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "upstream-key", + "base_url": "https://api.anthropic.com", + }, + Extra: map[string]any{"anthropic_passthrough": true}, + Status: StatusActive, + Schedulable: true, + } + + err := svc.ForwardCountTokens(context.Background(), c, account, parsed) + require.NoError(t, err) + + sentBody := upstream.lastBody + require.False(t, gjson.GetBytes(sentBody, "temperature").Exists()) + require.False(t, gjson.GetBytes(sentBody, "top_p").Exists()) + require.False(t, gjson.GetBytes(sentBody, "top_k").Exists()) + require.False(t, gjson.GetBytes(sentBody, "stream").Exists()) + require.False(t, gjson.GetBytes(sentBody, "stop_sequences").Exists()) + require.Equal(t, "claude-sonnet-4-20250514", gjson.GetBytes(sentBody, "model").String()) + require.Equal(t, "sys", gjson.GetBytes(sentBody, "system.0.text").String()) + require.Equal(t, "hello", gjson.GetBytes(sentBody, "messages.0.content").String()) + require.Equal(t, "tool", gjson.GetBytes(sentBody, "tools.0.name").String()) + require.Equal(t, int64(1024), gjson.GetBytes(sentBody, "max_tokens").Int()) + require.Equal(t, "enabled", gjson.GetBytes(sentBody, "thinking.type").String()) +} + // TestGatewayService_AnthropicAPIKeyPassthrough_EmptyModelSkipsMapping // 确保空模型名不会触发映射逻辑 func TestGatewayService_AnthropicAPIKeyPassthrough_EmptyModelSkipsMapping(t *testing.T) { diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index 4a8175a4..a787e3eb 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -9311,6 +9311,7 @@ func (s *GatewayService) buildCountTokensRequestAnthropicAPIKeyPassthrough( } targetURL = validatedURL + "/v1/messages/count_tokens?beta=true" } + body = sanitizeCountTokensRequestBody(body) req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body)) if err != nil { @@ -9405,6 +9406,7 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con if ctEnableCCH { body = signBillingHeaderCCH(body) } + body = sanitizeCountTokensRequestBody(body) req, err := http.NewRequestWithContext(ctx, "POST", targetURL, bytes.NewReader(body)) if err != nil { @@ -9501,6 +9503,25 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con return req, nil } +func sanitizeCountTokensRequestBody(body []byte) []byte { + out := body + for _, path := range []string{ + "temperature", + "top_p", + "top_k", + "stream", + "stop_sequences", + "stop", + } { + if gjson.GetBytes(out, path).Exists() { + if next, ok := deleteJSONPathBytes(out, path); ok { + out = next + } + } + } + return out +} + // countTokensError 返回 count_tokens 错误响应 func (s *GatewayService) countTokensError(c *gin.Context, status int, errType, message string) { c.JSON(status, gin.H{ From b15375dfb4b3fa40aaec53012385106c8c29d3ae Mon Sep 17 00:00:00 2001 From: wucm667 Date: Thu, 28 May 2026 17:27:01 +0800 Subject: [PATCH 20/79] fix(admin): handle already up-to-date updates --- .../internal/handler/admin/system_handler.go | 26 +++- .../handler/admin/system_handler_test.go | 144 ++++++++++++++++++ backend/internal/service/update_service.go | 8 +- .../internal/service/update_service_test.go | 64 ++++++++ 4 files changed, 239 insertions(+), 3 deletions(-) create mode 100644 backend/internal/handler/admin/system_handler_test.go create mode 100644 backend/internal/service/update_service_test.go diff --git a/backend/internal/handler/admin/system_handler.go b/backend/internal/handler/admin/system_handler.go index 3e2022c7..fb6c0ef7 100644 --- a/backend/internal/handler/admin/system_handler.go +++ b/backend/internal/handler/admin/system_handler.go @@ -2,6 +2,7 @@ package admin import ( "context" + "errors" "net/http" "strconv" "strings" @@ -17,12 +18,18 @@ import ( // SystemHandler handles system-related operations type SystemHandler struct { - updateSvc *service.UpdateService + updateSvc systemUpdateService lockSvc *service.SystemOperationLockService } +type systemUpdateService interface { + CheckUpdate(ctx context.Context, force bool) (*service.UpdateInfo, error) + PerformUpdate(ctx context.Context) error + Rollback() error +} + // NewSystemHandler creates a new SystemHandler -func NewSystemHandler(updateSvc *service.UpdateService, lockSvc *service.SystemOperationLockService) *SystemHandler { +func NewSystemHandler(updateSvc systemUpdateService, lockSvc *service.SystemOperationLockService) *SystemHandler { return &SystemHandler{ updateSvc: updateSvc, lockSvc: lockSvc, @@ -67,6 +74,21 @@ func (h *SystemHandler) PerformUpdate(c *gin.Context) { }() if err := h.updateSvc.PerformUpdate(ctx); err != nil { + if errors.Is(err, service.ErrNoUpdateAvailable) { + info, checkErr := h.updateSvc.CheckUpdate(ctx, false) + if checkErr != nil { + releaseReason = "SYSTEM_UPDATE_FAILED" + return nil, checkErr + } + succeeded = true + return gin.H{ + "message": "Already up to date", + "already_up_to_date": true, + "current_version": info.CurrentVersion, + "latest_version": info.LatestVersion, + "operation_id": lock.OperationID(), + }, nil + } releaseReason = "SYSTEM_UPDATE_FAILED" return nil, err } diff --git a/backend/internal/handler/admin/system_handler_test.go b/backend/internal/handler/admin/system_handler_test.go new file mode 100644 index 00000000..0f33a452 --- /dev/null +++ b/backend/internal/handler/admin/system_handler_test.go @@ -0,0 +1,144 @@ +//go:build unit + +package admin + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +type systemHandlerUpdateServiceStub struct { + performErr error + updateInfo *service.UpdateInfo + checkErr error + checkForces []bool + performCall int +} + +func (s *systemHandlerUpdateServiceStub) CheckUpdate(_ context.Context, force bool) (*service.UpdateInfo, error) { + s.checkForces = append(s.checkForces, force) + return s.updateInfo, s.checkErr +} + +func (s *systemHandlerUpdateServiceStub) PerformUpdate(context.Context) error { + s.performCall++ + return s.performErr +} + +func (s *systemHandlerUpdateServiceStub) Rollback() error { + return nil +} + +type systemUpdateResponseEnvelope struct { + Code int `json:"code"` + Message string `json:"message"` + Data struct { + Message string `json:"message"` + AlreadyUpToDate bool `json:"already_up_to_date"` + CurrentVersion string `json:"current_version"` + LatestVersion string `json:"latest_version"` + OperationID string `json:"operation_id"` + } `json:"data"` +} + +type systemUpdateErrorEnvelope struct { + Code int `json:"code"` + Message string `json:"message"` +} + +func newSystemHandlerTestRouter(t *testing.T, updateSvc *systemHandlerUpdateServiceStub, repo *memoryIdempotencyRepoStub) *gin.Engine { + t.Helper() + gin.SetMode(gin.TestMode) + service.SetDefaultIdempotencyCoordinator(nil) + t.Cleanup(func() { + service.SetDefaultIdempotencyCoordinator(nil) + }) + + lockSvc := service.NewSystemOperationLockService(repo, service.IdempotencyConfig{ + ProcessingTimeout: time.Second, + SystemOperationTTL: time.Minute, + }) + handler := NewSystemHandler(updateSvc, lockSvc) + + router := gin.New() + router.POST("/api/v1/admin/system/update", handler.PerformUpdate) + return router +} + +func requireSystemLockStatus(t *testing.T, repo *memoryIdempotencyRepoStub, wantStatus string) { + t.Helper() + repo.mu.Lock() + defer repo.mu.Unlock() + + for _, record := range repo.data { + if record.Status == wantStatus { + return + } + } + t.Fatalf("system lock status %q not found in records: %#v", wantStatus, repo.data) +} + +func TestSystemHandlerPerformUpdateAlreadyUpToDateReturnsOK(t *testing.T) { + updateSvc := &systemHandlerUpdateServiceStub{ + performErr: service.ErrNoUpdateAvailable, + updateInfo: &service.UpdateInfo{ + CurrentVersion: "0.1.132", + LatestVersion: "0.1.132", + HasUpdate: false, + }, + } + repo := newMemoryIdempotencyRepoStub() + router := newSystemHandlerTestRouter(t, updateSvc, repo) + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/system/update", nil) + req.Header.Set("Idempotency-Key", "already-up-to-date") + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, 1, updateSvc.performCall) + require.Equal(t, []bool{false}, updateSvc.checkForces) + requireSystemLockStatus(t, repo, service.IdempotencyStatusSucceeded) + + var body systemUpdateResponseEnvelope + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &body)) + require.Equal(t, 0, body.Code) + require.Equal(t, "success", body.Message) + require.Equal(t, "Already up to date", body.Data.Message) + require.True(t, body.Data.AlreadyUpToDate) + require.Equal(t, "0.1.132", body.Data.CurrentVersion) + require.Equal(t, "0.1.132", body.Data.LatestVersion) + require.NotEmpty(t, body.Data.OperationID) +} + +func TestSystemHandlerPerformUpdateFailureStillReturnsInternalError(t *testing.T) { + updateSvc := &systemHandlerUpdateServiceStub{ + performErr: errors.New("download failed"), + } + repo := newMemoryIdempotencyRepoStub() + router := newSystemHandlerTestRouter(t, updateSvc, repo) + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/system/update", nil) + req.Header.Set("Idempotency-Key", "real-failure") + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusInternalServerError, rec.Code) + require.Equal(t, 1, updateSvc.performCall) + require.Empty(t, updateSvc.checkForces) + requireSystemLockStatus(t, repo, service.IdempotencyStatusFailedRetryable) + + var body systemUpdateErrorEnvelope + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &body)) + require.Equal(t, http.StatusInternalServerError, body.Code) + require.Equal(t, "internal error", body.Message) +} diff --git a/backend/internal/service/update_service.go b/backend/internal/service/update_service.go index 34ad4610..de8c5e16 100644 --- a/backend/internal/service/update_service.go +++ b/backend/internal/service/update_service.go @@ -17,6 +17,12 @@ import ( "strconv" "strings" "time" + + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" +) + +var ( + ErrNoUpdateAvailable = infraerrors.Conflict("ALREADY_UP_TO_DATE", "no update available; current version is latest") ) const ( @@ -146,7 +152,7 @@ func (s *UpdateService) PerformUpdate(ctx context.Context) error { } if !info.HasUpdate { - return fmt.Errorf("no update available") + return ErrNoUpdateAvailable } // Find matching archive and checksum for current platform diff --git a/backend/internal/service/update_service_test.go b/backend/internal/service/update_service_test.go new file mode 100644 index 00000000..8d8310d4 --- /dev/null +++ b/backend/internal/service/update_service_test.go @@ -0,0 +1,64 @@ +//go:build unit + +package service + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +type updateServiceCacheStub struct { + data string +} + +func (s *updateServiceCacheStub) GetUpdateInfo(context.Context) (string, error) { + if s.data == "" { + return "", errors.New("cache miss") + } + return s.data, nil +} + +func (s *updateServiceCacheStub) SetUpdateInfo(_ context.Context, data string, _ time.Duration) error { + s.data = data + return nil +} + +type updateServiceGitHubClientStub struct { + release *GitHubRelease +} + +func (s *updateServiceGitHubClientStub) FetchLatestRelease(context.Context, string) (*GitHubRelease, error) { + return s.release, nil +} + +func (s *updateServiceGitHubClientStub) DownloadFile(context.Context, string, string, int64) error { + panic("DownloadFile should not be called when no update is available") +} + +func (s *updateServiceGitHubClientStub) FetchChecksumFile(context.Context, string) ([]byte, error) { + panic("FetchChecksumFile should not be called when no update is available") +} + +func TestUpdateServicePerformUpdateNoUpdateReturnsSentinel(t *testing.T) { + svc := NewUpdateService( + &updateServiceCacheStub{}, + &updateServiceGitHubClientStub{ + release: &GitHubRelease{ + TagName: "v0.1.132", + Name: "v0.1.132", + }, + }, + "0.1.132", + "release", + ) + + err := svc.PerformUpdate(context.Background()) + + require.Error(t, err) + require.True(t, errors.Is(err, ErrNoUpdateAvailable)) + require.ErrorIs(t, err, ErrNoUpdateAvailable) +} From e9a2db8e80b70091e64fcd88a3cda96db80d22c1 Mon Sep 17 00:00:00 2001 From: haichuan Date: Thu, 28 May 2026 18:02:18 +0800 Subject: [PATCH 21/79] fix: normalize responses streaming terminal output --- .../service/openai_gateway_service.go | 87 ++++++++++++++++--- .../service/openai_gateway_service_test.go | 79 +++++++++++++++++ 2 files changed, 153 insertions(+), 13 deletions(-) diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index f93cc221..8b7e837b 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -4454,6 +4454,9 @@ func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp } needModelReplace := originalModel != mappedModel + streamOutputAccumulator := apicompat.NewBufferedResponseAccumulator() + streamImageOutputs := make([]json.RawMessage, 0, 1) + streamSeenImages := make(map[string]struct{}) resultWithUsage := func() *openaiStreamingResult { return &openaiStreamingResult{ usage: usage, @@ -4532,13 +4535,6 @@ func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp } // Extract data from SSE line (supports both "data: " and "data:" formats) if data, ok := extractOpenAISSEDataLine(line); ok { - - // Replace model in response if needed. - // Fast path: most events do not contain model field values. - if needModelReplace && mappedModel != "" && strings.Contains(data, mappedModel) { - line = s.replaceModelInSSELine(line, mappedModel, originalModel) - } - dataBytes := []byte(data) if openAIStreamEventIsTerminal(data) { sawTerminalEvent = true @@ -4564,6 +4560,26 @@ func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp line = "data: " + data eventType = strings.TrimSpace(gjson.GetBytes(dataBytes, "type").String()) } + if imageOutput, ok := extractImageGenerationOutputFromSSEData(dataBytes, streamSeenImages); ok { + streamImageOutputs = append(streamImageOutputs, imageOutput) + } + if responsesStreamEventMayContributeToOutput(eventType) { + var streamEvent apicompat.ResponsesStreamEvent + if err := json.Unmarshal(dataBytes, &streamEvent); err == nil { + streamOutputAccumulator.ProcessEvent(&streamEvent) + } + } + if normalizedData, normalized := normalizeResponsesStreamingTerminalOutput(dataBytes, streamOutputAccumulator, streamImageOutputs); normalized { + dataBytes = normalizedData + data = string(normalizedData) + line = "data: " + data + eventType = strings.TrimSpace(gjson.GetBytes(dataBytes, "type").String()) + } + // Replace model in response if needed. + // Fast path: most events do not contain model field values. + if needModelReplace && mappedModel != "" && strings.Contains(line, mappedModel) { + line = s.replaceModelInSSELine(line, mappedModel, originalModel) + } startsClientOutput := forceFlushFailedEvent || openAIStreamDataStartsClientOutput(data, eventType) // 写入客户端(客户端断开后继续 drain 上游) @@ -5099,6 +5115,45 @@ func extractCodexFinalResponse(body string) ([]byte, bool) { return nil, false } +func normalizeResponsesStreamingTerminalOutput(data []byte, acc *apicompat.BufferedResponseAccumulator, imageOutputs []json.RawMessage) ([]byte, bool) { + eventType := strings.TrimSpace(gjson.GetBytes(data, "type").String()) + switch eventType { + case "response.completed", "response.done", "response.incomplete", "response.cancelled", "response.canceled": + default: + return data, false + } + + output := gjson.GetBytes(data, "response.output") + hasAccumulatedOutput := (acc != nil && acc.HasContent()) || len(imageOutputs) > 0 + if output.Exists() && output.IsArray() { + if len(output.Array()) > 0 || !hasAccumulatedOutput { + return data, false + } + } + + outputJSON := []byte("[]") + if reconstructed, ok := buildResponsesOutputJSON(acc, imageOutputs); ok { + outputJSON = reconstructed + } + updated, err := sjson.SetRawBytes(data, "response.output", outputJSON) + if err != nil { + return data, false + } + return updated, true +} + +func responsesStreamEventMayContributeToOutput(eventType string) bool { + switch eventType { + case "response.output_text.delta", + "response.output_item.added", + "response.function_call_arguments.delta", + "response.reasoning_summary_text.delta": + return true + default: + return false + } +} + // reconstructResponseOutputFromSSE scans raw SSE body text for delta events and // returns a JSON-encoded output array reconstructed from accumulated deltas. // Returns (nil, false) if no content was found in deltas. @@ -5110,17 +5165,23 @@ func reconstructResponseOutputFromSSE(bodyText string) ([]byte, bool) { if imageOutput, ok := extractImageGenerationOutputFromSSEData(data, seenImages); ok { imageOutputs = append(imageOutputs, imageOutput) } - var event apicompat.ResponsesStreamEvent - if err := json.Unmarshal(data, &event); err == nil { - acc.ProcessEvent(&event) + eventType := strings.TrimSpace(gjson.GetBytes(data, "type").String()) + if responsesStreamEventMayContributeToOutput(eventType) { + var event apicompat.ResponsesStreamEvent + if err := json.Unmarshal(data, &event); err == nil { + acc.ProcessEvent(&event) + } } }) - if !acc.HasContent() && len(imageOutputs) == 0 { + return buildResponsesOutputJSON(acc, imageOutputs) +} + +func buildResponsesOutputJSON(acc *apicompat.BufferedResponseAccumulator, imageOutputs []json.RawMessage) ([]byte, bool) { + if (acc == nil || !acc.HasContent()) && len(imageOutputs) == 0 { return nil, false } - var output []json.RawMessage - if acc.HasContent() { + if acc != nil && acc.HasContent() { outputJSON, err := json.Marshal(acc.BuildOutput()) if err == nil { _ = json.Unmarshal(outputJSON, &output) diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index 8bed920d..c642fcd4 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -1233,6 +1233,85 @@ func TestOpenAIStreamingPreambleKeepaliveUsesDownstreamIdle(t *testing.T) { require.Contains(t, rec.Body.String(), "response.completed") } +func TestOpenAIStreamingNormalizesTerminalOutputFromDeltas(t *testing.T) { + gin.SetMode(gin.TestMode) + cfg := &config.Config{ + Gateway: config.GatewayConfig{ + StreamDataIntervalTimeout: 0, + StreamKeepaliveInterval: 0, + MaxLineSize: defaultMaxLineSize, + }, + } + svc := &OpenAIGatewayService{cfg: cfg} + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/", nil) + + resp := &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(strings.NewReader(strings.Join([]string{ + `data: {"type":"response.created","response":{"id":"resp_sdk_parse"}}`, + "", + `data: {"type":"response.output_text.delta","delta":"pon"}`, + "", + `data: {"type":"response.output_text.delta","delta":"g"}`, + "", + `data: {"type":"response.completed","response":{"id":"resp_sdk_parse","status":"completed","output":null,"usage":{"input_tokens":1,"output_tokens":1}}}`, + "", + }, "\n"))), + Header: http.Header{"X-Request-Id": []string{"rid-sdk-parse"}}, + } + + result, err := svc.handleStreamingResponse(c.Request.Context(), resp, c, &Account{ID: 1, Platform: PlatformOpenAI, Name: "acc"}, time.Now(), "model", "model") + require.NoError(t, err) + require.NotNil(t, result) + + terminalType, terminalPayload, ok := extractOpenAISSETerminalEvent(rec.Body.String()) + require.True(t, ok) + require.Equal(t, "response.completed", terminalType) + output := gjson.GetBytes(terminalPayload, "response.output") + require.True(t, output.IsArray()) + require.Len(t, output.Array(), 1) + require.Equal(t, "pong", gjson.GetBytes(terminalPayload, "response.output.0.content.0.text").String()) +} + +func TestOpenAIStreamingNormalizesTerminalOutputToEmptyArray(t *testing.T) { + gin.SetMode(gin.TestMode) + cfg := &config.Config{ + Gateway: config.GatewayConfig{ + StreamDataIntervalTimeout: 0, + StreamKeepaliveInterval: 0, + MaxLineSize: defaultMaxLineSize, + }, + } + svc := &OpenAIGatewayService{cfg: cfg} + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/", nil) + + resp := &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(strings.NewReader(strings.Join([]string{ + `data: {"type":"response.completed","response":{"id":"resp_empty","status":"completed","output":null,"usage":{"input_tokens":1,"output_tokens":0}}}`, + "", + }, "\n"))), + Header: http.Header{"X-Request-Id": []string{"rid-empty-output"}}, + } + + result, err := svc.handleStreamingResponse(c.Request.Context(), resp, c, &Account{ID: 1, Platform: PlatformOpenAI, Name: "acc"}, time.Now(), "model", "model") + require.NoError(t, err) + require.NotNil(t, result) + + terminalType, terminalPayload, ok := extractOpenAISSETerminalEvent(rec.Body.String()) + require.True(t, ok) + require.Equal(t, "response.completed", terminalType) + output := gjson.GetBytes(terminalPayload, "response.output") + require.True(t, output.IsArray()) + require.Len(t, output.Array(), 0) +} + func TestOpenAIStreamingPolicyResponseFailedBeforeOutputPassesThrough(t *testing.T) { gin.SetMode(gin.TestMode) cfg := &config.Config{ From 56e96fdd8c0fbb42c833c15808d953a08d3fa40e Mon Sep 17 00:00:00 2001 From: gaoren002 Date: Thu, 28 May 2026 10:03:41 +0000 Subject: [PATCH 22/79] fix: classify concurrency acquire failures --- .../handler/concurrency_error_response.go | 27 ++++++++ .../concurrency_error_response_test.go | 63 +++++++++++++++++++ backend/internal/handler/gateway_handler.go | 6 +- backend/internal/handler/gateway_helper.go | 3 + .../handler/gateway_helper_hotpath_test.go | 19 ++++++ .../handler/openai_gateway_handler.go | 6 +- 6 files changed, 118 insertions(+), 6 deletions(-) create mode 100644 backend/internal/handler/concurrency_error_response.go create mode 100644 backend/internal/handler/concurrency_error_response_test.go diff --git a/backend/internal/handler/concurrency_error_response.go b/backend/internal/handler/concurrency_error_response.go new file mode 100644 index 00000000..52abf735 --- /dev/null +++ b/backend/internal/handler/concurrency_error_response.go @@ -0,0 +1,27 @@ +package handler + +import ( + "context" + "errors" + "fmt" + "net/http" +) + +const statusClientClosedRequest = 499 + +func concurrencyErrorResponse(err error, slotType string) (int, string, string) { + var concurrencyErr *ConcurrencyError + if errors.As(err, &concurrencyErr) { + if concurrencyErr.SlotType != "" { + slotType = concurrencyErr.SlotType + } + return http.StatusTooManyRequests, "rate_limit_error", + fmt.Sprintf("Concurrency limit exceeded for %s, please retry later", slotType) + } + + if errors.Is(err, context.Canceled) { + return statusClientClosedRequest, "api_error", "context canceled" + } + + return http.StatusServiceUnavailable, "api_error", "Service temporarily unavailable, please retry later" +} diff --git a/backend/internal/handler/concurrency_error_response_test.go b/backend/internal/handler/concurrency_error_response_test.go new file mode 100644 index 00000000..a2e6b9ab --- /dev/null +++ b/backend/internal/handler/concurrency_error_response_test.go @@ -0,0 +1,63 @@ +package handler + +import ( + "context" + "errors" + "net/http" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestConcurrencyErrorResponse(t *testing.T) { + tests := []struct { + name string + err error + slotType string + wantStatus int + wantType string + wantMessage string + }{ + { + name: "true concurrency timeout remains rate limit", + err: &ConcurrencyError{SlotType: "account", IsTimeout: true}, + slotType: "user", + wantStatus: http.StatusTooManyRequests, + wantType: "rate_limit_error", + wantMessage: "Concurrency limit exceeded for account, please retry later", + }, + { + name: "client cancellation is not classified as concurrency limit", + err: context.Canceled, + slotType: "user", + wantStatus: statusClientClosedRequest, + wantType: "api_error", + wantMessage: "context canceled", + }, + { + name: "deadline exceeded is service unavailable", + err: context.DeadlineExceeded, + slotType: "user", + wantStatus: http.StatusServiceUnavailable, + wantType: "api_error", + wantMessage: "Service temporarily unavailable, please retry later", + }, + { + name: "redis acquire error is service unavailable", + err: errors.New("redis unavailable"), + slotType: "user", + wantStatus: http.StatusServiceUnavailable, + wantType: "api_error", + wantMessage: "Service temporarily unavailable, please retry later", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + status, errType, message := concurrencyErrorResponse(tt.err, tt.slotType) + require.Equal(t, tt.wantStatus, status) + require.Equal(t, tt.wantType, errType) + require.Equal(t, tt.wantMessage, message) + }) + } +} diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index 4695a791..a24611f9 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -1471,10 +1471,10 @@ func (h *GatewayHandler) calculateSubscriptionRemaining(group *service.Group, su return min } -// handleConcurrencyError handles concurrency-related errors with proper 429 response +// handleConcurrencyError handles concurrency-related acquire errors. func (h *GatewayHandler) handleConcurrencyError(c *gin.Context, err error, slotType string, streamStarted bool) { - h.handleStreamingAwareError(c, http.StatusTooManyRequests, "rate_limit_error", - fmt.Sprintf("Concurrency limit exceeded for %s, please retry later", slotType), streamStarted) + status, errType, message := concurrencyErrorResponse(err, slotType) + h.handleStreamingAwareError(c, status, errType, message, streamStarted) } func (h *GatewayHandler) handleFailoverExhausted(c *gin.Context, failoverErr *service.UpstreamFailoverError, platform string, streamStarted bool) { diff --git a/backend/internal/handler/gateway_helper.go b/backend/internal/handler/gateway_helper.go index 09e6c09b..e4897502 100644 --- a/backend/internal/handler/gateway_helper.go +++ b/backend/internal/handler/gateway_helper.go @@ -336,6 +336,9 @@ func (h *ConcurrencyHelper) waitForSlotWithPingTimeout(c *gin.Context, slotType for { select { case <-ctx.Done(): + if parentErr := c.Request.Context().Err(); parentErr != nil { + return nil, parentErr + } return nil, &ConcurrencyError{ SlotType: slotType, IsTimeout: true, diff --git a/backend/internal/handler/gateway_helper_hotpath_test.go b/backend/internal/handler/gateway_helper_hotpath_test.go index 4a677199..d57c396c 100644 --- a/backend/internal/handler/gateway_helper_hotpath_test.go +++ b/backend/internal/handler/gateway_helper_hotpath_test.go @@ -280,6 +280,25 @@ func TestWaitForSlotWithPingTimeout_TimeoutAndStreamPing(t *testing.T) { }) } +func TestWaitForSlotWithPingTimeout_ParentContextCanceled(t *testing.T) { + cache := &helperConcurrencyCacheStub{ + accountSeq: []bool{false}, + } + concurrency := service.NewConcurrencyService(cache) + helper := NewConcurrencyHelper(concurrency, SSEPingFormatNone, 5*time.Millisecond) + c, _ := newHelperTestContext(http.MethodPost, "/v1/messages") + reqCtx, cancel := context.WithCancel(c.Request.Context()) + c.Request = c.Request.WithContext(reqCtx) + cancel() + + streamStarted := false + release, err := helper.waitForSlotWithPingTimeout(c, "account", 101, 2, time.Second, false, &streamStarted, true) + require.Nil(t, release) + require.ErrorIs(t, err, context.Canceled) + var cErr *ConcurrencyError + require.False(t, errors.As(err, &cErr)) +} + func TestWaitForSlotWithPingTimeout_AcquireError(t *testing.T) { errCache := &helperConcurrencyCacheStubWithError{ err: errors.New("redis unavailable"), diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index a51eee86..2d0524cd 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -1685,10 +1685,10 @@ func (h *OpenAIGatewayHandler) acquireImageGenerationSlot(c *gin.Context, stream return nil, false } -// handleConcurrencyError handles concurrency-related errors with proper 429 response +// handleConcurrencyError handles concurrency-related acquire errors. func (h *OpenAIGatewayHandler) handleConcurrencyError(c *gin.Context, err error, slotType string, streamStarted bool) { - h.handleStreamingAwareError(c, http.StatusTooManyRequests, "rate_limit_error", - fmt.Sprintf("Concurrency limit exceeded for %s, please retry later", slotType), streamStarted) + status, errType, message := concurrencyErrorResponse(err, slotType) + h.handleStreamingAwareError(c, status, errType, message, streamStarted) } func (h *OpenAIGatewayHandler) handleFailoverExhausted(c *gin.Context, failoverErr *service.UpstreamFailoverError, streamStarted bool) { From ccace69d4e074929d2470b9fa995a9e9639b9b06 Mon Sep 17 00:00:00 2001 From: Wey Gu Date: Thu, 28 May 2026 19:39:52 +0800 Subject: [PATCH 23/79] Add OpenAI embeddings gateway --- backend/internal/handler/endpoint.go | 5 +- backend/internal/handler/endpoint_test.go | 2 + backend/internal/handler/openai_embeddings.go | 253 ++++++++++++++++++ backend/internal/server/routes/gateway.go | 26 ++ backend/internal/service/openai_embeddings.go | 240 +++++++++++++++++ .../service/openai_embeddings_test.go | 106 ++++++++ 6 files changed, 631 insertions(+), 1 deletion(-) create mode 100644 backend/internal/handler/openai_embeddings.go create mode 100644 backend/internal/service/openai_embeddings.go create mode 100644 backend/internal/service/openai_embeddings_test.go diff --git a/backend/internal/handler/endpoint.go b/backend/internal/handler/endpoint.go index db29618a..0d6f4b3c 100644 --- a/backend/internal/handler/endpoint.go +++ b/backend/internal/handler/endpoint.go @@ -17,6 +17,7 @@ import ( const ( EndpointMessages = "/v1/messages" EndpointChatCompletions = "/v1/chat/completions" + EndpointEmbeddings = "/v1/embeddings" EndpointResponses = "/v1/responses" EndpointImagesGenerations = "/v1/images/generations" EndpointImagesEdits = "/v1/images/edits" @@ -42,6 +43,8 @@ const ( func NormalizeInboundEndpoint(path string) string { path = strings.TrimSpace(path) switch { + case strings.Contains(path, EndpointEmbeddings): + return EndpointEmbeddings case strings.Contains(path, EndpointChatCompletions): return EndpointChatCompletions case strings.Contains(path, EndpointMessages): @@ -75,7 +78,7 @@ func DeriveUpstreamEndpoint(inbound, rawRequestPath, platform string) string { switch platform { case service.PlatformOpenAI: - if inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits { + if inbound == EndpointEmbeddings || inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits { return inbound } // OpenAI forwards everything to the Responses API. diff --git a/backend/internal/handler/endpoint_test.go b/backend/internal/handler/endpoint_test.go index 369c5fa7..42b6d6e7 100644 --- a/backend/internal/handler/endpoint_test.go +++ b/backend/internal/handler/endpoint_test.go @@ -24,6 +24,7 @@ func TestNormalizeInboundEndpoint(t *testing.T) { // Direct canonical paths. {"/v1/messages", EndpointMessages}, {"/v1/chat/completions", EndpointChatCompletions}, + {"/v1/embeddings", EndpointEmbeddings}, {"/v1/responses", EndpointResponses}, {"/v1/images/generations", EndpointImagesGenerations}, {"/v1/images/edits", EndpointImagesEdits}, @@ -77,6 +78,7 @@ func TestDeriveUpstreamEndpoint(t *testing.T) { {"openai responses nested", EndpointResponses, "/openai/v1/responses/compact/detail", service.PlatformOpenAI, "/v1/responses/compact/detail"}, {"openai from messages", EndpointMessages, "/v1/messages", service.PlatformOpenAI, EndpointResponses}, {"openai from completions", EndpointChatCompletions, "/v1/chat/completions", service.PlatformOpenAI, EndpointResponses}, + {"openai embeddings", EndpointEmbeddings, "/v1/embeddings", service.PlatformOpenAI, EndpointEmbeddings}, {"openai image generations", EndpointImagesGenerations, "/v1/images/generations", service.PlatformOpenAI, EndpointImagesGenerations}, {"openai image edits", EndpointImagesEdits, "/openai/v1/images/edits", service.PlatformOpenAI, EndpointImagesEdits}, diff --git a/backend/internal/handler/openai_embeddings.go b/backend/internal/handler/openai_embeddings.go new file mode 100644 index 00000000..bbb67044 --- /dev/null +++ b/backend/internal/handler/openai_embeddings.go @@ -0,0 +1,253 @@ +package handler + +import ( + "context" + "errors" + "net/http" + "strconv" + "strings" + "time" + + pkghttputil "github.com/Wei-Shaw/sub2api/internal/pkg/httputil" + "github.com/Wei-Shaw/sub2api/internal/pkg/ip" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/tidwall/gjson" + "go.uber.org/zap" +) + +// Embeddings handles the OpenAI-compatible Embeddings API. +// POST /v1/embeddings +func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) { + streamStarted := false + requestStart := time.Now() + + apiKey, ok := middleware2.GetAPIKeyFromContext(c) + if !ok { + h.errorResponse(c, http.StatusUnauthorized, "authentication_error", "Invalid API key") + return + } + + subject, ok := middleware2.GetAuthSubjectFromContext(c) + if !ok { + h.errorResponse(c, http.StatusInternalServerError, "api_error", "User context not found") + return + } + reqLog := requestLogger( + c, + "handler.openai_gateway.embeddings", + zap.Int64("user_id", subject.UserID), + zap.Int64("api_key_id", apiKey.ID), + zap.Any("group_id", apiKey.GroupID), + ) + if !h.ensureResponsesDependencies(c, reqLog) { + return + } + + body, err := pkghttputil.ReadRequestBodyWithPrealloc(c.Request) + if err != nil { + if maxErr, ok := extractMaxBytesError(err); ok { + h.errorResponse(c, http.StatusRequestEntityTooLarge, "invalid_request_error", buildBodyTooLargeMessage(maxErr.Limit)) + return + } + h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to read request body") + return + } + if len(body) == 0 { + h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Request body is empty") + return + } + if !gjson.ValidBytes(body) { + h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body") + return + } + + modelResult := gjson.GetBytes(body, "model") + if !modelResult.Exists() || modelResult.Type != gjson.String || strings.TrimSpace(modelResult.String()) == "" { + h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "model is required") + return + } + reqModel := modelResult.String() + reqLog = reqLog.With(zap.String("model", reqModel)) + setOpsRequestContext(c, reqModel, false) + setOpsEndpointContext(c, "", int16(service.RequestTypeSync)) + + channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, reqModel) + + subscription, _ := middleware2.GetSubscriptionFromContext(c) + service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds()) + + userReleaseFunc, acquired := h.acquireResponsesUserSlot(c, subject.UserID, subject.Concurrency, false, &streamStarted, reqLog) + if !acquired { + return + } + if userReleaseFunc != nil { + defer userReleaseFunc() + } + + if err := h.billingCacheService.CheckBillingEligibility(c.Request.Context(), apiKey.User, apiKey, apiKey.Group, subscription, service.QuotaPlatform(c.Request.Context(), apiKey)); err != nil { + reqLog.Info("openai_embeddings.billing_check_failed", zap.Error(err)) + status, code, message, retryAfter := billingErrorDetails(err) + if retryAfter > 0 { + c.Header("Retry-After", strconv.Itoa(retryAfter)) + } + h.errorResponse(c, status, code, message) + return + } + + failedAccountIDs := make(map[int64]struct{}) + var lastFailoverErr *service.UpstreamFailoverError + switchCount := 0 + maxAccountSwitches := h.maxAccountSwitches + if maxAccountSwitches <= 0 { + maxAccountSwitches = 3 + } + routingStart := time.Now() + + for { + selection, _, err := h.gatewayService.SelectAccountWithScheduler( + c.Request.Context(), + apiKey.GroupID, + "", + "", + reqModel, + failedAccountIDs, + service.OpenAIUpstreamTransportHTTPSSE, + false, + ) + if err != nil { + reqLog.Warn("openai_embeddings.account_select_failed", + zap.Error(err), + zap.Int("excluded_account_count", len(failedAccountIDs)), + ) + if len(failedAccountIDs) == 0 { + markOpsRoutingCapacityLimitedIfNoAvailable(c, err) + h.errorResponse(c, http.StatusServiceUnavailable, "api_error", "Service temporarily unavailable") + return + } + if lastFailoverErr != nil { + h.handleFailoverExhausted(c, lastFailoverErr, false) + } else { + h.errorResponse(c, http.StatusBadGateway, "api_error", "Upstream request failed") + } + return + } + if selection == nil || selection.Account == nil { + markOpsRoutingCapacityLimited(c) + h.errorResponse(c, http.StatusServiceUnavailable, "api_error", "No available accounts") + return + } + account := selection.Account + if account.Type != service.AccountTypeAPIKey { + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } + failedAccountIDs[account.ID] = struct{}{} + continue + } + setOpsSelectedAccount(c, account.ID, account.Platform) + + accountReleaseFunc, accountAcquired := h.acquireResponsesAccountSlot(c, apiKey.GroupID, "", selection, false, &streamStarted, reqLog) + if !accountAcquired { + return + } + + service.SetOpsLatencyMs(c, service.OpsRoutingLatencyMsKey, time.Since(routingStart).Milliseconds()) + forwardStart := time.Now() + + forwardBody := body + if channelMapping.Mapped { + forwardBody = h.gatewayService.ReplaceModelInBody(body, channelMapping.MappedModel) + } + writerSizeBeforeForward := c.Writer.Size() + result, err := func() (*service.OpenAIForwardResult, error) { + defer func() { + if accountReleaseFunc != nil { + accountReleaseFunc() + } + }() + return h.gatewayService.ForwardEmbeddings(c.Request.Context(), c, account, forwardBody, "") + }() + + forwardDurationMs := time.Since(forwardStart).Milliseconds() + upstreamLatencyMs, _ := getContextInt64(c, service.OpsUpstreamLatencyMsKey) + responseLatencyMs := forwardDurationMs + if upstreamLatencyMs > 0 && forwardDurationMs > upstreamLatencyMs { + responseLatencyMs = forwardDurationMs - upstreamLatencyMs + } + service.SetOpsLatencyMs(c, service.OpsResponseLatencyMsKey, responseLatencyMs) + + if err != nil { + var failoverErr *service.UpstreamFailoverError + if errors.As(err, &failoverErr) { + if c.Writer.Size() != writerSizeBeforeForward { + h.handleFailoverExhausted(c, failoverErr, true) + return + } + h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) + h.gatewayService.RecordOpenAIAccountSwitch() + failedAccountIDs[account.ID] = struct{}{} + lastFailoverErr = failoverErr + if switchCount >= maxAccountSwitches { + h.handleFailoverExhausted(c, failoverErr, false) + return + } + switchCount++ + reqLog.Warn("openai_embeddings.upstream_failover_switching", + zap.Int64("account_id", account.ID), + zap.Int("upstream_status", failoverErr.StatusCode), + zap.Int("switch_count", switchCount), + zap.Int("max_switches", maxAccountSwitches), + ) + continue + } + h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) + if c.Writer.Size() == writerSizeBeforeForward { + h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Upstream request failed") + } + reqLog.Warn("openai_embeddings.forward_failed", + zap.Int64("account_id", account.ID), + zap.Error(err), + ) + return + } + + h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, true, nil) + userAgent := c.GetHeader("User-Agent") + clientIP := ip.GetClientIP(c) + inboundEndpoint := GetInboundEndpoint(c) + upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) + + h.submitOpenAIUsageRecordTask(result, func(ctx context.Context) { + if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{ + Result: result, + APIKey: apiKey, + User: apiKey.User, + Account: account, + Subscription: subscription, + InboundEndpoint: inboundEndpoint, + UpstreamEndpoint: upstreamEndpoint, + UserAgent: userAgent, + IPAddress: clientIP, + APIKeyService: h.apiKeyService, + ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel), + }); err != nil { + logger.L().With( + zap.String("component", "handler.openai_gateway.embeddings"), + zap.Int64("user_id", subject.UserID), + zap.Int64("api_key_id", apiKey.ID), + zap.Any("group_id", apiKey.GroupID), + zap.String("model", reqModel), + zap.Int64("account_id", account.ID), + ).Error("openai_embeddings.record_usage_failed", zap.Error(err)) + } + }) + reqLog.Debug("openai_embeddings.request_completed", + zap.Int64("account_id", account.ID), + zap.Int("switch_count", switchCount), + ) + return + } +} diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index efc0687f..b039a6ec 100644 --- a/backend/internal/server/routes/gateway.go +++ b/backend/internal/server/routes/gateway.go @@ -89,6 +89,19 @@ func RegisterGatewayRoutes( } h.Gateway.ChatCompletions(c) }) + gateway.POST("/embeddings", func(c *gin.Context) { + if getGroupPlatform(c) != service.PlatformOpenAI { + service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) + c.JSON(http.StatusNotFound, gin.H{ + "error": gin.H{ + "type": "not_found_error", + "message": "Embeddings API is not supported for this platform", + }, + }) + return + } + h.OpenAIGateway.Embeddings(c) + }) gateway.POST("/images/generations", func(c *gin.Context) { if getGroupPlatform(c) != service.PlatformOpenAI { service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) @@ -158,6 +171,19 @@ func RegisterGatewayRoutes( } h.Gateway.ChatCompletions(c) }) + r.POST("/embeddings", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, func(c *gin.Context) { + if getGroupPlatform(c) != service.PlatformOpenAI { + service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) + c.JSON(http.StatusNotFound, gin.H{ + "error": gin.H{ + "type": "not_found_error", + "message": "Embeddings API is not supported for this platform", + }, + }) + return + } + h.OpenAIGateway.Embeddings(c) + }) r.POST("/images/generations", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, func(c *gin.Context) { if getGroupPlatform(c) != service.PlatformOpenAI { service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) diff --git a/backend/internal/service/openai_embeddings.go b/backend/internal/service/openai_embeddings.go new file mode 100644 index 00000000..359df3bb --- /dev/null +++ b/backend/internal/service/openai_embeddings.go @@ -0,0 +1,240 @@ +package service + +import ( + "bytes" + "context" + "errors" + "fmt" + "io" + "net/http" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" + "github.com/gin-gonic/gin" + "github.com/tidwall/gjson" + "go.uber.org/zap" +) + +func (s *OpenAIGatewayService) ForwardEmbeddings( + ctx context.Context, + c *gin.Context, + account *Account, + body []byte, + defaultMappedModel string, +) (*OpenAIForwardResult, error) { + startTime := time.Now() + + originalModel := strings.TrimSpace(gjson.GetBytes(body, "model").String()) + if originalModel == "" { + writeOpenAIEmbeddingsError(c, http.StatusBadRequest, "invalid_request_error", "model is required") + return nil, fmt.Errorf("missing model in request") + } + + billingModel := resolveOpenAIForwardModel(account, originalModel, defaultMappedModel) + upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel) + upstreamBody := body + if upstreamModel != originalModel { + upstreamBody = ReplaceModelInBody(body, upstreamModel) + } + + logger.L().Debug("openai embeddings: forwarding", + zap.Int64("account_id", account.ID), + zap.String("original_model", originalModel), + zap.String("billing_model", billingModel), + zap.String("upstream_model", upstreamModel), + ) + + apiKey := account.GetOpenAIApiKey() + if apiKey == "" { + return nil, fmt.Errorf("account %d missing api_key", account.ID) + } + baseURL := account.GetOpenAIBaseURL() + if baseURL == "" { + baseURL = "https://api.openai.com" + } + validatedURL, err := s.validateUpstreamBaseURL(baseURL) + if err != nil { + return nil, fmt.Errorf("invalid base_url: %w", err) + } + targetURL := buildOpenAIEmbeddingsURL(validatedURL) + + upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) + upstreamReq, err := http.NewRequestWithContext(upstreamCtx, http.MethodPost, targetURL, bytes.NewReader(upstreamBody)) + releaseUpstreamCtx() + if err != nil { + return nil, fmt.Errorf("build upstream request: %w", err) + } + upstreamReq = upstreamReq.WithContext(WithHTTPUpstreamProfile(upstreamReq.Context(), HTTPUpstreamProfileOpenAI)) + upstreamReq.Header.Set("Content-Type", "application/json") + upstreamReq.Header.Set("Authorization", "Bearer "+apiKey) + upstreamReq.Header.Set("Accept", "application/json") + for key, values := range c.Request.Header { + lowerKey := strings.ToLower(key) + if openaiCCRawAllowedHeaders[lowerKey] { + for _, v := range values { + upstreamReq.Header.Add(key, v) + } + } + } + if customUA := account.GetOpenAIUserAgent(); customUA != "" { + upstreamReq.Header.Set("user-agent", customUA) + } + + proxyURL := "" + if account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency) + if err != nil { + safeErr := sanitizeUpstreamErrorMessage(err.Error()) + setOpsUpstreamError(c, 0, safeErr, "") + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: 0, + Kind: "request_error", + Message: safeErr, + }) + writeOpenAIEmbeddingsError(c, http.StatusBadGateway, "upstream_error", "Upstream request failed") + return nil, fmt.Errorf("upstream request failed: %s", safeErr) + } + defer func() { _ = resp.Body.Close() }() + + if resp.StatusCode >= 400 { + respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20)) + _ = resp.Body.Close() + resp.Body = io.NopCloser(bytes.NewReader(respBody)) + + upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(respBody)) + upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) + if s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody) { + upstreamDetail := "" + if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { + maxBytes := s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes + if maxBytes <= 0 { + maxBytes = 2048 + } + upstreamDetail = truncateString(string(respBody), maxBytes) + } + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: resp.StatusCode, + UpstreamRequestID: resp.Header.Get("x-request-id"), + Kind: "failover", + Message: upstreamMsg, + Detail: upstreamDetail, + }) + s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, upstreamModel) + return nil, &UpstreamFailoverError{ + StatusCode: resp.StatusCode, + ResponseBody: respBody, + RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), + } + } + writeOpenAIEmbeddingsUpstreamResponse(c, resp, respBody, s.responseHeaderFilter) + return nil, fmt.Errorf("upstream returned status %d", resp.StatusCode) + } + + respBody, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError) + if err != nil { + if !errors.Is(err, ErrUpstreamResponseBodyTooLarge) { + writeOpenAIEmbeddingsError(c, http.StatusBadGateway, "api_error", "Failed to read upstream response") + } + return nil, fmt.Errorf("read upstream body: %w", err) + } + + writeOpenAIEmbeddingsUpstreamResponse(c, resp, respBody, s.responseHeaderFilter) + + return &OpenAIForwardResult{ + RequestID: firstNonEmptyString(resp.Header.Get("x-request-id"), resp.Header.Get("request-id")), + Usage: extractOpenAIEmbeddingsUsage(respBody), + Model: originalModel, + BillingModel: billingModel, + UpstreamModel: upstreamModel, + Stream: false, + Duration: time.Since(startTime), + }, nil +} + +func writeOpenAIEmbeddingsUpstreamResponse(c *gin.Context, resp *http.Response, body []byte, filter *responseheaders.CompiledHeaderFilter) { + if c == nil || resp == nil { + return + } + if c.Writer.Written() { + return + } + if resp.Header != nil { + responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, filter) + } + if ct := resp.Header.Get("Content-Type"); ct != "" { + c.Writer.Header().Set("Content-Type", ct) + } else { + c.Writer.Header().Set("Content-Type", "application/json") + } + c.Writer.WriteHeader(resp.StatusCode) + _, _ = c.Writer.Write(body) +} + +func writeOpenAIEmbeddingsError(c *gin.Context, statusCode int, errType, message string) { + c.JSON(statusCode, gin.H{ + "error": gin.H{ + "type": errType, + "message": message, + }, + }) +} + +func extractOpenAIEmbeddingsUsage(body []byte) OpenAIUsage { + usage := gjson.GetBytes(body, "usage") + if !usage.Exists() || !usage.IsObject() { + return OpenAIUsage{} + } + inputTokens := firstPositiveGJSONInt( + usage.Get("prompt_tokens"), + usage.Get("input_tokens"), + usage.Get("total_tokens"), + ) + outputTokens := firstPositiveGJSONInt( + usage.Get("completion_tokens"), + usage.Get("output_tokens"), + ) + cacheReadTokens := firstPositiveGJSONInt( + usage.Get("prompt_tokens_details.cached_tokens"), + usage.Get("input_tokens_details.cached_tokens"), + usage.Get("cache_read_tokens"), + usage.Get("cache_read_input_tokens"), + ) + cacheCreationTokens := firstPositiveGJSONInt( + usage.Get("cache_creation_tokens"), + usage.Get("cache_creation_input_tokens"), + usage.Get("input_tokens_details.cache_creation_tokens"), + ) + return OpenAIUsage{ + InputTokens: inputTokens, + OutputTokens: outputTokens, + CacheReadInputTokens: cacheReadTokens, + CacheCreationInputTokens: cacheCreationTokens, + } +} + +func firstPositiveGJSONInt(values ...gjson.Result) int { + for _, value := range values { + if !value.Exists() { + continue + } + n := int(value.Int()) + if n > 0 { + return n + } + } + return 0 +} + +func buildOpenAIEmbeddingsURL(base string) string { + return buildOpenAIEndpointURL(base, "/v1/embeddings") +} diff --git a/backend/internal/service/openai_embeddings_test.go b/backend/internal/service/openai_embeddings_test.go new file mode 100644 index 00000000..c7e89d64 --- /dev/null +++ b/backend/internal/service/openai_embeddings_test.go @@ -0,0 +1,106 @@ +package service + +import ( + "bytes" + "context" + "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 TestBuildOpenAIEmbeddingsURL(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + base string + want string + }{ + {"bare domain", "https://api.openai.com", "https://api.openai.com/v1/embeddings"}, + {"bare /v1", "https://api.openai.com/v1", "https://api.openai.com/v1/embeddings"}, + {"already embeddings", "https://api.openai.com/v1/embeddings", "https://api.openai.com/v1/embeddings"}, + {"third-party versioned path", "https://open.bigmodel.cn/api/paas/v4", "https://open.bigmodel.cn/api/paas/v4/embeddings"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + require.Equal(t, tt.want, buildOpenAIEmbeddingsURL(tt.base)) + }) + } +} + +func TestForwardEmbeddings_APIKeyPassthroughRecordsUsageAndBatchInput(t *testing.T) { + gin.SetMode(gin.TestMode) + + reqBody := []byte(`{ + "model":"nowledge-embedding", + "input":["hello","world"], + "encoding_format":"float", + "dimensions":256 + }`) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/embeddings", bytes.NewReader(reqBody)) + c.Request.Header.Set("Content-Type", "application/json") + + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"application/json"}, + "X-Request-Id": []string{"emb-rid"}, + }, + Body: io.NopCloser(strings.NewReader(`{ + "object":"list", + "data":[ + {"object":"embedding","index":0,"embedding":[0.1,0.2]}, + {"object":"embedding","index":1,"embedding":[0.3,0.4]} + ], + "model":"jina-embeddings-v5-text-small", + "usage":{"prompt_tokens":13,"total_tokens":13} + }`)), + }} + svc := &OpenAIGatewayService{ + cfg: &config.Config{}, + httpUpstream: upstream, + } + account := &Account{ + ID: 42, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://api.jina.ai", + "model_mapping": map[string]any{ + "nowledge-embedding": "jina-embeddings-v5-text-small", + }, + }, + } + + result, err := svc.ForwardEmbeddings(context.Background(), c, account, reqBody, "") + + require.NoError(t, err) + require.Equal(t, http.StatusOK, rec.Code) + require.NotNil(t, result) + require.Equal(t, "emb-rid", result.RequestID) + require.Equal(t, "nowledge-embedding", result.Model) + require.Equal(t, "jina-embeddings-v5-text-small", result.BillingModel) + require.Equal(t, "jina-embeddings-v5-text-small", result.UpstreamModel) + require.Equal(t, 13, result.Usage.InputTokens) + require.Equal(t, 0, result.Usage.OutputTokens) + require.Equal(t, "https://api.jina.ai/v1/embeddings", upstream.lastReq.URL.String()) + require.Equal(t, "Bearer sk-test", upstream.lastReq.Header.Get("Authorization")) + require.Equal(t, "jina-embeddings-v5-text-small", gjson.GetBytes(upstream.lastBody, "model").String()) + require.Equal(t, int64(2), gjson.GetBytes(upstream.lastBody, "input.#").Int()) + require.Equal(t, "hello", gjson.GetBytes(upstream.lastBody, "input.0").String()) + require.Equal(t, "world", gjson.GetBytes(upstream.lastBody, "input.1").String()) + require.Equal(t, "float", gjson.GetBytes(upstream.lastBody, "encoding_format").String()) + require.Equal(t, int64(256), gjson.GetBytes(upstream.lastBody, "dimensions").Int()) +} From 1b2d8873b0c1624b985e1fe7dfbce1fffff5b324 Mon Sep 17 00:00:00 2001 From: lyen1688 Date: Thu, 28 May 2026 20:05:24 +0800 Subject: [PATCH 24/79] =?UTF-8?q?feat:=20=E5=AE=8C=E5=96=84=E5=89=8D?= =?UTF-8?q?=E7=BD=AE=E6=8B=A6=E6=88=AA=E5=AE=A1=E6=A0=B8=E8=BF=90=E8=A1=8C?= =?UTF-8?q?=E6=80=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../handler/openai_gateway_handler_test.go | 30 +- .../repository/content_moderation_repo.go | 3 +- .../content_moderation_repo_test.go | 40 ++ .../internal/service/content_moderation.go | 422 +++++++++++++++--- .../service/content_moderation_test.go | 378 ++++++++++++++-- frontend/src/api/admin/riskControl.ts | 24 + .../i18n/__tests__/riskControlLocales.spec.ts | 24 + frontend/src/i18n/locales/en.ts | 27 +- frontend/src/i18n/locales/zh.ts | 27 +- frontend/src/views/admin/RiskControlView.vue | 176 +++++++- .../admin/__tests__/RiskControlView.spec.ts | 147 +++++- 11 files changed, 1193 insertions(+), 105 deletions(-) create mode 100644 backend/internal/repository/content_moderation_repo_test.go create mode 100644 frontend/src/i18n/__tests__/riskControlLocales.spec.ts diff --git a/backend/internal/handler/openai_gateway_handler_test.go b/backend/internal/handler/openai_gateway_handler_test.go index d7d21fac..7de30e9c 100644 --- a/backend/internal/handler/openai_gateway_handler_test.go +++ b/backend/internal/handler/openai_gateway_handler_test.go @@ -7,6 +7,7 @@ import ( "net/http" "net/http/httptest" "strings" + "sync" "testing" "time" @@ -740,16 +741,31 @@ func (r *contentModerationHandlerSettingRepo) Delete(ctx context.Context, key st } type contentModerationHandlerTestRepo struct { + mu sync.Mutex logs []service.ContentModerationLog } func (r *contentModerationHandlerTestRepo) CreateLog(ctx context.Context, log *service.ContentModerationLog) error { if log != nil { + r.mu.Lock() + defer r.mu.Unlock() r.logs = append(r.logs, *log) } return nil } +func (r *contentModerationHandlerTestRepo) resetLogs() { + r.mu.Lock() + defer r.mu.Unlock() + r.logs = nil +} + +func (r *contentModerationHandlerTestRepo) logSnapshot() []service.ContentModerationLog { + r.mu.Lock() + defer r.mu.Unlock() + return append([]service.ContentModerationLog(nil), r.logs...) +} + func (r *contentModerationHandlerTestRepo) ListLogs(ctx context.Context, filter service.ContentModerationLogFilter) ([]service.ContentModerationLog, *pagination.PaginationResult, error) { return nil, nil, nil } @@ -808,7 +824,10 @@ func TestOpenAIResponsesWebSocket_ContentModerationBlocksFirstFrame(t *testing.T }) require.NoError(t, err) require.True(t, decision.Blocked) - repo.logs = nil + require.Eventually(t, func() bool { + return len(repo.logSnapshot()) == 1 + }, time.Second, 10*time.Millisecond) + repo.resetLogs() h := &OpenAIGatewayHandler{ gatewayService: &service.OpenAIGatewayService{}, billingCacheService: &service.BillingCacheService{}, @@ -848,10 +867,11 @@ func TestOpenAIResponsesWebSocket_ContentModerationBlocksFirstFrame(t *testing.T require.Equal(t, coderws.StatusPolicyViolation, closeErr.Code) require.Contains(t, closeErr.Reason, "内容审计测试阻断") } - require.Len(t, repo.logs, 1) - require.True(t, repo.logs[0].Flagged) - require.Equal(t, service.ContentModerationActionBlock, repo.logs[0].Action) - require.Equal(t, "bad prompt", repo.logs[0].InputExcerpt) + logs := repo.logSnapshot() + require.Len(t, logs, 1) + require.True(t, logs[0].Flagged) + require.Equal(t, service.ContentModerationActionBlock, logs[0].Action) + require.Equal(t, "bad prompt", logs[0].InputExcerpt) } func TestOpenAIResponsesWebSocket_PassthroughUsageLogPersistsUserAgentAndReasoningEffort(t *testing.T) { diff --git a/backend/internal/repository/content_moderation_repo.go b/backend/internal/repository/content_moderation_repo.go index 6ada004a..9b19cce9 100644 --- a/backend/internal/repository/content_moderation_repo.go +++ b/backend/internal/repository/content_moderation_repo.go @@ -192,6 +192,7 @@ SELECT COUNT(*) FROM content_moderation_logs WHERE user_id = $1 AND flagged = TRUE + AND action <> 'hash_block' AND created_at >= $2 AND created_at > COALESCE((SELECT at FROM last_auto_ban), '-infinity'::timestamptz) `, userID, since).Scan(&count) @@ -246,7 +247,7 @@ func buildContentModerationLogWhere(filter service.ContentModerationLogFilter) ( case "hit", "flagged": where = append(where, "l.flagged = TRUE") case "blocked", "block": - where = append(where, "l.action = 'block'") + where = append(where, "l.action IN ('block', 'keyword_block', 'hash_block')") case "pass", "allow": where = append(where, "l.flagged = FALSE AND l.error = ''") case "error": diff --git a/backend/internal/repository/content_moderation_repo_test.go b/backend/internal/repository/content_moderation_repo_test.go new file mode 100644 index 00000000..6d5faa12 --- /dev/null +++ b/backend/internal/repository/content_moderation_repo_test.go @@ -0,0 +1,40 @@ +package repository + +import ( + "context" + "regexp" + "strings" + "testing" + "time" + + sqlmock "github.com/DATA-DOG/go-sqlmock" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +func TestBuildContentModerationLogWhere_BlockedIncludesAllBlockActions(t *testing.T) { + where, args := buildContentModerationLogWhere(service.ContentModerationLogFilter{Result: "blocked"}) + + require.Empty(t, args) + sql := strings.Join(where, " AND ") + require.Contains(t, sql, "l.action IN ('block', 'keyword_block', 'hash_block')") + require.NotContains(t, sql, "l.action = 'block'") +} + +func TestContentModerationRepositoryCountFlaggedByUserSince_ExcludesHashBlock(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer func() { _ = db.Close() }() + + repo := NewContentModerationRepository(db) + since := time.Now().Add(-time.Hour) + mock.ExpectQuery(regexp.QuoteMeta("AND action <> 'hash_block'")). + WithArgs(int64(1001), since). + WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(2)) + + count, err := repo.CountFlaggedByUserSince(context.Background(), 1001, since) + + require.NoError(t, err) + require.Equal(t, 2, count) + require.NoError(t, mock.ExpectationsWereMet()) +} diff --git a/backend/internal/service/content_moderation.go b/backend/internal/service/content_moderation.go index a5a84d7b..ee1fca41 100644 --- a/backend/internal/service/content_moderation.go +++ b/backend/internal/service/content_moderation.go @@ -211,6 +211,20 @@ type ContentModerationAPIKeyStatus struct { Configured bool `json:"configured"` } +type ContentModerationAPIKeyLoad struct { + Index int `json:"index"` + KeyHash string `json:"key_hash"` + Masked string `json:"masked"` + Status string `json:"status"` + Active int64 `json:"active"` + Total int64 `json:"total"` + Success int64 `json:"success"` + Errors int64 `json:"errors"` + AvgLatencyMS int64 `json:"avg_latency_ms"` + LastLatencyMS int `json:"last_latency_ms"` + LastHTTPStatus int `json:"last_http_status"` +} + type TestContentModerationAPIKeysInput struct { APIKeys []string `json:"api_keys"` BaseURL string `json:"base_url"` @@ -399,25 +413,35 @@ type ContentModerationCleanupResult struct { } type ContentModerationRuntimeStatus struct { - Enabled bool `json:"enabled"` - RiskControlEnabled bool `json:"risk_control_enabled"` - Mode string `json:"mode"` - WorkerCount int `json:"worker_count"` - MaxWorkers int `json:"max_workers"` - ActiveWorkers int `json:"active_workers"` - IdleWorkers int `json:"idle_workers"` - QueueSize int `json:"queue_size"` - QueueLength int `json:"queue_length"` - QueueUsagePercent float64 `json:"queue_usage_percent"` - Enqueued int64 `json:"enqueued"` - Dropped int64 `json:"dropped"` - Processed int64 `json:"processed"` - Errors int64 `json:"errors"` - APIKeyStatuses []ContentModerationAPIKeyStatus `json:"api_key_statuses"` - FlaggedHashCount int64 `json:"flagged_hash_count"` - LastCleanupAt *time.Time `json:"last_cleanup_at,omitempty"` - LastCleanupDeletedHit int64 `json:"last_cleanup_deleted_hit"` - LastCleanupDeletedNonHit int64 `json:"last_cleanup_deleted_non_hit"` + Enabled bool `json:"enabled"` + RiskControlEnabled bool `json:"risk_control_enabled"` + Mode string `json:"mode"` + WorkerCount int `json:"worker_count"` + MaxWorkers int `json:"max_workers"` + ActiveWorkers int `json:"active_workers"` + IdleWorkers int `json:"idle_workers"` + QueueSize int `json:"queue_size"` + QueueLength int `json:"queue_length"` + QueueUsagePercent float64 `json:"queue_usage_percent"` + Enqueued int64 `json:"enqueued"` + Dropped int64 `json:"dropped"` + Processed int64 `json:"processed"` + Errors int64 `json:"errors"` + PreBlockActive int `json:"pre_block_active"` + PreBlockChecked int64 `json:"pre_block_checked"` + PreBlockAllowed int64 `json:"pre_block_allowed"` + PreBlockBlocked int64 `json:"pre_block_blocked"` + PreBlockErrors int64 `json:"pre_block_errors"` + PreBlockAvgLatencyMS int64 `json:"pre_block_avg_latency_ms"` + PreBlockAPIKeyActive int64 `json:"pre_block_api_key_active"` + PreBlockAPIKeyAvailableCount int64 `json:"pre_block_api_key_available_count"` + PreBlockAPIKeyTotalCalls int64 `json:"pre_block_api_key_total_calls"` + PreBlockAPIKeyLoads []ContentModerationAPIKeyLoad `json:"pre_block_api_key_loads"` + APIKeyStatuses []ContentModerationAPIKeyStatus `json:"api_key_statuses"` + FlaggedHashCount int64 `json:"flagged_hash_count"` + LastCleanupAt *time.Time `json:"last_cleanup_at,omitempty"` + LastCleanupDeletedHit int64 `json:"last_cleanup_deleted_hit"` + LastCleanupDeletedNonHit int64 `json:"last_cleanup_deleted_non_hit"` } type ContentModerationUnbanUserResult struct { @@ -466,6 +490,12 @@ type ContentModerationService struct { asyncDropped atomic.Int64 asyncProcessed atomic.Int64 asyncErrors atomic.Int64 + preBlockActive atomic.Int64 + preBlockChecked atomic.Int64 + preBlockAllowed atomic.Int64 + preBlockBlocked atomic.Int64 + preBlockErrors atomic.Int64 + preBlockLatencyTotalMS atomic.Int64 lastCleanupUnix atomic.Int64 lastCleanupDeletedHit atomic.Int64 lastCleanupDeletedNonHit atomic.Int64 @@ -474,10 +504,14 @@ type ContentModerationService struct { } type contentModerationTask struct { - input ContentModerationCheckInput - content ContentModerationInput - inputHash string - enqueuedAt time.Time + input ContentModerationCheckInput + content ContentModerationInput + inputHash string + log *ContentModerationLog + config *ContentModerationConfig + recordHash bool + applySideEffects bool + enqueuedAt time.Time } type contentModerationKeyHealth struct { @@ -491,6 +525,11 @@ type contentModerationKeyHealth struct { LastLatencyMS int LastHTTPStatus int LastTested bool + SyncActive int64 + SyncTotal int64 + SyncSuccess int64 + SyncErrors int64 + SyncLatencyMS int64 } func NewContentModerationService( @@ -827,9 +866,11 @@ func (s *ContentModerationService) Check(ctx context.Context, input ContentModer "protocol", input.Protocol, "text_runes", len([]rune(content.Text)), "image_count", len(content.Images)) + hashText := content.Hash() if cfg.Mode == ContentModerationModePreBlock { if cfg.KeywordBlockingMode != ContentModerationKeywordModeAPIOnly && len(cfg.BlockedKeywords) > 0 { if keyword, hit := matchBlockedKeyword(content.Text, cfg.BlockedKeywords); hit { + s.recordPreBlockSyncMetric(0, ContentModerationActionKeywordBlock) slog.Info("content_moderation.keyword_block", "user_id", input.UserID, "api_key_id", input.APIKeyID, @@ -840,8 +881,7 @@ func (s *ContentModerationService) Check(ctx context.Context, input ContentModer "keyword", keyword) scores := map[string]float64{contentModerationKeywordCategory: 1.0} log := s.buildLog(input, cfg, ContentModerationActionKeywordBlock, true, contentModerationKeywordCategory, 1.0, scores, content.ExcerptText(), nil, nil, "") - s.applyFlaggedSideEffects(ctx, cfg, log) - _ = s.repo.CreateLog(ctx, log) + s.enqueueRecord(input, cfg, log, hashText, false, true) return &ContentModerationDecision{ Allowed: false, Blocked: true, @@ -856,6 +896,7 @@ func (s *ContentModerationService) Check(ctx context.Context, input ContentModer } } if cfg.KeywordBlockingMode == ContentModerationKeywordModeKeywordOnly { + s.recordPreBlockSyncMetric(0, ContentModerationActionAllow) slog.Info("content_moderation.skip_api_keyword_only", "user_id", input.UserID, "api_key_id", input.APIKeyID, @@ -865,13 +906,15 @@ func (s *ContentModerationService) Check(ctx context.Context, input ContentModer return allow, nil } } - hashText := content.Hash() if cfg.PreHashCheckEnabled && s.hashCache != nil { matched, err := s.hashCache.HasFlaggedInputHash(ctx, hashText) if err != nil { slog.Warn("content_moderation.hash_check_failed", "user_id", input.UserID, "endpoint", input.Endpoint, "error", err) } if matched { + if cfg.Mode == ContentModerationModePreBlock { + s.recordPreBlockSyncMetric(0, ContentModerationActionHashBlock) + } slog.Info("content_moderation.hash_block", "user_id", input.UserID, "api_key_id", input.APIKeyID, @@ -883,6 +926,9 @@ func (s *ContentModerationService) Check(ctx context.Context, input ContentModer if message != "" { message = fmt.Sprintf("%s(hash: %s)", message, hashText) } + scores := map[string]float64{"hash": 1.0} + log := s.buildLog(input, cfg, ContentModerationActionHashBlock, true, "hash", 1.0, scores, content.ExcerptText(), nil, nil, "") + s.enqueueRecord(input, cfg, log, hashText, false, false) return &ContentModerationDecision{ Allowed: false, Blocked: true, @@ -895,6 +941,9 @@ func (s *ContentModerationService) Check(ctx context.Context, input ContentModer } } if !cfg.shouldSample(hashText) { + if cfg.Mode == ContentModerationModePreBlock { + s.recordPreBlockSyncMetric(0, ContentModerationActionAllow) + } slog.Info("content_moderation.skip_sample_rate", "user_id", input.UserID, "api_key_id", input.APIKeyID, @@ -905,6 +954,9 @@ func (s *ContentModerationService) Check(ctx context.Context, input ContentModer return allow, nil } if len(cfg.apiKeys()) == 0 { + if cfg.Mode == ContentModerationModePreBlock { + s.recordPreBlockSyncMetric(0, ContentModerationActionError) + } slog.Warn("content_moderation.skip_no_audit_api_keys", "user_id", input.UserID, "api_key_id", input.APIKeyID, @@ -930,10 +982,18 @@ func (s *ContentModerationService) Check(ctx context.Context, input ContentModer func (s *ContentModerationService) checkSync(ctx context.Context, input ContentModerationCheckInput, cfg *ContentModerationConfig, content ContentModerationInput, hashText string, queueDelay *int, allowBlock bool) *ContentModerationDecision { allow := &ContentModerationDecision{Allowed: true, Action: ContentModerationActionAllow} + trackPreBlock := queueDelay == nil && allowBlock && cfg != nil && cfg.Mode == ContentModerationModePreBlock + if trackPreBlock { + s.preBlockActive.Add(1) + defer s.preBlockActive.Add(-1) + } start := time.Now() - result, err := s.callModeration(ctx, cfg, content.ModerationInput()) + result, err := s.callModeration(ctx, cfg, content.ModerationInput(), trackPreBlock) latency := int(time.Since(start).Milliseconds()) if err != nil { + if trackPreBlock { + s.recordPreBlockSyncMetric(latency, ContentModerationActionError) + } slog.Warn("content_moderation.audit_api_failed", "user_id", input.UserID, "api_key_id", input.APIKeyID, @@ -962,6 +1022,9 @@ func (s *ContentModerationService) checkSync(ctx context.Context, input ContentM action = ContentModerationActionBlock blocked = true } + if trackPreBlock { + s.recordPreBlockSyncMetric(latency, action) + } slog.Info("content_moderation.audit_result", "user_id", input.UserID, "api_key_id", input.APIKeyID, @@ -980,13 +1043,11 @@ func (s *ContentModerationService) checkSync(ctx context.Context, input ContentM "queue_delay_ms", queueDelay) if flagged || cfg.RecordNonHits { log := s.buildLog(input, cfg, action, flagged, highestCategory, highestScore, result.CategoryScores, content.ExcerptText(), &latency, queueDelay, "") - if flagged && s.hashCache != nil { - if err := s.hashCache.RecordFlaggedInputHash(ctx, hashText); err != nil { - slog.Warn("content_moderation.record_hash_failed", "user_id", input.UserID, "endpoint", input.Endpoint, "error", err) - } + if queueDelay == nil && cfg.Mode == ContentModerationModePreBlock { + s.enqueueRecord(input, cfg, log, hashText, flagged, flagged) + } else { + s.persistContentModerationLog(ctx, cfg, log, hashText, flagged, flagged) } - s.applyFlaggedSideEffects(ctx, cfg, log) - _ = s.repo.CreateLog(ctx, log) } if blocked { return &ContentModerationDecision{ @@ -1012,6 +1073,25 @@ func (s *ContentModerationService) checkSync(ctx context.Context, input ContentM } } +func (s *ContentModerationService) recordPreBlockSyncMetric(latencyMS int, action string) { + if s == nil { + return + } + s.preBlockChecked.Add(1) + if latencyMS < 0 { + latencyMS = 0 + } + s.preBlockLatencyTotalMS.Add(int64(latencyMS)) + switch action { + case ContentModerationActionBlock, ContentModerationActionHashBlock, ContentModerationActionKeywordBlock: + s.preBlockBlocked.Add(1) + case ContentModerationActionError: + s.preBlockErrors.Add(1) + default: + s.preBlockAllowed.Add(1) + } +} + func (s *ContentModerationService) enqueueAsync(input ContentModerationCheckInput, cfg *ContentModerationConfig, content ContentModerationInput, hashText string) { if s == nil || s.asyncQueue == nil { return @@ -1040,11 +1120,49 @@ func (s *ContentModerationService) enqueueAsync(input ContentModerationCheckInpu } } +func (s *ContentModerationService) enqueueRecord(input ContentModerationCheckInput, cfg *ContentModerationConfig, log *ContentModerationLog, inputHash string, recordHash bool, applySideEffects bool) { + if s == nil || s.asyncQueue == nil || log == nil { + return + } + queueSize := defaultContentModerationQueueSize + if cfg != nil && cfg.QueueSize > 0 { + queueSize = cfg.QueueSize + } + if len(s.asyncQueue) >= queueSize { + slog.Warn("content_moderation.record_queue_full", + "user_id", input.UserID, + "endpoint", input.Endpoint, + "action", log.Action, + "queue_size", queueSize) + s.asyncDropped.Add(1) + return + } + task := contentModerationTask{ + input: input, + inputHash: inputHash, + log: log, + config: cloneContentModerationConfig(cfg), + recordHash: recordHash, + applySideEffects: applySideEffects, + enqueuedAt: time.Now(), + } + select { + case s.asyncQueue <- task: + s.asyncEnqueued.Add(1) + default: + slog.Warn("content_moderation.record_queue_full", + "user_id", input.UserID, + "endpoint", input.Endpoint, + "action", log.Action) + s.asyncDropped.Add(1) + } +} + func (s *ContentModerationService) worker(id int) { for { ctx, cancel := context.WithTimeout(context.Background(), maxContentModerationTimeoutMS*time.Millisecond+10*time.Second) cfg, err := s.loadConfig(ctx) - if err != nil || !cfg.Enabled || cfg.Mode == ContentModerationModeOff || len(cfg.apiKeys()) == 0 || id >= cfg.WorkerCount { + if err != nil || id >= cfg.WorkerCount { cancel() time.Sleep(time.Second) continue @@ -1061,6 +1179,22 @@ func (s *ContentModerationService) worker(id int) { slog.Error("content_moderation.worker_panic", "worker_id", id, "recover", r) } }() + if task.log != nil { + s.asyncActive.Add(1) + defer s.asyncActive.Add(-1) + queueDelay := int(time.Since(task.enqueuedAt).Milliseconds()) + task.log.QueueDelayMS = &queueDelay + taskCfg := task.config + if taskCfg == nil { + taskCfg = cfg + } + s.persistContentModerationLog(ctx, taskCfg, task.log, task.inputHash, task.recordHash, task.applySideEffects) + s.asyncProcessed.Add(1) + return + } + if !cfg.Enabled || cfg.Mode == ContentModerationModeOff || len(cfg.apiKeys()) == 0 { + return + } if !cfg.includesGroup(task.input.GroupID) { return } @@ -1186,6 +1320,15 @@ func (s *ContentModerationService) GetStatus(ctx context.Context) (*ContentModer if active > cfg.WorkerCount { active = cfg.WorkerCount } + preBlockActive := int(s.preBlockActive.Load()) + if preBlockActive < 0 { + preBlockActive = 0 + } + preBlockChecked := s.preBlockChecked.Load() + preBlockAvgLatency := int64(0) + if preBlockChecked > 0 { + preBlockAvgLatency = s.preBlockLatencyTotalMS.Load() / preBlockChecked + } queueLength := 0 if s.asyncQueue != nil { queueLength = len(s.asyncQueue) @@ -1208,25 +1351,35 @@ func (s *ContentModerationService) GetStatus(ctx context.Context) (*ContentModer lastCleanupAt = &t } return &ContentModerationRuntimeStatus{ - Enabled: cfg.Enabled, - RiskControlEnabled: riskEnabled, - Mode: cfg.Mode, - WorkerCount: cfg.WorkerCount, - MaxWorkers: maxContentModerationWorkerCount, - ActiveWorkers: active, - IdleWorkers: cfg.WorkerCount - active, - QueueSize: cfg.QueueSize, - QueueLength: queueLength, - QueueUsagePercent: queueUsage, - Enqueued: s.asyncEnqueued.Load(), - Dropped: s.asyncDropped.Load(), - Processed: s.asyncProcessed.Load(), - Errors: s.asyncErrors.Load(), - APIKeyStatuses: s.apiKeyStatuses(cfg.apiKeys()), - FlaggedHashCount: flaggedHashCount, - LastCleanupAt: lastCleanupAt, - LastCleanupDeletedHit: s.lastCleanupDeletedHit.Load(), - LastCleanupDeletedNonHit: s.lastCleanupDeletedNonHit.Load(), + Enabled: cfg.Enabled, + RiskControlEnabled: riskEnabled, + Mode: cfg.Mode, + WorkerCount: cfg.WorkerCount, + MaxWorkers: maxContentModerationWorkerCount, + ActiveWorkers: active, + IdleWorkers: cfg.WorkerCount - active, + QueueSize: cfg.QueueSize, + QueueLength: queueLength, + QueueUsagePercent: queueUsage, + Enqueued: s.asyncEnqueued.Load(), + Dropped: s.asyncDropped.Load(), + Processed: s.asyncProcessed.Load(), + Errors: s.asyncErrors.Load(), + PreBlockActive: preBlockActive, + PreBlockChecked: preBlockChecked, + PreBlockAllowed: s.preBlockAllowed.Load(), + PreBlockBlocked: s.preBlockBlocked.Load(), + PreBlockErrors: s.preBlockErrors.Load(), + PreBlockAvgLatencyMS: preBlockAvgLatency, + PreBlockAPIKeyActive: s.preBlockAPIKeyActive(cfg.apiKeys()), + PreBlockAPIKeyAvailableCount: s.preBlockAPIKeyAvailableCount(cfg.apiKeys()), + PreBlockAPIKeyTotalCalls: s.preBlockAPIKeyTotalCalls(cfg.apiKeys()), + PreBlockAPIKeyLoads: s.preBlockAPIKeyLoads(cfg.apiKeys()), + APIKeyStatuses: s.apiKeyStatuses(cfg.apiKeys()), + FlaggedHashCount: flaggedHashCount, + LastCleanupAt: lastCleanupAt, + LastCleanupDeletedHit: s.lastCleanupDeletedHit.Load(), + LastCleanupDeletedNonHit: s.lastCleanupDeletedNonHit.Load(), }, nil } @@ -1325,7 +1478,7 @@ func (s *ContentModerationService) validateConfig(ctx context.Context, cfg *Cont return nil } -func (s *ContentModerationService) callModeration(ctx context.Context, cfg *ContentModerationConfig, input any) (*moderationAPIResult, error) { +func (s *ContentModerationService) callModeration(ctx context.Context, cfg *ContentModerationConfig, input any, trackKeyLoad ...bool) (*moderationAPIResult, error) { attempts := cfg.RetryCount + 1 if attempts <= 0 { attempts = 1 @@ -1333,6 +1486,7 @@ func (s *ContentModerationService) callModeration(ctx context.Context, cfg *Cont if attempts > maxContentModerationRetryCount+1 { attempts = maxContentModerationRetryCount + 1 } + trackLoad := len(trackKeyLoad) > 0 && trackKeyLoad[0] var lastErr error for attempt := 0; attempt < attempts; attempt++ { key, ok := s.nextUsableAPIKey(cfg) @@ -1340,14 +1494,23 @@ func (s *ContentModerationService) callModeration(ctx context.Context, cfg *Cont lastErr = errors.New("no moderation api key available") break } + if trackLoad { + s.beginModerationAPIKeyCall(key) + } start := time.Now() httpStatus := 0 result, err := s.callModerationOnceWithInput(ctx, cfg, key, input, &httpStatus) latency := int(time.Since(start).Milliseconds()) if err == nil { + if trackLoad { + s.finishModerationAPIKeyCall(key, latency, true) + } s.markAPIKeySuccess(key, latency, httpStatus) return result, nil } + if trackLoad { + s.finishModerationAPIKeyCall(key, latency, false) + } s.markAPIKeyError(key, err.Error(), latency, httpStatus) lastErr = err if httpStatus == http.StatusBadRequest { @@ -1452,10 +1615,32 @@ func (s *ContentModerationService) buildLog(input ContentModerationCheckInput, c } } -func (s *ContentModerationService) applyFlaggedSideEffects(ctx context.Context, cfg *ContentModerationConfig, log *ContentModerationLog) { - if s == nil || cfg == nil || log == nil || !log.Flagged || log.UserID == nil || *log.UserID <= 0 { +func (s *ContentModerationService) persistContentModerationLog(ctx context.Context, cfg *ContentModerationConfig, log *ContentModerationLog, hashText string, recordHash bool, applySideEffects bool) { + if s == nil || log == nil { return } + if recordHash && s.hashCache != nil { + if err := s.hashCache.RecordFlaggedInputHash(ctx, hashText); err != nil { + slog.Warn("content_moderation.record_hash_failed", "user_id", contentModerationEmailUserID(log), "endpoint", log.Endpoint, "error", err) + } + } + autoBanJustApplied := false + if applySideEffects { + autoBanJustApplied = s.applyFlaggedAccountSideEffects(ctx, cfg, log) + s.sendFlaggedNotificationSideEffects(ctx, cfg, log, autoBanJustApplied) + } + if s.repo != nil { + if err := s.repo.CreateLog(ctx, log); err != nil { + slog.Warn("content_moderation.create_log_failed", "user_id", contentModerationEmailUserID(log), "endpoint", log.Endpoint, "action", log.Action, "error", err) + return + } + } +} + +func (s *ContentModerationService) applyFlaggedAccountSideEffects(ctx context.Context, cfg *ContentModerationConfig, log *ContentModerationLog) bool { + if s == nil || cfg == nil || log == nil || !log.Flagged || log.UserID == nil || *log.UserID <= 0 { + return false + } count := 1 if s.repo != nil && cfg.ViolationWindowHours > 0 { since := time.Now().Add(-time.Duration(cfg.ViolationWindowHours) * time.Hour) @@ -1469,13 +1654,13 @@ func (s *ContentModerationService) applyFlaggedSideEffects(ctx context.Context, user, err := s.userRepo.GetByID(ctx, *log.UserID) if err != nil { slog.Warn("content_moderation.ban_get_user_failed", "user_id", *log.UserID, "error", err) - return + return false } if user.Status != StatusDisabled { user.Status = StatusDisabled if err := s.userRepo.Update(ctx, user); err != nil { slog.Warn("content_moderation.ban_update_user_failed", "user_id", *log.UserID, "error", err) - return + return false } if s.authCacheInvalidator != nil { s.authCacheInvalidator.InvalidateAuthCacheByUserID(ctx, *log.UserID) @@ -1484,7 +1669,13 @@ func (s *ContentModerationService) applyFlaggedSideEffects(ctx context.Context, } log.AutoBanned = true } + return autoBanJustApplied +} +func (s *ContentModerationService) sendFlaggedNotificationSideEffects(ctx context.Context, cfg *ContentModerationConfig, log *ContentModerationLog, autoBanJustApplied bool) { + if s == nil || cfg == nil || log == nil || !log.Flagged { + return + } if s.emailService == nil || strings.TrimSpace(log.UserEmail) == "" { return } @@ -1642,6 +1833,22 @@ func defaultContentModerationConfig() *ContentModerationConfig { } } +func cloneContentModerationConfig(cfg *ContentModerationConfig) *ContentModerationConfig { + if cfg == nil { + return nil + } + clone := *cfg + clone.APIKeys = append([]string(nil), cfg.APIKeys...) + clone.GroupIDs = append([]int64(nil), cfg.GroupIDs...) + clone.BlockedKeywords = append([]string(nil), cfg.BlockedKeywords...) + clone.Thresholds = cloneFloatMap(cfg.Thresholds) + clone.ModelFilter = ContentModerationModelFilter{ + Type: cfg.ModelFilter.Type, + Models: append([]string(nil), cfg.ModelFilter.Models...), + } + return &clone +} + func (cfg *ContentModerationConfig) normalize() { if cfg.APIKey != "" { cfg.APIKeys = normalizeModerationAPIKeys(append(cfg.APIKeys, cfg.APIKey)) @@ -1807,6 +2014,40 @@ func (s *ContentModerationService) isAPIKeyFrozen(key string, now time.Time) boo return state != nil && state.FrozenUntil.After(now) } +func (s *ContentModerationService) beginModerationAPIKeyCall(key string) { + hash := moderationAPIKeyHash(key) + if hash == "" || s == nil { + return + } + s.keyHealthMu.Lock() + defer s.keyHealthMu.Unlock() + state := s.ensureAPIKeyHealthLocked(hash, maskSecretTail(key)) + state.SyncActive++ +} + +func (s *ContentModerationService) finishModerationAPIKeyCall(key string, latencyMS int, success bool) { + hash := moderationAPIKeyHash(key) + if hash == "" || s == nil { + return + } + if latencyMS < 0 { + latencyMS = 0 + } + s.keyHealthMu.Lock() + defer s.keyHealthMu.Unlock() + state := s.ensureAPIKeyHealthLocked(hash, maskSecretTail(key)) + if state.SyncActive > 0 { + state.SyncActive-- + } + state.SyncTotal++ + state.SyncLatencyMS += int64(latencyMS) + if success { + state.SyncSuccess++ + return + } + state.SyncErrors++ +} + func (s *ContentModerationService) markAPIKeySuccess(key string, latencyMS int, httpStatus int) { hash := moderationAPIKeyHash(key) if hash == "" || s == nil { @@ -1926,6 +2167,71 @@ func (s *ContentModerationService) apiKeyStatuses(keys []string) []ContentModera return out } +func (s *ContentModerationService) preBlockAPIKeyLoads(keys []string) []ContentModerationAPIKeyLoad { + out := make([]ContentModerationAPIKeyLoad, 0, len(keys)) + for idx, key := range keys { + out = append(out, s.preBlockAPIKeyLoadForHash(idx, moderationAPIKeyHash(key), maskSecretTail(key))) + } + return out +} + +func (s *ContentModerationService) preBlockAPIKeyActive(keys []string) int64 { + var total int64 + for _, item := range s.preBlockAPIKeyLoads(keys) { + total += item.Active + } + return total +} + +func (s *ContentModerationService) preBlockAPIKeyAvailableCount(keys []string) int64 { + now := time.Now() + var count int64 + for _, key := range keys { + if !s.isAPIKeyFrozen(key, now) { + count++ + } + } + return count +} + +func (s *ContentModerationService) preBlockAPIKeyTotalCalls(keys []string) int64 { + var total int64 + for _, item := range s.preBlockAPIKeyLoads(keys) { + total += item.Total + } + return total +} + +func (s *ContentModerationService) preBlockAPIKeyLoadForHash(index int, hash string, masked string) ContentModerationAPIKeyLoad { + load := ContentModerationAPIKeyLoad{ + Index: index, + KeyHash: hash, + Masked: masked, + Status: "unknown", + } + status := s.apiKeyStatusForHash(index, hash, masked, true) + load.Status = status.Status + load.LastLatencyMS = status.LastLatencyMS + load.LastHTTPStatus = status.LastHTTPStatus + if hash == "" || s == nil { + return load + } + s.keyHealthMu.Lock() + defer s.keyHealthMu.Unlock() + state := s.keyHealth[hash] + if state == nil { + return load + } + load.Active = state.SyncActive + load.Total = state.SyncTotal + load.Success = state.SyncSuccess + load.Errors = state.SyncErrors + if state.SyncTotal > 0 { + load.AvgLatencyMS = state.SyncLatencyMS / state.SyncTotal + } + return load +} + func (s *ContentModerationService) apiKeyStatusForHash(index int, hash string, masked string, configured bool) ContentModerationAPIKeyStatus { status := ContentModerationAPIKeyStatus{ Index: index, diff --git a/backend/internal/service/content_moderation_test.go b/backend/internal/service/content_moderation_test.go index 20fce3ec..1fb72f36 100644 --- a/backend/internal/service/content_moderation_test.go +++ b/backend/internal/service/content_moderation_test.go @@ -3,9 +3,11 @@ package service import ( "context" "encoding/json" + "fmt" "net/http" "net/http/httptest" "strings" + "sync" "testing" "time" @@ -73,10 +75,13 @@ func (r *contentModerationTestSettingRepo) Delete(ctx context.Context, key strin } type contentModerationTestRepo struct { + mu sync.Mutex logs []ContentModerationLog } func (r *contentModerationTestRepo) CreateLog(ctx context.Context, log *ContentModerationLog) error { + r.mu.Lock() + defer r.mu.Unlock() if log != nil { r.logs = append(r.logs, *log) } @@ -88,14 +93,55 @@ func (r *contentModerationTestRepo) ListLogs(ctx context.Context, filter Content } func (r *contentModerationTestRepo) CountFlaggedByUserSince(ctx context.Context, userID int64, since time.Time) (int, error) { - return 0, nil + r.mu.Lock() + defer r.mu.Unlock() + count := 0 + for _, log := range r.logs { + if log.UserID == nil || *log.UserID != userID || !log.Flagged || log.Action == ContentModerationActionHashBlock { + continue + } + if log.CreatedAt.IsZero() || log.CreatedAt.Before(since) { + continue + } + count++ + } + return count, nil } func (r *contentModerationTestRepo) CleanupExpiredLogs(ctx context.Context, hitBefore time.Time, nonHitBefore time.Time) (*ContentModerationCleanupResult, error) { return &ContentModerationCleanupResult{}, nil } +func (r *contentModerationTestRepo) snapshotLogs() []ContentModerationLog { + r.mu.Lock() + defer r.mu.Unlock() + out := make([]ContentModerationLog, len(r.logs)) + copy(out, r.logs) + return out +} + +func requireContentModerationLogCount(t *testing.T, repo *contentModerationTestRepo, want int) []ContentModerationLog { + t.Helper() + var logs []ContentModerationLog + require.Eventually(t, func() bool { + logs = repo.snapshotLogs() + return len(logs) == want + }, time.Second, 10*time.Millisecond) + return logs +} + +func requireRecordedHashCount(t *testing.T, cache *contentModerationTestHashCache, want int) []string { + t.Helper() + var hashes []string + require.Eventually(t, func() bool { + hashes = cache.snapshotRecorded() + return len(hashes) == want + }, time.Second, 10*time.Millisecond) + return hashes +} + type contentModerationTestHashCache struct { + mu sync.Mutex hashes map[string]struct{} recorded []string checked []string @@ -246,6 +292,8 @@ func (i *contentModerationTestAuthCacheInvalidator) InvalidateAuthCacheByGroupID } func (c *contentModerationTestHashCache) RecordFlaggedInputHash(ctx context.Context, inputHash string) error { + c.mu.Lock() + defer c.mu.Unlock() if c.hashes == nil { c.hashes = map[string]struct{}{} } @@ -255,6 +303,8 @@ func (c *contentModerationTestHashCache) RecordFlaggedInputHash(ctx context.Cont } func (c *contentModerationTestHashCache) HasFlaggedInputHash(ctx context.Context, inputHash string) (bool, error) { + c.mu.Lock() + defer c.mu.Unlock() c.checked = append(c.checked, inputHash) if c.hasResultUsed { return c.hasResult, nil @@ -264,6 +314,8 @@ func (c *contentModerationTestHashCache) HasFlaggedInputHash(ctx context.Context } func (c *contentModerationTestHashCache) DeleteFlaggedInputHash(ctx context.Context, inputHash string) (bool, error) { + c.mu.Lock() + defer c.mu.Unlock() c.deleted = append(c.deleted, inputHash) if c.hashes == nil { return false, nil @@ -276,15 +328,50 @@ func (c *contentModerationTestHashCache) DeleteFlaggedInputHash(ctx context.Cont } func (c *contentModerationTestHashCache) ClearFlaggedInputHashes(ctx context.Context) (int64, error) { + c.mu.Lock() + defer c.mu.Unlock() deleted := int64(len(c.hashes)) c.hashes = map[string]struct{}{} return deleted, nil } func (c *contentModerationTestHashCache) CountFlaggedInputHashes(ctx context.Context) (int64, error) { + c.mu.Lock() + defer c.mu.Unlock() return int64(len(c.hashes)), nil } +func (c *contentModerationTestHashCache) snapshotRecorded() []string { + c.mu.Lock() + defer c.mu.Unlock() + out := make([]string, len(c.recorded)) + copy(out, c.recorded) + return out +} + +func (c *contentModerationTestHashCache) snapshotChecked() []string { + c.mu.Lock() + defer c.mu.Unlock() + out := make([]string, len(c.checked)) + copy(out, c.checked) + return out +} + +func (c *contentModerationTestHashCache) hasHash(inputHash string) bool { + c.mu.Lock() + defer c.mu.Unlock() + _, ok := c.hashes[inputHash] + return ok +} + +func (c *contentModerationTestHashCache) snapshotDeleted() []string { + c.mu.Lock() + defer c.mu.Unlock() + out := make([]string, len(c.deleted)) + copy(out, c.deleted) + return out +} + func TestBuildContentModerationLog_RedactsInputExcerpt(t *testing.T) { svc := &ContentModerationService{} cfg := defaultContentModerationConfig() @@ -381,10 +468,10 @@ func TestContentModerationCheck_PreBlockKeywordHitSkipsUpstreamCall(t *testing.T require.True(t, decision.Blocked) require.Equal(t, ContentModerationActionKeywordBlock, decision.Action) require.False(t, upstreamCalled, "keyword block must short-circuit upstream moderation call") - require.Len(t, repo.logs, 1) - require.True(t, repo.logs[0].Flagged) - require.Equal(t, ContentModerationActionKeywordBlock, repo.logs[0].Action) - require.Equal(t, contentModerationKeywordCategory, repo.logs[0].HighestCategory) + logs := requireContentModerationLogCount(t, repo, 1) + require.True(t, logs[0].Flagged) + require.Equal(t, ContentModerationActionKeywordBlock, logs[0].Action) + require.Equal(t, contentModerationKeywordCategory, logs[0].HighestCategory) } func TestContentModerationCheck_KeywordsIgnoredInObserveMode(t *testing.T) { @@ -474,7 +561,7 @@ func TestContentModerationCheck_KeywordOnlyStrategySkipsAPIOnMiss(t *testing.T) require.NoError(t, err) require.True(t, decision.Allowed, "keyword-only must allow misses without calling the API") require.False(t, upstreamCalled, "keyword-only must not call the upstream moderation API") - require.Len(t, repo.logs, 0) + require.Len(t, repo.snapshotLogs(), 0) } func TestContentModerationCheck_APIOnlyStrategyIgnoresKeywordList(t *testing.T) { @@ -545,7 +632,7 @@ func TestContentModerationCheck_ModelFilterAllAuditsEveryModel(t *testing.T) { require.True(t, decision.Blocked) require.Equal(t, ContentModerationActionKeywordBlock, decision.Action) } - require.Len(t, repo.logs, 2) + requireContentModerationLogCount(t, repo, 2) } func TestContentModerationCheck_ModelFilterIncludeOnlyAuditsListedModels(t *testing.T) { @@ -571,8 +658,8 @@ func TestContentModerationCheck_ModelFilterIncludeOnlyAuditsListedModels(t *test require.True(t, decision.Allowed) require.False(t, decision.Blocked) require.Equal(t, ContentModerationActionAllow, decision.Action) - require.Len(t, repo.logs, 1) - require.Equal(t, "gpt-5.5", repo.logs[0].Model) + logs := requireContentModerationLogCount(t, repo, 1) + require.Equal(t, "gpt-5.5", logs[0].Model) } func TestContentModerationCheck_ModelFilterExcludeSkipsListedModels(t *testing.T) { @@ -598,8 +685,8 @@ func TestContentModerationCheck_ModelFilterExcludeSkipsListedModels(t *testing.T require.True(t, decision.Allowed) require.False(t, decision.Blocked) require.Equal(t, ContentModerationActionAllow, decision.Action) - require.Len(t, repo.logs, 1) - require.Equal(t, "gpt-5.5", repo.logs[0].Model) + logs := requireContentModerationLogCount(t, repo, 1) + require.Equal(t, "gpt-5.5", logs[0].Model) } func TestContentModerationLoadConfig_LegacyConfigDefaultsModelFilterToAll(t *testing.T) { @@ -639,8 +726,8 @@ func TestContentModerationCheck_ModelFilterUsesRequestedModelNotBodyModel(t *tes require.NoError(t, err) require.True(t, decision.Blocked) require.Equal(t, ContentModerationActionKeywordBlock, decision.Action) - require.Len(t, repo.logs, 1) - require.Equal(t, "gpt-5.5", repo.logs[0].Model) + logs := requireContentModerationLogCount(t, repo, 1) + require.Equal(t, "gpt-5.5", logs[0].Model) } func defaultContentModerationModelFilterTestConfig() *ContentModerationConfig { @@ -939,11 +1026,11 @@ func TestContentModerationCheck_OpenAIResponsesRecordsNonHitForCodexPayload(t *t require.NoError(t, err) require.False(t, decision.Blocked) - require.Len(t, repo.logs, 1) - require.False(t, repo.logs[0].Flagged) - require.Equal(t, ContentModerationActionAllow, repo.logs[0].Action) - require.Equal(t, "/responses", repo.logs[0].Endpoint) - require.Equal(t, "last user prompt", repo.logs[0].InputExcerpt) + logs := requireContentModerationLogCount(t, repo, 1) + require.False(t, logs[0].Flagged) + require.Equal(t, ContentModerationActionAllow, logs[0].Action) + require.Equal(t, "/responses", logs[0].Endpoint) + require.Equal(t, "last user prompt", logs[0].InputExcerpt) require.Equal(t, "last user prompt", moderationRequest.Input) } @@ -1007,14 +1094,164 @@ func TestContentModerationCheck_PreBlockBlocksCodexResponsesLatestUserInput(t *t require.Equal(t, ContentModerationActionBlock, decision.Action) require.Equal(t, http.StatusUnavailableForLegalReasons, decision.StatusCode) require.Equal(t, "内容审计测试阻断", decision.Message) - require.Len(t, repo.logs, 1) - require.True(t, repo.logs[0].Flagged) - require.Equal(t, ContentModerationActionBlock, repo.logs[0].Action) - require.Equal(t, ContentModerationModePreBlock, repo.logs[0].Mode) - require.Equal(t, "latest blocked prompt", repo.logs[0].InputExcerpt) + logs := requireContentModerationLogCount(t, repo, 1) + require.True(t, logs[0].Flagged) + require.Equal(t, ContentModerationActionBlock, logs[0].Action) + require.Equal(t, ContentModerationModePreBlock, logs[0].Mode) + require.Equal(t, "latest blocked prompt", logs[0].InputExcerpt) require.Equal(t, "latest blocked prompt", moderationRequest.Input) } +func TestContentModerationStatusTracksPreBlockSyncMetrics(t *testing.T) { + var requestCount int + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requestCount++ + score := 0.01 + if requestCount == 1 { + score = 0.9 + } + time.Sleep(5 * time.Millisecond) + _ = json.NewEncoder(w).Encode(moderationAPIResponse{ + Results: []moderationAPIResult{{ + CategoryScores: map[string]float64{"sexual": score}, + }}, + }) + })) + defer server.Close() + + cfg := defaultContentModerationConfig() + cfg.Enabled = true + cfg.Mode = ContentModerationModePreBlock + cfg.BaseURL = server.URL + cfg.APIKeys = []string{"sk-test"} + rawCfg, err := json.Marshal(cfg) + require.NoError(t, err) + + svc := NewContentModerationService( + &contentModerationTestSettingRepo{values: map[string]string{ + SettingKeyRiskControlEnabled: "true", + SettingKeyContentModerationConfig: string(rawCfg), + }}, + &contentModerationTestRepo{}, + &contentModerationTestHashCache{}, + nil, + nil, + nil, + nil, + ) + + for _, prompt := range []string{"blocked prompt", "clean prompt"} { + _, err := svc.Check(context.Background(), ContentModerationCheckInput{ + UserID: 1001, + Protocol: ContentModerationProtocolOpenAIChat, + Body: []byte(fmt.Sprintf(`{"messages":[{"role":"user","content":%q}]}`, prompt)), + }) + require.NoError(t, err) + } + + status, err := svc.GetStatus(context.Background()) + require.NoError(t, err) + require.Equal(t, int64(2), status.PreBlockChecked) + require.Equal(t, int64(1), status.PreBlockAllowed) + require.Equal(t, int64(1), status.PreBlockBlocked) + require.Equal(t, int64(0), status.PreBlockErrors) + require.Equal(t, 0, status.PreBlockActive) + require.GreaterOrEqual(t, status.PreBlockAvgLatencyMS, int64(1)) +} + +func TestContentModerationStatusTracksPreBlockAPIKeyLoad(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _ = json.NewEncoder(w).Encode(moderationAPIResponse{ + Results: []moderationAPIResult{{ + CategoryScores: map[string]float64{"sexual": 0.01}, + }}, + }) + })) + defer server.Close() + + cfg := defaultContentModerationConfig() + cfg.Enabled = true + cfg.Mode = ContentModerationModePreBlock + cfg.BaseURL = server.URL + cfg.APIKeys = []string{"sk-one", "sk-two"} + rawCfg, err := json.Marshal(cfg) + require.NoError(t, err) + + svc := NewContentModerationService( + &contentModerationTestSettingRepo{values: map[string]string{ + SettingKeyRiskControlEnabled: "true", + SettingKeyContentModerationConfig: string(rawCfg), + }}, + &contentModerationTestRepo{}, + &contentModerationTestHashCache{}, + nil, + nil, + nil, + nil, + ) + + for idx := 0; idx < 4; idx++ { + _, err := svc.Check(context.Background(), ContentModerationCheckInput{ + UserID: 1001, + Protocol: ContentModerationProtocolOpenAIChat, + Body: []byte(fmt.Sprintf(`{"messages":[{"role":"user","content":"prompt %d"}]}`, idx)), + }) + require.NoError(t, err) + } + + status, err := svc.GetStatus(context.Background()) + require.NoError(t, err) + require.Len(t, status.PreBlockAPIKeyLoads, 2) + require.Equal(t, int64(4), status.PreBlockAPIKeyTotalCalls) + require.Equal(t, int64(2), status.PreBlockAPIKeyAvailableCount) + require.Equal(t, int64(0), status.PreBlockAPIKeyActive) + require.Equal(t, int64(0), status.PreBlockAPIKeyLoads[0].Active) + require.Equal(t, int64(2), status.PreBlockAPIKeyLoads[0].Total) + require.Equal(t, int64(2), status.PreBlockAPIKeyLoads[0].Success) + require.Equal(t, int64(0), status.PreBlockAPIKeyLoads[0].Errors) + require.Equal(t, int64(2), status.PreBlockAPIKeyLoads[1].Total) + require.Equal(t, int64(2), status.PreBlockAPIKeyLoads[1].Success) +} + +func TestContentModerationStatusTracksPreBlockLocalBlocks(t *testing.T) { + cfg := defaultContentModerationConfig() + cfg.Enabled = true + cfg.Mode = ContentModerationModePreBlock + cfg.KeywordBlockingMode = ContentModerationKeywordModeKeywordOnly + cfg.BlockedKeywords = []string{"blocked"} + rawCfg, err := json.Marshal(cfg) + require.NoError(t, err) + + svc := NewContentModerationService( + &contentModerationTestSettingRepo{values: map[string]string{ + SettingKeyRiskControlEnabled: "true", + SettingKeyContentModerationConfig: string(rawCfg), + }}, + &contentModerationTestRepo{}, + &contentModerationTestHashCache{}, + nil, + nil, + nil, + nil, + ) + + for _, prompt := range []string{"blocked prompt", "clean prompt"} { + _, err := svc.Check(context.Background(), ContentModerationCheckInput{ + UserID: 1001, + Protocol: ContentModerationProtocolOpenAIChat, + Body: []byte(fmt.Sprintf(`{"messages":[{"role":"user","content":%q}]}`, prompt)), + }) + require.NoError(t, err) + } + + status, err := svc.GetStatus(context.Background()) + require.NoError(t, err) + require.Equal(t, int64(2), status.PreBlockChecked) + require.Equal(t, int64(1), status.PreBlockAllowed) + require.Equal(t, int64(1), status.PreBlockBlocked) + require.Equal(t, int64(0), status.PreBlockErrors) +} + func TestBuildContentModerationTestAuditResult_UsesConfiguredThresholdsOnly(t *testing.T) { result := buildContentModerationTestAuditResult(&moderationAPIResult{ Flagged: true, @@ -1137,6 +1374,8 @@ func TestContentModerationCheck_PreHashUsesRedisHashCache(t *testing.T) { cfg.APIKeys = []string{"sk-test"} cfg.BlockStatus = http.StatusConflict cfg.BlockMessage = "命中历史风险输入" + cfg.AutoBanEnabled = true + cfg.BanThreshold = 1 rawCfg, err := json.Marshal(cfg) require.NoError(t, err) @@ -1145,20 +1384,23 @@ func TestContentModerationCheck_PreHashUsesRedisHashCache(t *testing.T) { content.Normalize() hashCache.hashes[content.Hash()] = struct{}{} + repo := &contentModerationTestRepo{} + userRepo := &contentModerationTestUserRepo{user: &User{ID: 1001, Status: StatusActive}} svc := NewContentModerationService( &contentModerationTestSettingRepo{values: map[string]string{ SettingKeyRiskControlEnabled: "true", SettingKeyContentModerationConfig: string(rawCfg), }}, - &contentModerationTestRepo{}, + repo, hashCache, nil, - nil, + userRepo, nil, nil, ) decision, err := svc.Check(context.Background(), ContentModerationCheckInput{ + UserID: 1001, Protocol: ContentModerationProtocolOpenAIChat, Body: []byte(`{"messages":[{"role":"user","content":"blocked prompt"}]}`), }) @@ -1169,7 +1411,73 @@ func TestContentModerationCheck_PreHashUsesRedisHashCache(t *testing.T) { require.Equal(t, content.Hash(), decision.InputHash) require.Contains(t, decision.Message, "命中历史风险输入") require.Contains(t, decision.Message, content.Hash()) - require.Len(t, hashCache.checked, 1) + require.Len(t, hashCache.snapshotChecked(), 1) + logs := requireContentModerationLogCount(t, repo, 1) + require.True(t, logs[0].Flagged) + require.Equal(t, ContentModerationActionHashBlock, logs[0].Action) + require.Equal(t, 1.0, logs[0].CategoryScores["hash"]) + require.Equal(t, ContentModerationModePreBlock, logs[0].Mode) + require.Zero(t, logs[0].ViolationCount) + require.False(t, logs[0].AutoBanned) + require.Empty(t, userRepo.updated) +} + +func TestContentModerationCheck_HashBlockLogsDoNotIncreaseNextViolationCount(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _ = json.NewEncoder(w).Encode(moderationAPIResponse{ + Results: []moderationAPIResult{{ + CategoryScores: map[string]float64{"sexual": 0.9}, + }}, + }) + })) + defer server.Close() + + cfg := defaultContentModerationConfig() + cfg.Enabled = true + cfg.Mode = ContentModerationModePreBlock + cfg.BaseURL = server.URL + cfg.APIKeys = []string{"sk-test"} + cfg.AutoBanEnabled = false + rawCfg, err := json.Marshal(cfg) + require.NoError(t, err) + + userID := int64(1001) + repo := &contentModerationTestRepo{} + hashLog := &ContentModerationLog{ + UserID: &userID, + Action: ContentModerationActionHashBlock, + Flagged: true, + HighestCategory: "hash", + HighestScore: 1, + CreatedAt: time.Now(), + } + require.NoError(t, repo.CreateLog(context.Background(), hashLog)) + + svc := NewContentModerationService( + &contentModerationTestSettingRepo{values: map[string]string{ + SettingKeyRiskControlEnabled: "true", + SettingKeyContentModerationConfig: string(rawCfg), + }}, + repo, + &contentModerationTestHashCache{}, + nil, + nil, + nil, + nil, + ) + + decision, err := svc.Check(context.Background(), ContentModerationCheckInput{ + UserID: userID, + Protocol: ContentModerationProtocolOpenAIChat, + Body: []byte(`{"messages":[{"role":"user","content":"new blocked prompt"}]}`), + }) + + require.NoError(t, err) + require.True(t, decision.Blocked) + logs := requireContentModerationLogCount(t, repo, 2) + require.Equal(t, ContentModerationActionHashBlock, logs[0].Action) + require.Equal(t, ContentModerationActionBlock, logs[1].Action) + require.Equal(t, 1, logs[1].ViolationCount) } func TestContentModerationCheck_PreBlockFlaggedWritesRedisHashCache(t *testing.T) { @@ -1219,8 +1527,8 @@ func TestContentModerationCheck_PreBlockFlaggedWritesRedisHashCache(t *testing.T require.True(t, decision.Blocked) require.Equal(t, ContentModerationActionBlock, decision.Action) require.Equal(t, 1, requestCount) - require.Len(t, hashCache.recorded, 1) - require.Len(t, repo.logs, 1) + recorded := requireRecordedHashCount(t, hashCache, 1) + requireContentModerationLogCount(t, repo, 1) decision, err = svc.Check(context.Background(), ContentModerationCheckInput{ Protocol: ContentModerationProtocolOpenAIChat, @@ -1229,9 +1537,11 @@ func TestContentModerationCheck_PreBlockFlaggedWritesRedisHashCache(t *testing.T require.NoError(t, err) require.True(t, decision.Blocked) require.Equal(t, ContentModerationActionHashBlock, decision.Action) - require.Equal(t, hashCache.recorded[0], decision.InputHash) + require.Equal(t, recorded[0], decision.InputHash) require.Equal(t, 1, requestCount) - require.Len(t, repo.logs, 1) + logs := requireContentModerationLogCount(t, repo, 2) + require.Equal(t, ContentModerationActionBlock, logs[0].Action) + require.Equal(t, ContentModerationActionHashBlock, logs[1].Action) } func TestContentModerationDeleteFlaggedInputHash_NormalizesAndDeletes(t *testing.T) { @@ -1246,8 +1556,8 @@ func TestContentModerationDeleteFlaggedInputHash_NormalizesAndDeletes(t *testing require.NoError(t, err) require.Equal(t, existingHash, result.InputHash) require.True(t, result.Deleted) - require.NotContains(t, hashCache.hashes, existingHash) - require.Equal(t, []string{existingHash}, hashCache.deleted) + require.False(t, hashCache.hasHash(existingHash)) + require.Equal(t, []string{existingHash}, hashCache.snapshotDeleted()) result, err = svc.DeleteFlaggedInputHash(context.Background(), existingHash) @@ -1327,8 +1637,8 @@ func TestContentModerationCheck_AsyncFlaggedWritesRedisHashCache(t *testing.T) { }, cfg, ContentModerationInput{Text: "bad prompt"}, strings.Repeat("b", 64), contentModerationIntPtr(25), false) require.False(t, decision.Blocked) - require.Len(t, hashCache.recorded, 1) - require.Len(t, repo.logs, 1) + requireRecordedHashCount(t, hashCache, 1) + requireContentModerationLogCount(t, repo, 1) } func TestBuildContentModerationAccountDisabledEmailBody_ContainsBanDetails(t *testing.T) { diff --git a/frontend/src/api/admin/riskControl.ts b/frontend/src/api/admin/riskControl.ts index 521114c2..aefd1618 100644 --- a/frontend/src/api/admin/riskControl.ts +++ b/frontend/src/api/admin/riskControl.ts @@ -132,6 +132,16 @@ export interface ContentModerationRuntimeStatus { dropped: number processed: number errors: number + pre_block_active: number + pre_block_checked: number + pre_block_allowed: number + pre_block_blocked: number + pre_block_errors: number + pre_block_avg_latency_ms: number + pre_block_api_key_active: number + pre_block_api_key_available_count: number + pre_block_api_key_total_calls: number + pre_block_api_key_loads: ContentModerationAPIKeyLoad[] api_key_statuses: ContentModerationAPIKeyStatus[] flagged_hash_count: number last_cleanup_at?: string @@ -139,6 +149,20 @@ export interface ContentModerationRuntimeStatus { last_cleanup_deleted_non_hit: number } +export interface ContentModerationAPIKeyLoad { + index: number + key_hash: string + masked: string + status: ContentModerationAPIKeyStatusValue + active: number + total: number + success: number + errors: number + avg_latency_ms: number + last_latency_ms: number + last_http_status: number +} + export interface ContentModerationLog { id: number request_id: string diff --git a/frontend/src/i18n/__tests__/riskControlLocales.spec.ts b/frontend/src/i18n/__tests__/riskControlLocales.spec.ts new file mode 100644 index 00000000..eab94fe6 --- /dev/null +++ b/frontend/src/i18n/__tests__/riskControlLocales.spec.ts @@ -0,0 +1,24 @@ +import { describe, expect, it } from 'vitest' + +import en from '../locales/en' +import zh from '../locales/zh' + +describe('risk control locale copy', () => { + it('describes worker runtime as audit and pre-block record processing', () => { + expect(zh.admin.riskControl.workerStatusHint).toContain('前置拦截记录任务') + expect(zh.admin.riskControl.workerStatusHint).not.toContain('异步观察任务') + expect(en.admin.riskControl.workerStatusHint).toContain('pre-block record tasks') + expect(en.admin.riskControl.workerStatusHint).not.toContain('observation tasks') + }) + + it('keeps pre-block audit key summary aware of async worker load', () => { + expect(zh.admin.riskControl.preBlockAPIKeyLoadSummary).toContain('worker:{workerActive} / {workerTotal}') + expect(en.admin.riskControl.preBlockAPIKeyLoadSummary).toContain('worker: {workerActive} / {workerTotal}') + }) + + it('does not describe pre-block audit key polling as bypassing the worker pool', () => { + expect(zh.admin.riskControl.preBlockAPIKeyLoadHint).toBe('同步前置拦截直接轮询可用审核 Key。') + expect(zh.admin.riskControl.preBlockAPIKeyLoadHint).not.toContain('Worker 池') + expect(en.admin.riskControl.preBlockAPIKeyLoadHint).not.toContain('worker pool') + }) +}) diff --git a/frontend/src/i18n/locales/en.ts b/frontend/src/i18n/locales/en.ts index ff5ea651..41c3c495 100644 --- a/frontend/src/i18n/locales/en.ts +++ b/frontend/src/i18n/locales/en.ts @@ -2599,14 +2599,37 @@ export default { modelFilterIncludeSummary: 'Applies to {count} models', modelFilterExcludeSummary: 'Excludes {count} models', emptyLogs: 'No audit records', + preBlockSyncStatus: 'Pre-Block Sync Status', + preBlockSyncHint: 'Live counters for the synchronous moderation path, excluding async record tasks.', + preBlockActive: 'Sync Processing', + preBlockActiveHint: 'Currently checking', + preBlockChecked: 'Checked', + preBlockCheckedHint: 'Entered pre-block path', + preBlockAllowed: 'Allowed', + preBlockAllowedHint: 'No block triggered', + preBlockBlocked: 'Blocked', + preBlockBlockedHint: 'Rejected after hit', + preBlockErrors: 'Audit Errors', + preBlockErrorsHint: 'Failed or no usable key', + preBlockAvgLatency: 'Avg Latency', + preBlockAvgLatencyHint: 'Synchronous path average', + preBlockAPIKeyLoad: 'Audit Key Load', + preBlockAPIKeyLoadHint: 'Synchronous pre-block checks round-robin usable audit keys directly.', + preBlockAPIKeyLoadSummary: 'Sync active {active} / usable keys {available}, {total} total, worker: {workerActive} / {workerTotal}', + preBlockAPIKeyTotals: 'Total {total}, success {success}, errors {errors}', + preBlockAPIKeyLoadEmpty: 'No audit key load data yet', + preBlockKeyActiveShort: 'Active', + preBlockKeyTotalShort: 'Total', + preBlockKeyAvgShort: 'Avg', + preBlockKeyLastShort: 'Last', workerStatus: 'Worker Runtime', - workerStatusHint: 'Queue and worker pool status for asynchronous observation tasks.', + workerStatusHint: 'Queue and worker pool status for async audit tasks and pre-block record tasks, excluding synchronous pre-block checks.', workerPool: 'Worker Pool', workerPoolMeta: '{active} processing, {idle} idle and ready, {total} total', queueUsage: 'Queue Usage', activeWorkers: 'Processing', idleWorkers: 'Idle Ready', - workerActive: 'Processing an asynchronous audit task', + workerActive: 'Processing an async audit or record task', workerIdle: 'Started, idle and ready', workerDisabled: 'Risk control or content audit is disabled', processed: 'Processed', diff --git a/frontend/src/i18n/locales/zh.ts b/frontend/src/i18n/locales/zh.ts index b8ac7d2c..8ff8ea80 100644 --- a/frontend/src/i18n/locales/zh.ts +++ b/frontend/src/i18n/locales/zh.ts @@ -2676,14 +2676,37 @@ export default { modelFilterIncludeSummary: '仅 {count} 个模型生效', modelFilterExcludeSummary: '排除 {count} 个模型', emptyLogs: '暂无审核记录', + preBlockSyncStatus: '前置拦截同步状态', + preBlockSyncHint: '同步审核链路的实时计数,不包含异步写记录任务。', + preBlockActive: '同步处理中', + preBlockActiveHint: '当前正在审核', + preBlockChecked: '已检查', + preBlockCheckedHint: '进入前置拦截链路', + preBlockAllowed: '已放行', + preBlockAllowedHint: '未触发拦截', + preBlockBlocked: '已拦截', + preBlockBlockedHint: '命中后拒绝请求', + preBlockErrors: '审核异常', + preBlockErrorsHint: '失败或无可用 Key', + preBlockAvgLatency: '平均耗时', + preBlockAvgLatencyHint: '同步链路平均值', + preBlockAPIKeyLoad: '审核 Key 负载', + preBlockAPIKeyLoadHint: '同步前置拦截直接轮询可用审核 Key。', + preBlockAPIKeyLoadSummary: '同步并发 {active} / 可用 Key {available},累计 {total} 次,worker:{workerActive} / {workerTotal}', + preBlockAPIKeyTotals: '累计 {total},成功 {success},异常 {errors}', + preBlockAPIKeyLoadEmpty: '暂无审核 Key 负载数据', + preBlockKeyActiveShort: '并发', + preBlockKeyTotalShort: '累计', + preBlockKeyAvgShort: '平均', + preBlockKeyLastShort: '最近', workerStatus: 'Worker 运行状态', - workerStatusHint: '异步观察任务的队列和 worker 池状态。', + workerStatusHint: '异步审计任务和前置拦截记录任务的队列与 Worker 池状态,不包含同步前置拦截审核请求。', workerPool: 'Worker 池', workerPoolMeta: '{active} 个处理中,{idle} 个空闲可用,共 {total} 个', queueUsage: '队列占用', activeWorkers: '处理中', idleWorkers: '空闲可用', - workerActive: '正在处理异步审计任务', + workerActive: '正在处理异步审计或记录任务', workerIdle: '已启动,当前空闲可用', workerDisabled: '风控或内容审计未启用', processed: '已处理', diff --git a/frontend/src/views/admin/RiskControlView.vue b/frontend/src/views/admin/RiskControlView.vue index 36a04756..b6d62767 100644 --- a/frontend/src/views/admin/RiskControlView.vue +++ b/frontend/src/views/admin/RiskControlView.vue @@ -53,7 +53,105 @@ -
+
+
+
+
+

{{ t('admin.riskControl.preBlockSyncStatus') }}

+

{{ t('admin.riskControl.preBlockSyncHint') }}

+
+ + {{ modeLabel(status?.mode ?? configForm.mode) }} + +
+ +
+
+
+

{{ item.label }}

+

{{ item.value }}

+

{{ item.meta }}

+
+
+
+
+ +
+
+
+

{{ t('admin.riskControl.preBlockAPIKeyLoad') }}

+

+ {{ t('admin.riskControl.preBlockAPIKeyLoadHint') }} +

+
+ + {{ preBlockAPIKeyLoadSummaryText }} + +
+ +
+
+
+
+
+
+ #{{ item.index + 1 }} + {{ item.masked || '-' }} + +
+

+ {{ t('admin.riskControl.preBlockAPIKeyTotals', { total: formatNumber(item.total), success: formatNumber(item.success), errors: formatNumber(item.errors) }) }} +

+
+
+
+

{{ t('admin.riskControl.preBlockKeyActiveShort') }}

+

{{ formatNumber(item.active) }}

+
+
+

{{ t('admin.riskControl.preBlockKeyTotalShort') }}

+

{{ formatNumber(item.total) }}

+
+
+

{{ t('admin.riskControl.preBlockKeyAvgShort') }}

+

{{ formatNumber(item.avg_latency_ms) }} ms

+
+
+

{{ t('admin.riskControl.preBlockKeyLastShort') }}

+

{{ formatNumber(item.last_latency_ms) }} ms

+
+
+
+
+
+
+
+
+

+ {{ t('admin.riskControl.preBlockAPIKeyLoadEmpty') }} +

+
+
+
+ +

{{ t('admin.riskControl.workerStatus') }}

@@ -1013,6 +1111,7 @@ import Pagination from '@/components/common/Pagination.vue' import ModelWhitelistSelector from '@/components/account/ModelWhitelistSelector.vue' import { adminAPI } from '@/api/admin' import type { + ContentModerationAPIKeyLoad, ContentModerationAPIKeyStatus, ContentModerationConfig, ContentModerationLog, @@ -1472,6 +1571,81 @@ const queueUsageStyle = computed(() => ({ width: queueUsagePercent.value, })) +const runtimeMode = computed(() => status.value?.mode ?? configForm.mode) + +const showPreBlockRuntimeCard = computed(() => runtimeMode.value === 'pre_block') + +const showWorkerRuntimeCard = computed(() => runtimeMode.value === 'observe') + +const preBlockMetricItems = computed(() => [ + { + key: 'active', + label: t('admin.riskControl.preBlockActive'), + value: formatNumber(status.value?.pre_block_active ?? 0), + meta: t('admin.riskControl.preBlockActiveHint'), + class: 'bg-sky-50 dark:bg-sky-900/10', + valueClass: 'text-sky-700 dark:text-sky-300', + }, + { + key: 'checked', + label: t('admin.riskControl.preBlockChecked'), + value: formatNumber(status.value?.pre_block_checked ?? 0), + meta: t('admin.riskControl.preBlockCheckedHint'), + class: 'bg-gray-50 dark:bg-dark-700/50', + valueClass: 'text-gray-900 dark:text-white', + }, + { + key: 'allowed', + label: t('admin.riskControl.preBlockAllowed'), + value: formatNumber(status.value?.pre_block_allowed ?? 0), + meta: t('admin.riskControl.preBlockAllowedHint'), + class: 'bg-emerald-50 dark:bg-emerald-900/10', + valueClass: 'text-emerald-700 dark:text-emerald-300', + }, + { + key: 'blocked', + label: t('admin.riskControl.preBlockBlocked'), + value: formatNumber(status.value?.pre_block_blocked ?? 0), + meta: t('admin.riskControl.preBlockBlockedHint'), + class: 'bg-rose-50 dark:bg-rose-900/10', + valueClass: 'text-rose-700 dark:text-rose-300', + }, + { + key: 'errors', + label: t('admin.riskControl.preBlockErrors'), + value: formatNumber(status.value?.pre_block_errors ?? 0), + meta: t('admin.riskControl.preBlockErrorsHint'), + class: 'bg-amber-50 dark:bg-amber-900/10', + valueClass: 'text-amber-700 dark:text-amber-300', + }, + { + key: 'latency', + label: t('admin.riskControl.preBlockAvgLatency'), + value: `${formatNumber(status.value?.pre_block_avg_latency_ms ?? 0)} ms`, + meta: t('admin.riskControl.preBlockAvgLatencyHint'), + class: 'bg-violet-50 dark:bg-violet-900/10', + valueClass: 'text-violet-700 dark:text-violet-300', + }, +]) + +const preBlockAPIKeyLoads = computed(() => ( + [...(status.value?.pre_block_api_key_loads ?? [])].sort((a, b) => a.index - b.index) +)) + +const preBlockAPIKeyMaxTotal = computed(() => Math.max(1, ...preBlockAPIKeyLoads.value.map((item) => item.total || 0))) + +const preBlockAPIKeyLoadSummaryText = computed(() => t('admin.riskControl.preBlockAPIKeyLoadSummary', { + active: formatNumber(status.value?.pre_block_api_key_active ?? 0), + available: formatNumber(status.value?.pre_block_api_key_available_count ?? 0), + total: formatNumber(status.value?.pre_block_api_key_total_calls ?? 0), + workerActive: formatNumber(status.value?.active_workers ?? 0), + workerTotal: formatNumber(status.value?.worker_count ?? configForm.worker_count), +})) + +function preBlockAPIKeyLoadWidth(total: number): string { + return `${Math.min(100, Math.max(0, (total / preBlockAPIKeyMaxTotal.value) * 100)).toFixed(1)}%` +} + const workerSlots = computed(() => { const total = Math.max(0, status.value?.worker_count ?? configForm.worker_count) const active = Math.max(0, status.value?.active_workers ?? 0) diff --git a/frontend/src/views/admin/__tests__/RiskControlView.spec.ts b/frontend/src/views/admin/__tests__/RiskControlView.spec.ts index 3c6aa0e9..5f1798a4 100644 --- a/frontend/src/views/admin/__tests__/RiskControlView.spec.ts +++ b/frontend/src/views/admin/__tests__/RiskControlView.spec.ts @@ -58,8 +58,12 @@ vi.mock('vue-i18n', async () => { return { ...actual, useI18n: () => ({ - t: (key: string, params?: Record) => - key.replace(/\{(\w+)\}/g, (_, token) => String(params?.[token] ?? `{${token}}`)), + t: (key: string, params?: Record) => { + if (key === 'admin.riskControl.preBlockAPIKeyLoadSummary') { + return `同步并发 ${params?.active} / 可用 Key ${params?.available},累计 ${params?.total} 次,worker:${params?.workerActive} / ${params?.workerTotal}` + } + return key.replace(/\{(\w+)\}/g, (_, token) => String(params?.[token] ?? `{${token}}`)) + }, }), } }) @@ -118,6 +122,16 @@ const runtimeStatus = () => ({ dropped: 0, processed: 0, errors: 0, + pre_block_active: 0, + pre_block_checked: 0, + pre_block_allowed: 0, + pre_block_blocked: 0, + pre_block_errors: 0, + pre_block_avg_latency_ms: 0, + pre_block_api_key_active: 0, + pre_block_api_key_available_count: 0, + pre_block_api_key_total_calls: 0, + pre_block_api_key_loads: [], api_key_statuses: [], flagged_hash_count: 0, last_cleanup_deleted_hit: 0, @@ -261,4 +275,133 @@ describe('admin RiskControlView', () => { })) expect(showError).not.toHaveBeenCalled() }) + + it('describes worker runtime as async audit and pre-block record processing', async () => { + getStatus.mockResolvedValue({ + ...runtimeStatus(), + mode: 'observe', + processed: 12, + queue_length: 2, + }) + + const wrapper = mount(RiskControlView, { + global: { + stubs: { + AppLayout: AppLayoutStub, + BaseDialog: BaseDialogStub, + Icon: true, + Select: true, + Toggle: true, + Pagination: true, + ModelWhitelistSelector: ModelWhitelistSelectorStub, + }, + }, + }) + + await flushPromises() + + expect(wrapper.text()).toContain('admin.riskControl.workerStatusHint') + expect(wrapper.text()).not.toContain('admin.riskControl.preBlockSyncStatus') + expect(wrapper.text()).toContain('admin.riskControl.records') + expect(wrapper.text()).toContain('12') + expect(wrapper.text()).toContain('2 / 32,768') + }) + + it('shows pre-block synchronous moderation metrics separately from worker queue', async () => { + getStatus.mockResolvedValue({ + ...runtimeStatus(), + pre_block_active: 2, + pre_block_checked: 128, + pre_block_allowed: 120, + pre_block_blocked: 8, + pre_block_errors: 1, + pre_block_avg_latency_ms: 86, + pre_block_api_key_active: 2, + pre_block_api_key_available_count: 2, + pre_block_api_key_total_calls: 128, + active_workers: 3, + worker_count: 7, + pre_block_api_key_loads: [ + { + index: 0, + key_hash: 'hash-one', + masked: 'sk-...one', + status: 'ok', + active: 1, + total: 72, + success: 70, + errors: 2, + avg_latency_ms: 84, + last_latency_ms: 80, + last_http_status: 200, + }, + { + index: 1, + key_hash: 'hash-two', + masked: 'sk-...two', + status: 'ok', + active: 1, + total: 56, + success: 56, + errors: 0, + avg_latency_ms: 90, + last_latency_ms: 92, + last_http_status: 200, + }, + ], + }) + + const wrapper = mount(RiskControlView, { + global: { + stubs: { + AppLayout: AppLayoutStub, + BaseDialog: BaseDialogStub, + Icon: true, + Select: true, + Toggle: true, + Pagination: true, + ModelWhitelistSelector: ModelWhitelistSelectorStub, + }, + }, + }) + + await flushPromises() + + expect(wrapper.text()).toContain('admin.riskControl.preBlockSyncStatus') + expect(wrapper.text()).toContain('admin.riskControl.preBlockSyncHint') + expect(wrapper.text()).not.toContain('admin.riskControl.workerStatus') + expect(wrapper.text()).toContain('admin.riskControl.records') + expect(wrapper.text()).toContain('128') + expect(wrapper.text()).toContain('120') + expect(wrapper.text()).toContain('8') + expect(wrapper.text()).toContain('86 ms') + expect(wrapper.text()).toContain('admin.riskControl.preBlockAPIKeyLoad') + expect(wrapper.text()).toContain('sk-...one') + expect(wrapper.text()).toContain('sk-...two') + expect(wrapper.text()).toContain('72') + expect(wrapper.text()).toContain('56') + expect(wrapper.text()).toContain('同步并发 2 / 可用 Key 2,累计 128 次,worker:3 / 7') + + const runtimeCards = wrapper.get('[data-test="pre-block-runtime-cards"]') + const syncCard = wrapper.get('[data-test="pre-block-sync-card"]') + const apiKeyLoadCard = wrapper.get('[data-test="pre-block-api-key-load-card"]') + + expect(runtimeCards.classes()).toEqual(expect.arrayContaining([ + 'grid', + 'grid-cols-1', + 'xl:grid-cols-[minmax(0,520px)_minmax(0,1fr)]', + ])) + expect(syncCard.element.parentElement).toBe(runtimeCards.element) + expect(apiKeyLoadCard.element.parentElement).toBe(runtimeCards.element) + expect(syncCard.classes()).toContain('card') + expect(apiKeyLoadCard.classes()).toContain('card') + expect(syncCard.get('h2').text()).toBe('admin.riskControl.preBlockSyncStatus') + expect(syncCard.text()).toContain('admin.riskControl.preBlockSyncHint') + expect(apiKeyLoadCard.get('h2').text()).toBe('admin.riskControl.preBlockAPIKeyLoad') + expect(apiKeyLoadCard.text()).toContain('admin.riskControl.preBlockAPIKeyLoadHint') + expect(wrapper.get('[data-test="pre-block-api-key-load-list"]').classes()).toEqual(expect.arrayContaining([ + 'max-h-[280px]', + 'overflow-y-auto', + ])) + }) }) From 6aec505016de92c4860316edb37b66e14f6ff5d2 Mon Sep 17 00:00:00 2001 From: fofoj <163302188+fofoj@users.noreply.github.com> Date: Thu, 28 May 2026 20:05:38 +0800 Subject: [PATCH 25/79] fix(oauth): don't overwrite credentials JSONB in 401 handler The 401 handler in RateLimitService.HandleUpstreamError set account.Credentials["expires_at"] = time.Now() and then persisted the full credentials map via persistAccountCredentials, which routes through accountRepository.UpdateCredentials -> ent SetCredentials and replaces the entire JSONB column. The account passed to the handler is the request-start snapshot taken by the gateway at SelectAccount time. When another worker has just rotated refresh_token via oauth_refresh_api.RefreshIfNeeded, the snapshot still holds the old refresh_token; writing the full snapshot back rolls refresh_token in the DB back to the stale value. The next refresh cycle then calls the upstream with the stale token, receives invalid_grant, and tryRecoverFromRefreshRace re-reads the DB only to find currentRT == usedRT (because the 401 handler just poisoned the DB), returns false, and the account is incorrectly disabled. Drop the credentials write. InvalidateToken + SetTempUnschedulable is sufficient: the account is held out of scheduling during the cooldown, and after the cooldown the next request goes through token_provider's NeedsRefresh check, which routes through the locked, DB-re-reading RefreshIfNeeded path. The "force background refresh by setting expires_at = now" semantic is intentionally dropped. token_refresh_service will naturally pick the account up when the real expires_at enters the refresh window, and if the real expires_at has already passed by the time the account becomes schedulable again, token_provider's NeedsRefresh returns true and RefreshIfNeeded fires synchronously on the next request. --- backend/internal/service/ratelimit_service.go | 20 +++++++++---------- 1 file changed, 9 insertions(+), 11 deletions(-) diff --git a/backend/internal/service/ratelimit_service.go b/backend/internal/service/ratelimit_service.go index d12824ec..ecbd86d1 100644 --- a/backend/internal/service/ratelimit_service.go +++ b/backend/internal/service/ratelimit_service.go @@ -248,17 +248,15 @@ func (s *RateLimitService) HandleUpstreamError(ctx context.Context, account *Acc shouldDisable = true break } - // 2. 设置 expires_at 为当前时间,强制下次请求刷新 token - if account.Credentials == nil { - account.Credentials = make(map[string]any) - } - account.Credentials["expires_at"] = time.Now().Format(time.RFC3339) - if err := persistAccountCredentials(ctx, s.accountRepo, account, account.Credentials); err != nil { - slog.Warn("oauth_401_force_refresh_update_failed", "account_id", account.ID, "error", err) - } else { - slog.Info("oauth_401_force_refresh_set", "account_id", account.ID, "platform", account.Platform) - } - // 3. 临时不可调度,替代 SetError(保持 status=active 让刷新服务能拾取) + // 2. 临时不可调度,替代 SetError(保持 status=active 让刷新服务能拾取) + // 注意:此处不再写回 account.Credentials/expires_at。 + // 原实现使用请求开始时的 account 快照整列覆盖 credentials JSONB(见 + // persistAccountCredentials → accountRepository.UpdateCredentials → SetCredentials), + // 在另一个 worker 刚刷新完 refresh_token 的窄窗口内会把新 refresh_token 回滚为旧值, + // 导致下一周期用旧 refresh_token 调上游拿到 invalid_grant 后, + // tryRecoverFromRefreshRace 重读 DB 发现 currentRT == usedRT 也救不回来,账号被错误 disable。 + // 这里仅依赖 InvalidateToken + SetTempUnschedulable 让账号在冷却期内不被调度, + // 冷却结束后由 token_provider 的 NeedsRefresh / token_refresh_service 走带分布式锁的正路刷新。 msg := "Authentication failed (401): invalid or expired credentials" if upstreamMsg != "" { msg = "OAuth 401: " + upstreamMsg From be3613593b1b984389a8de031eed3e2ae04300c8 Mon Sep 17 00:00:00 2001 From: fofoj <163302188+fofoj@users.noreply.github.com> Date: Thu, 28 May 2026 20:32:16 +0800 Subject: [PATCH 26/79] test(oauth): update OAuth 401 tests to match new no-write behavior Two tests in ratelimit_service_401_test.go were encoding the bug behavior itself: - OAuth401InvalidatorError asserted updateCredentialsCalls == 1 - OAuth401UsesCredentialsUpdater asserted updateCredentialsCalls == 1 and lastCredentials["expires_at"] non-empty Both assertions exercised the exact write-back this PR removes. Update them to reflect the new contract and guard against regression: - OAuth401InvalidatorError: assert updateCredentialsCalls == 0 - OAuth401UsesCredentialsUpdater is renamed to OAuth401DoesNotOverwriteCredentials with reversed assertions, so it now serves as a regression test ensuring the 401 handler never writes credentials back from the request-start snapshot. --- .../service/ratelimit_service_401_test.go | 19 ++++++++++++++----- 1 file changed, 14 insertions(+), 5 deletions(-) diff --git a/backend/internal/service/ratelimit_service_401_test.go b/backend/internal/service/ratelimit_service_401_test.go index a964775e..873aaf33 100644 --- a/backend/internal/service/ratelimit_service_401_test.go +++ b/backend/internal/service/ratelimit_service_401_test.go @@ -129,7 +129,10 @@ func TestRateLimitService_HandleUpstreamError_OAuth401SetsTempUnschedulable(t *t } // TestRateLimitService_HandleUpstreamError_OAuth401InvalidatorError -// OpenAI OAuth 401 缓存失效出错时仍走 temp_unschedulable +// OpenAI OAuth 401 缓存失效出错时仍走 temp_unschedulable。 +// 注意:401 handler 不再回写 credentials(避免请求开始时的快照整列覆盖 DB +// 把另一个 worker 刚刷新出来的新 refresh_token 回滚为旧值), +// 因此 updateCredentialsCalls 应当为 0。 func TestRateLimitService_HandleUpstreamError_OAuth401InvalidatorError(t *testing.T) { repo := &rateLimitAccountRepoStub{} invalidator := &tokenCacheInvalidatorRecorder{err: errors.New("boom")} @@ -149,7 +152,7 @@ func TestRateLimitService_HandleUpstreamError_OAuth401InvalidatorError(t *testin require.True(t, shouldDisable) require.Equal(t, 0, repo.setErrorCalls) require.Equal(t, 1, repo.tempCalls) - require.Equal(t, 1, repo.updateCredentialsCalls) + require.Equal(t, 0, repo.updateCredentialsCalls) require.Len(t, invalidator.accounts, 1) } @@ -171,7 +174,12 @@ func TestRateLimitService_HandleUpstreamError_NonOAuth401(t *testing.T) { require.Empty(t, invalidator.accounts) } -func TestRateLimitService_HandleUpstreamError_OAuth401UsesCredentialsUpdater(t *testing.T) { +// TestRateLimitService_HandleUpstreamError_OAuth401DoesNotOverwriteCredentials +// 回归测试:确保 401 handler 不再使用请求开始时的 account 快照写回 credentials。 +// 原实现会通过 persistAccountCredentials → UpdateCredentials → SetCredentials +// 整列覆盖 credentials JSONB,在另一个 worker 刚刷新完 refresh_token 的窄窗口内 +// 会把新 refresh_token 回滚为快照中的旧值,导致下一周期拿 invalid_grant 被错误 disable。 +func TestRateLimitService_HandleUpstreamError_OAuth401DoesNotOverwriteCredentials(t *testing.T) { repo := &rateLimitAccountRepoStub{} service := NewRateLimitService(repo, nil, &config.Config{}, nil, nil) account := &Account{ @@ -187,8 +195,9 @@ func TestRateLimitService_HandleUpstreamError_OAuth401UsesCredentialsUpdater(t * shouldDisable := service.HandleUpstreamError(context.Background(), account, 401, http.Header{}, []byte("unauthorized")) require.True(t, shouldDisable) - require.Equal(t, 1, repo.updateCredentialsCalls) - require.NotEmpty(t, repo.lastCredentials["expires_at"]) + require.Equal(t, 0, repo.updateCredentialsCalls, "401 handler must not write credentials back from the request-start snapshot") + require.Equal(t, 1, repo.tempCalls, "401 handler should still set temp-unschedulable cooldown") + require.Nil(t, repo.lastCredentials, "no credentials should have been persisted") } // 缺少 refresh_token 的 OAuth 账号 401 应直接 SetError 永久禁用, From 2bd3125d0fe23515fd42ee3c6651bbe38e47593c Mon Sep 17 00:00:00 2001 From: Wey Gu Date: Thu, 28 May 2026 22:08:02 +0800 Subject: [PATCH 27/79] Preserve usage request context --- backend/internal/handler/gateway_handler.go | 7 +-- .../gateway_handler_chat_completions.go | 2 +- .../handler/gateway_handler_responses.go | 2 +- .../internal/handler/gemini_v1beta_handler.go | 2 +- .../handler/openai_chat_completions.go | 2 +- backend/internal/handler/openai_embeddings.go | 2 +- .../handler/openai_gateway_handler.go | 44 +++++++++++++--- .../openai_gateway_usage_context_test.go | 41 +++++++++++++++ backend/internal/handler/openai_images.go | 2 +- .../handler/usage_record_submit_task_test.go | 24 ++++----- .../server/middleware/client_request_id.go | 6 ++- .../middleware/client_request_id_test.go | 50 +++++++++++++++++++ 12 files changed, 154 insertions(+), 30 deletions(-) create mode 100644 backend/internal/handler/openai_gateway_usage_context_test.go create mode 100644 backend/internal/server/middleware/client_request_id_test.go diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index 4695a791..a6749191 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -510,7 +510,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) { // 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。 quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) - h.submitUsageRecordTask(func(ctx context.Context) { + h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) { if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{ Result: result, ParsedRequest: parsedReq, @@ -905,7 +905,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) { // 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。 quotaPlatform := service.QuotaPlatform(c.Request.Context(), currentAPIKey) - h.submitUsageRecordTask(func(ctx context.Context) { + h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) { if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{ Result: result, ParsedRequest: parsedReq, @@ -2056,10 +2056,11 @@ func (h *GatewayHandler) maybeLogCompatibilityFallbackMetrics(reqLog *zap.Logger ) } -func (h *GatewayHandler) submitUsageRecordTask(task service.UsageRecordTask) { +func (h *GatewayHandler) submitUsageRecordTask(parent context.Context, task service.UsageRecordTask) { if task == nil { return } + task = wrapUsageRecordTaskContext(parent, task) if h.usageRecordWorkerPool != nil { h.usageRecordWorkerPool.Submit(task) return diff --git a/backend/internal/handler/gateway_handler_chat_completions.go b/backend/internal/handler/gateway_handler_chat_completions.go index acbdc261..daf6e6ea 100644 --- a/backend/internal/handler/gateway_handler_chat_completions.go +++ b/backend/internal/handler/gateway_handler_chat_completions.go @@ -292,7 +292,7 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) { upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) - h.submitUsageRecordTask(func(ctx context.Context) { + h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) { if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{ Result: result, QuotaPlatform: quotaPlatform, diff --git a/backend/internal/handler/gateway_handler_responses.go b/backend/internal/handler/gateway_handler_responses.go index 6a083f31..f57b9989 100644 --- a/backend/internal/handler/gateway_handler_responses.go +++ b/backend/internal/handler/gateway_handler_responses.go @@ -267,7 +267,7 @@ func (h *GatewayHandler) Responses(c *gin.Context) { upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) - h.submitUsageRecordTask(func(ctx context.Context) { + h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) { if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{ Result: result, QuotaPlatform: quotaPlatform, diff --git a/backend/internal/handler/gemini_v1beta_handler.go b/backend/internal/handler/gemini_v1beta_handler.go index 27ea4404..0b33ca3e 100644 --- a/backend/internal/handler/gemini_v1beta_handler.go +++ b/backend/internal/handler/gemini_v1beta_handler.go @@ -528,7 +528,7 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) { inboundEndpoint := GetInboundEndpoint(c) upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) - h.submitUsageRecordTask(func(ctx context.Context) { + h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) { if err := h.gatewayService.RecordUsageWithLongContext(ctx, &service.RecordUsageLongContextInput{ Result: result, QuotaPlatform: quotaPlatform, diff --git a/backend/internal/handler/openai_chat_completions.go b/backend/internal/handler/openai_chat_completions.go index 17f0d47e..9f63ef1f 100644 --- a/backend/internal/handler/openai_chat_completions.go +++ b/backend/internal/handler/openai_chat_completions.go @@ -273,7 +273,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { inboundEndpoint := GetInboundEndpoint(c) upstreamEndpoint := resolveRawCCUpstreamEndpoint(c, account) - h.submitOpenAIUsageRecordTask(result, func(ctx context.Context) { + h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) { if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{ Result: result, APIKey: apiKey, diff --git a/backend/internal/handler/openai_embeddings.go b/backend/internal/handler/openai_embeddings.go index bbb67044..b64ac41d 100644 --- a/backend/internal/handler/openai_embeddings.go +++ b/backend/internal/handler/openai_embeddings.go @@ -220,7 +220,7 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) { inboundEndpoint := GetInboundEndpoint(c) upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) - h.submitOpenAIUsageRecordTask(result, func(ctx context.Context) { + h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) { if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{ Result: result, APIKey: apiKey, diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index a51eee86..86503f30 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -12,6 +12,7 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" pkghttputil "github.com/Wei-Shaw/sub2api/internal/pkg/httputil" "github.com/Wei-Shaw/sub2api/internal/pkg/ip" "github.com/Wei-Shaw/sub2api/internal/pkg/logger" @@ -46,6 +47,31 @@ func resolveOpenAIMessagesDispatchMappedModel(apiKey *service.APIKey, requestedM return strings.TrimSpace(apiKey.Group.ResolveMessagesDispatchModel(requestedModel)) } +func usageRecordContext(parent context.Context, base context.Context) context.Context { + if base == nil { + base = context.Background() + } + if parent == nil { + return base + } + if clientRequestID, _ := parent.Value(ctxkey.ClientRequestID).(string); strings.TrimSpace(clientRequestID) != "" { + base = context.WithValue(base, ctxkey.ClientRequestID, strings.TrimSpace(clientRequestID)) + } + if requestID, _ := parent.Value(ctxkey.RequestID).(string); strings.TrimSpace(requestID) != "" { + base = context.WithValue(base, ctxkey.RequestID, strings.TrimSpace(requestID)) + } + return base +} + +func wrapUsageRecordTaskContext(parent context.Context, task service.UsageRecordTask) service.UsageRecordTask { + if task == nil { + return nil + } + return func(ctx context.Context) { + task(usageRecordContext(parent, ctx)) + } +} + // NewOpenAIGatewayHandler creates a new OpenAIGatewayHandler func NewOpenAIGatewayHandler( gatewayService *service.OpenAIGatewayService, @@ -437,7 +463,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) // 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。 - h.submitOpenAIUsageRecordTask(result, func(ctx context.Context) { + h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) { if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{ Result: result, APIKey: apiKey, @@ -821,7 +847,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { inboundEndpoint := GetInboundEndpoint(c) upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) - h.submitOpenAIUsageRecordTask(result, func(ctx context.Context) { + h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) { if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{ Result: result, APIKey: apiKey, @@ -1424,7 +1450,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, true, result.FirstTokenMs) inboundEndpoint := GetInboundEndpoint(c) upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) - h.submitOpenAIUsageRecordTask(result, func(taskCtx context.Context) { + h.submitOpenAIUsageRecordTask(ctx, result, func(taskCtx context.Context) { if err := h.gatewayService.RecordUsage(taskCtx, &service.OpenAIRecordUsageInput{ Result: result, APIKey: apiKey, @@ -1609,10 +1635,11 @@ func getContextInt64(c *gin.Context, key string) (int64, bool) { } } -func (h *OpenAIGatewayHandler) submitUsageRecordTask(task service.UsageRecordTask) { +func (h *OpenAIGatewayHandler) submitUsageRecordTask(parent context.Context, task service.UsageRecordTask) { if task == nil { return } + task = wrapUsageRecordTaskContext(parent, task) if h.usageRecordWorkerPool != nil { h.usageRecordWorkerPool.Submit(task) return @@ -1631,18 +1658,19 @@ func (h *OpenAIGatewayHandler) submitUsageRecordTask(task service.UsageRecordTas task(ctx) } -func (h *OpenAIGatewayHandler) submitOpenAIUsageRecordTask(result *service.OpenAIForwardResult, task service.UsageRecordTask) { +func (h *OpenAIGatewayHandler) submitOpenAIUsageRecordTask(parent context.Context, result *service.OpenAIForwardResult, task service.UsageRecordTask) { if result != nil && result.ImageCount > 0 { - h.submitMandatoryUsageRecordTask(task) + h.submitMandatoryUsageRecordTask(parent, task) return } - h.submitUsageRecordTask(task) + h.submitUsageRecordTask(parent, task) } -func (h *OpenAIGatewayHandler) submitMandatoryUsageRecordTask(task service.UsageRecordTask) { +func (h *OpenAIGatewayHandler) submitMandatoryUsageRecordTask(parent context.Context, task service.UsageRecordTask) { if task == nil { return } + task = wrapUsageRecordTaskContext(parent, task) if h.usageRecordWorkerPool != nil { if mode := h.usageRecordWorkerPool.Submit(task); mode != service.UsageRecordSubmitModeDropped { return diff --git a/backend/internal/handler/openai_gateway_usage_context_test.go b/backend/internal/handler/openai_gateway_usage_context_test.go new file mode 100644 index 00000000..7091c9c0 --- /dev/null +++ b/backend/internal/handler/openai_gateway_usage_context_test.go @@ -0,0 +1,41 @@ +package handler + +import ( + "context" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" + "github.com/stretchr/testify/require" +) + +func TestSubmitUsageRecordTaskCopiesRequestContext(t *testing.T) { + parent := context.WithValue(context.Background(), ctxkey.ClientRequestID, "client-request-123") + parent = context.WithValue(parent, ctxkey.RequestID, "request-456") + + var gotClientRequestID string + var gotRequestID string + h := &GatewayHandler{} + h.submitUsageRecordTask(parent, func(ctx context.Context) { + gotClientRequestID, _ = ctx.Value(ctxkey.ClientRequestID).(string) + gotRequestID, _ = ctx.Value(ctxkey.RequestID).(string) + }) + + require.Equal(t, "client-request-123", gotClientRequestID) + require.Equal(t, "request-456", gotRequestID) +} + +func TestOpenAISubmitUsageRecordTaskCopiesRequestContext(t *testing.T) { + parent := context.WithValue(context.Background(), ctxkey.ClientRequestID, "openai-client-request-123") + parent = context.WithValue(parent, ctxkey.RequestID, "openai-request-456") + + var gotClientRequestID string + var gotRequestID string + h := &OpenAIGatewayHandler{} + h.submitUsageRecordTask(parent, func(ctx context.Context) { + gotClientRequestID, _ = ctx.Value(ctxkey.ClientRequestID).(string) + gotRequestID, _ = ctx.Value(ctxkey.RequestID).(string) + }) + + require.Equal(t, "openai-client-request-123", gotClientRequestID) + require.Equal(t, "openai-request-456", gotRequestID) +} diff --git a/backend/internal/handler/openai_images.go b/backend/internal/handler/openai_images.go index bbb08014..36339d4b 100644 --- a/backend/internal/handler/openai_images.go +++ b/backend/internal/handler/openai_images.go @@ -311,7 +311,7 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) { if result != nil { upstreamModel = result.UpstreamModel } - h.submitMandatoryUsageRecordTask(func(ctx context.Context) { + h.submitMandatoryUsageRecordTask(c.Request.Context(), func(ctx context.Context) { if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{ Result: result, APIKey: apiKey, diff --git a/backend/internal/handler/usage_record_submit_task_test.go b/backend/internal/handler/usage_record_submit_task_test.go index e4c2837a..ebe5c3df 100644 --- a/backend/internal/handler/usage_record_submit_task_test.go +++ b/backend/internal/handler/usage_record_submit_task_test.go @@ -29,7 +29,7 @@ func TestGatewayHandlerSubmitUsageRecordTask_WithPool(t *testing.T) { h := &GatewayHandler{usageRecordWorkerPool: pool} done := make(chan struct{}) - h.submitUsageRecordTask(func(ctx context.Context) { + h.submitUsageRecordTask(context.Background(), func(ctx context.Context) { close(done) }) @@ -44,7 +44,7 @@ func TestGatewayHandlerSubmitUsageRecordTask_WithoutPoolSyncFallback(t *testing. h := &GatewayHandler{} var called atomic.Bool - h.submitUsageRecordTask(func(ctx context.Context) { + h.submitUsageRecordTask(context.Background(), func(ctx context.Context) { if _, ok := ctx.Deadline(); !ok { t.Fatal("expected deadline in fallback context") } @@ -57,7 +57,7 @@ func TestGatewayHandlerSubmitUsageRecordTask_WithoutPoolSyncFallback(t *testing. func TestGatewayHandlerSubmitUsageRecordTask_NilTask(t *testing.T) { h := &GatewayHandler{} require.NotPanics(t, func() { - h.submitUsageRecordTask(nil) + h.submitUsageRecordTask(context.Background(), nil) }) } @@ -66,12 +66,12 @@ func TestGatewayHandlerSubmitUsageRecordTask_WithoutPool_TaskPanicRecovered(t *t var called atomic.Bool require.NotPanics(t, func() { - h.submitUsageRecordTask(func(ctx context.Context) { + h.submitUsageRecordTask(context.Background(), func(ctx context.Context) { panic("usage task panic") }) }) - h.submitUsageRecordTask(func(ctx context.Context) { + h.submitUsageRecordTask(context.Background(), func(ctx context.Context) { called.Store(true) }) require.True(t, called.Load(), "panic 后后续任务应仍可执行") @@ -82,7 +82,7 @@ func TestOpenAIGatewayHandlerSubmitUsageRecordTask_WithPool(t *testing.T) { h := &OpenAIGatewayHandler{usageRecordWorkerPool: pool} done := make(chan struct{}) - h.submitUsageRecordTask(func(ctx context.Context) { + h.submitUsageRecordTask(context.Background(), func(ctx context.Context) { close(done) }) @@ -97,7 +97,7 @@ func TestOpenAIGatewayHandlerSubmitUsageRecordTask_WithoutPoolSyncFallback(t *te h := &OpenAIGatewayHandler{} var called atomic.Bool - h.submitUsageRecordTask(func(ctx context.Context) { + h.submitUsageRecordTask(context.Background(), func(ctx context.Context) { if _, ok := ctx.Deadline(); !ok { t.Fatal("expected deadline in fallback context") } @@ -110,7 +110,7 @@ func TestOpenAIGatewayHandlerSubmitUsageRecordTask_WithoutPoolSyncFallback(t *te func TestOpenAIGatewayHandlerSubmitUsageRecordTask_NilTask(t *testing.T) { h := &OpenAIGatewayHandler{} require.NotPanics(t, func() { - h.submitUsageRecordTask(nil) + h.submitUsageRecordTask(context.Background(), nil) }) } @@ -119,12 +119,12 @@ func TestOpenAIGatewayHandlerSubmitUsageRecordTask_WithoutPool_TaskPanicRecovere var called atomic.Bool require.NotPanics(t, func() { - h.submitUsageRecordTask(func(ctx context.Context) { + h.submitUsageRecordTask(context.Background(), func(ctx context.Context) { panic("usage task panic") }) }) - h.submitUsageRecordTask(func(ctx context.Context) { + h.submitUsageRecordTask(context.Background(), func(ctx context.Context) { called.Store(true) }) require.True(t, called.Load(), "panic 后后续任务应仍可执行") @@ -152,7 +152,7 @@ func TestOpenAIGatewayHandlerSubmitMandatoryUsageRecordTask_DroppedTaskSyncFallb pool.Submit(func(ctx context.Context) {}) var called atomic.Bool - h.submitMandatoryUsageRecordTask(func(ctx context.Context) { + h.submitMandatoryUsageRecordTask(context.Background(), func(ctx context.Context) { called.Store(true) }) close(release) @@ -182,7 +182,7 @@ func TestOpenAIGatewayHandlerSubmitOpenAIUsageRecordTask_ImageResultUsesMandator pool.Submit(func(ctx context.Context) {}) var called atomic.Bool - h.submitOpenAIUsageRecordTask(&service.OpenAIForwardResult{ImageCount: 1}, func(ctx context.Context) { + h.submitOpenAIUsageRecordTask(context.Background(), &service.OpenAIForwardResult{ImageCount: 1}, func(ctx context.Context) { called.Store(true) }) close(release) diff --git a/backend/internal/server/middleware/client_request_id.go b/backend/internal/server/middleware/client_request_id.go index 6838d6af..5f886646 100644 --- a/backend/internal/server/middleware/client_request_id.go +++ b/backend/internal/server/middleware/client_request_id.go @@ -11,6 +11,8 @@ import ( "go.uber.org/zap" ) +const clientRequestIDHeader = "X-Client-Request-ID" + // ClientRequestID ensures every request has a unique client_request_id in request.Context(). // // This is used by the Ops monitoring module for end-to-end request correlation. @@ -21,12 +23,14 @@ func ClientRequestID() gin.HandlerFunc { return } - if v := c.Request.Context().Value(ctxkey.ClientRequestID); v != nil { + if v, _ := c.Request.Context().Value(ctxkey.ClientRequestID).(string); strings.TrimSpace(v) != "" { + c.Header(clientRequestIDHeader, strings.TrimSpace(v)) c.Next() return } id := uuid.New().String() + c.Header(clientRequestIDHeader, id) ctx := context.WithValue(c.Request.Context(), ctxkey.ClientRequestID, id) requestLogger := logger.FromContext(ctx).With(zap.String("client_request_id", strings.TrimSpace(id))) ctx = logger.IntoContext(ctx, requestLogger) diff --git a/backend/internal/server/middleware/client_request_id_test.go b/backend/internal/server/middleware/client_request_id_test.go new file mode 100644 index 00000000..394c1612 --- /dev/null +++ b/backend/internal/server/middleware/client_request_id_test.go @@ -0,0 +1,50 @@ +package middleware + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestClientRequestIDGeneratesAndExposesID(t *testing.T) { + gin.SetMode(gin.TestMode) + router := gin.New() + router.Use(ClientRequestID()) + router.GET("/", func(c *gin.Context) { + value, _ := c.Request.Context().Value(ctxkey.ClientRequestID).(string) + c.String(http.StatusOK, value) + }) + + req := httptest.NewRequest(http.MethodGet, "/", nil) + w := httptest.NewRecorder() + + router.ServeHTTP(w, req) + + require.Equal(t, http.StatusOK, w.Code) + require.NotEmpty(t, w.Body.String()) + require.Equal(t, w.Body.String(), w.Header().Get(clientRequestIDHeader)) +} + +func TestClientRequestIDPreservesExistingContextID(t *testing.T) { + gin.SetMode(gin.TestMode) + router := gin.New() + router.Use(ClientRequestID()) + router.GET("/", func(c *gin.Context) { + value, _ := c.Request.Context().Value(ctxkey.ClientRequestID).(string) + c.String(http.StatusOK, value) + }) + + w := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/", nil) + req = req.WithContext(context.WithValue(req.Context(), ctxkey.ClientRequestID, "existing-client-request-id")) + router.ServeHTTP(w, req) + + require.Equal(t, http.StatusOK, w.Code) + require.Equal(t, "existing-client-request-id", w.Body.String()) + require.Equal(t, "existing-client-request-id", w.Header().Get(clientRequestIDHeader)) +} From ed1b57c5975578f62a88e95c3eaad3fa21e0efc0 Mon Sep 17 00:00:00 2001 From: shaw Date: Fri, 29 May 2026 08:58:10 +0800 Subject: [PATCH 28/79] fix(openai): gate routing by endpoint capability --- .../handler/openai_chat_completions.go | 3 +- backend/internal/handler/openai_embeddings.go | 10 +- .../handler/openai_gateway_handler.go | 9 +- backend/internal/service/account.go | 83 ++++++++ .../service/openai_account_scheduler.go | 55 +++-- .../service/openai_account_scheduler_test.go | 193 ++++++++++++++++++ .../service/openai_gateway_service.go | 62 +++--- .../internal/service/openai_images_test.go | 73 +++++++ .../service/openai_ws_account_sticky_test.go | 46 +++++ .../internal/service/openai_ws_forwarder.go | 39 +++- .../components/account/CreateAccountModal.vue | 69 ++++++- .../components/account/EditAccountModal.vue | 95 ++++++++- .../__tests__/EditAccountModal.spec.ts | 57 ++++++ frontend/src/i18n/locales/en.ts | 5 + frontend/src/i18n/locales/zh.ts | 5 + frontend/src/types/index.ts | 1 + 16 files changed, 740 insertions(+), 65 deletions(-) diff --git a/backend/internal/handler/openai_chat_completions.go b/backend/internal/handler/openai_chat_completions.go index 17f0d47e..9805bf8a 100644 --- a/backend/internal/handler/openai_chat_completions.go +++ b/backend/internal/handler/openai_chat_completions.go @@ -127,7 +127,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { for { reqLog.Debug("openai_chat_completions.account_selecting", zap.Int("excluded_account_count", len(failedAccountIDs))) - selection, scheduleDecision, err := h.gatewayService.SelectAccountWithScheduler( + selection, scheduleDecision, err := h.gatewayService.SelectAccountWithSchedulerForCapability( c.Request.Context(), apiKey.GroupID, "", @@ -135,6 +135,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { reqModel, failedAccountIDs, service.OpenAIUpstreamTransportAny, + service.OpenAIEndpointCapabilityChatCompletions, false, ) if err != nil { diff --git a/backend/internal/handler/openai_embeddings.go b/backend/internal/handler/openai_embeddings.go index bbb67044..81713f7f 100644 --- a/backend/internal/handler/openai_embeddings.go +++ b/backend/internal/handler/openai_embeddings.go @@ -107,7 +107,7 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) { routingStart := time.Now() for { - selection, _, err := h.gatewayService.SelectAccountWithScheduler( + selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability( c.Request.Context(), apiKey.GroupID, "", @@ -115,6 +115,7 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) { reqModel, failedAccountIDs, service.OpenAIUpstreamTransportHTTPSSE, + service.OpenAIEndpointCapabilityEmbeddings, false, ) if err != nil { @@ -140,13 +141,6 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) { return } account := selection.Account - if account.Type != service.AccountTypeAPIKey { - if selection.ReleaseFunc != nil { - selection.ReleaseFunc() - } - failedAccountIDs[account.ID] = struct{}{} - continue - } setOpsSelectedAccount(c, account.ID, account.Platform) accountReleaseFunc, accountAcquired := h.acquireResponsesAccountSlot(c, apiKey.GroupID, "", selection, false, &streamStarted, reqLog) diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index a51eee86..1d661748 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -266,7 +266,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { for { // Select account supporting the requested model reqLog.Debug("openai.account_selecting", zap.Int("excluded_account_count", len(failedAccountIDs))) - selection, scheduleDecision, err := h.gatewayService.SelectAccountWithScheduler( + selection, scheduleDecision, err := h.gatewayService.SelectAccountWithSchedulerForCapability( c.Request.Context(), apiKey.GroupID, previousResponseID, @@ -274,6 +274,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { reqModel, failedAccountIDs, service.OpenAIUpstreamTransportAny, + service.OpenAIEndpointCapabilityChatCompletions, requireCompact, ) if err != nil { @@ -675,7 +676,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { currentRoutingModel = effectiveMappedModel } reqLog.Debug("openai_messages.account_selecting", zap.Int("excluded_account_count", len(failedAccountIDs))) - selection, scheduleDecision, err := h.gatewayService.SelectAccountWithScheduler( + selection, scheduleDecision, err := h.gatewayService.SelectAccountWithSchedulerForCapability( c.Request.Context(), apiKey.GroupID, "", // no previous_response_id @@ -683,6 +684,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { currentRoutingModel, failedAccountIDs, service.OpenAIUpstreamTransportAny, + service.OpenAIEndpointCapabilityChatCompletions, false, ) if err != nil { @@ -1273,7 +1275,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { for { reqLog.Debug("openai.websocket_account_selecting", zap.Int("excluded_account_count", len(failedAccountIDs))) - selection, scheduleDecision, err := h.gatewayService.SelectAccountWithScheduler( + selection, scheduleDecision, err := h.gatewayService.SelectAccountWithSchedulerForCapability( ctx, apiKey.GroupID, previousResponseID, @@ -1281,6 +1283,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { reqModel, failedAccountIDs, service.OpenAIUpstreamTransportResponsesWebsocketV2, + service.OpenAIEndpointCapabilityChatCompletions, false, ) if err != nil { diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go index d488aa75..e3ca9c5d 100644 --- a/backend/internal/service/account.go +++ b/backend/internal/service/account.go @@ -66,6 +66,15 @@ type Account struct { modelMappingCacheRawSig uint64 } +type OpenAIEndpointCapability string + +const ( + OpenAIEndpointCapabilityChatCompletions OpenAIEndpointCapability = "chat_completions" + OpenAIEndpointCapabilityEmbeddings OpenAIEndpointCapability = "embeddings" +) + +const openAIEndpointCapabilitiesCredentialKey = "openai_capabilities" + type TempUnschedulableRule struct { ErrorCode int `json:"error_code"` Keywords []string `json:"keywords"` @@ -1122,6 +1131,80 @@ func (a *Account) GetOpenAISessionID() string { return strings.TrimSpace(a.GetExtraString("openai_session_id")) } +func (a *Account) SupportsOpenAIEndpointCapability(capability OpenAIEndpointCapability) bool { + if a == nil { + return false + } + if capability == "" { + return true + } + if !a.IsOpenAI() { + return false + } + switch capability { + case OpenAIEndpointCapabilityChatCompletions: + case OpenAIEndpointCapabilityEmbeddings: + if a.Type != AccountTypeAPIKey { + return false + } + default: + return false + } + + configured, found := a.openAIEndpointCapabilitySet() + if !found { + return true + } + return configured[string(capability)] +} + +func (a *Account) openAIEndpointCapabilitySet() (map[string]bool, bool) { + if a == nil || a.Credentials == nil { + return nil, false + } + raw, found := a.Credentials[openAIEndpointCapabilitiesCredentialKey] + if !found || raw == nil { + return nil, false + } + + result := make(map[string]bool) + add := func(value string) { + value = strings.ToLower(strings.TrimSpace(value)) + if value == "" { + return + } + result[value] = true + } + + switch capabilities := raw.(type) { + case []any: + for _, item := range capabilities { + if value, ok := item.(string); ok { + add(value) + } + } + case []string: + for _, value := range capabilities { + add(value) + } + case map[string]any: + for key, value := range capabilities { + enabled, ok := value.(bool) + if ok && enabled { + add(key) + } + } + case map[string]bool: + for key, enabled := range capabilities { + if enabled { + add(key) + } + } + } + + return result, true +} + func (a *Account) SupportsOpenAIImageCapability(capability OpenAIImagesCapability) bool { if !a.IsOpenAI() { return false diff --git a/backend/internal/service/openai_account_scheduler.go b/backend/internal/service/openai_account_scheduler.go index a8ac391a..1eca08b1 100644 --- a/backend/internal/service/openai_account_scheduler.go +++ b/backend/internal/service/openai_account_scheduler.go @@ -44,6 +44,7 @@ type OpenAIAccountScheduleRequest struct { PreviousResponseID string RequestedModel string RequiredTransport OpenAIUpstreamTransport + RequiredCapability OpenAIEndpointCapability RequiredImageCapability OpenAIImagesCapability RequireCompact bool ExcludedIDs map[int64]struct{} @@ -263,12 +264,13 @@ func (s *defaultOpenAIAccountScheduler) Select( previousResponseID := strings.TrimSpace(req.PreviousResponseID) if previousResponseID != "" { - selection, err := s.service.SelectAccountByPreviousResponseID( + selection, err := s.service.selectAccountByPreviousResponseIDForCapability( ctx, req.GroupID, previousResponseID, req.RequestedModel, req.ExcludedIDs, + req.RequiredCapability, req.RequireCompact, ) if err != nil { @@ -363,7 +365,7 @@ func (s *defaultOpenAIAccountScheduler) selectBySessionHash( _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash) return nil, nil } - account = s.service.recheckSelectedOpenAIAccountFromDB(ctx, account, req.RequestedModel, req.RequireCompact) + account = s.service.recheckSelectedOpenAIAccountFromDB(ctx, account, req.RequestedModel, req.RequireCompact, req.RequiredCapability) if account == nil || !s.isAccountTransportCompatible(account, req.RequiredTransport) { _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash) return nil, nil @@ -791,11 +793,11 @@ func (s *defaultOpenAIAccountScheduler) tryAcquireOpenAISelectionOrder( compactBlocked := false for i := 0; i < len(selectionOrder); i++ { candidate := selectionOrder[i] - fresh := s.service.resolveFreshSchedulableOpenAIAccount(ctx, candidate.account, req.RequestedModel, false) + fresh := s.service.resolveFreshSchedulableOpenAIAccount(ctx, candidate.account, req.RequestedModel, false, req.RequiredCapability) if fresh == nil || !s.isAccountTransportCompatible(fresh, req.RequiredTransport) || !s.isAccountRequestCompatible(ctx, fresh, req) { continue } - fresh = s.service.recheckSelectedOpenAIAccountFromDB(ctx, fresh, req.RequestedModel, false) + fresh = s.service.recheckSelectedOpenAIAccountFromDB(ctx, fresh, req.RequestedModel, false, req.RequiredCapability) if fresh == nil || !s.isAccountTransportCompatible(fresh, req.RequiredTransport) || !s.isAccountRequestCompatible(ctx, fresh, req) { continue } @@ -930,11 +932,11 @@ func (s *defaultOpenAIAccountScheduler) selectByLoadBalance( cfg := s.service.schedulingConfig() // WaitPlan.MaxConcurrency 使用 Concurrency(非 EffectiveLoadFactor),因为 WaitPlan 控制的是 Redis 实际并发槽位等待。 for _, candidate := range selectionOrder { - fresh := s.service.resolveFreshSchedulableOpenAIAccount(ctx, candidate.account, req.RequestedModel, false) + fresh := s.service.resolveFreshSchedulableOpenAIAccount(ctx, candidate.account, req.RequestedModel, false, req.RequiredCapability) if fresh == nil || !s.isAccountTransportCompatible(fresh, req.RequiredTransport) || !s.isAccountRequestCompatible(ctx, fresh, req) { continue } - fresh = s.service.recheckSelectedOpenAIAccountFromDB(ctx, fresh, req.RequestedModel, false) + fresh = s.service.recheckSelectedOpenAIAccountFromDB(ctx, fresh, req.RequestedModel, false, req.RequiredCapability) if fresh == nil || !s.isAccountTransportCompatible(fresh, req.RequiredTransport) || !s.isAccountRequestCompatible(ctx, fresh, req) { continue } @@ -981,7 +983,7 @@ func (s *defaultOpenAIAccountScheduler) isAccountRequestCompatible(ctx context.C s.service.isUpstreamModelRestrictedByChannel(ctx, *req.GroupID, account, req.RequestedModel, req.RequireCompact) { return false } - return account.SupportsOpenAIImageCapability(req.RequiredImageCapability) + return accountSupportsOpenAICapabilities(account, req.RequiredCapability, req.RequiredImageCapability) } func (s *defaultOpenAIAccountScheduler) ReportResult(accountID int64, success bool, firstTokenMs *int) { @@ -1104,7 +1106,21 @@ func (s *OpenAIGatewayService) SelectAccountWithScheduler( requiredTransport OpenAIUpstreamTransport, requireCompact bool, ) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) { - return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, "", requireCompact) + return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, "", "", requireCompact) +} + +func (s *OpenAIGatewayService) SelectAccountWithSchedulerForCapability( + ctx context.Context, + groupID *int64, + previousResponseID string, + sessionHash string, + requestedModel string, + excludedIDs map[int64]struct{}, + requiredTransport OpenAIUpstreamTransport, + requiredCapability OpenAIEndpointCapability, + requireCompact bool, +) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) { + return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, requiredCapability, "", requireCompact) } func (s *OpenAIGatewayService) SelectAccountWithSchedulerForImages( @@ -1115,13 +1131,13 @@ func (s *OpenAIGatewayService) SelectAccountWithSchedulerForImages( excludedIDs map[int64]struct{}, requiredCapability OpenAIImagesCapability, ) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) { - selection, decision, err := s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, requiredCapability, false) + selection, decision, err := s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, "", requiredCapability, false) if err == nil && selection != nil && selection.Account != nil { return selection, decision, nil } // 如果要求 native 能力(如指定了模型)但没有可用的 APIKey 账号,回退到 basic(OAuth 账号) if requiredCapability == OpenAIImagesCapabilityNative { - return s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, OpenAIImagesCapabilityBasic, false) + return s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, "", OpenAIImagesCapabilityBasic, false) } return selection, decision, err } @@ -1134,6 +1150,7 @@ func (s *OpenAIGatewayService) selectAccountWithScheduler( requestedModel string, excludedIDs map[int64]struct{}, requiredTransport OpenAIUpstreamTransport, + requiredCapability OpenAIEndpointCapability, requiredImageCapability OpenAIImagesCapability, requireCompact bool, ) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) { @@ -1144,14 +1161,14 @@ func (s *OpenAIGatewayService) selectAccountWithScheduler( if requiredTransport == OpenAIUpstreamTransportAny || requiredTransport == OpenAIUpstreamTransportHTTPSSE { effectiveExcludedIDs := cloneExcludedAccountIDs(excludedIDs) for { - selection, err := s.selectAccountWithLoadAwareness(ctx, groupID, sessionHash, requestedModel, effectiveExcludedIDs, requireCompact) + selection, err := s.selectAccountWithLoadAwareness(ctx, groupID, sessionHash, requestedModel, effectiveExcludedIDs, requireCompact, requiredCapability) if err != nil { return nil, decision, err } if selection == nil || selection.Account == nil { return selection, decision, nil } - if selection.Account.SupportsOpenAIImageCapability(requiredImageCapability) { + if accountSupportsOpenAICapabilities(selection.Account, requiredCapability, requiredImageCapability) { return selection, decision, nil } if selection.ReleaseFunc != nil { @@ -1169,14 +1186,15 @@ func (s *OpenAIGatewayService) selectAccountWithScheduler( effectiveExcludedIDs := cloneExcludedAccountIDs(excludedIDs) for { - selection, err := s.selectAccountWithLoadAwareness(ctx, groupID, sessionHash, requestedModel, effectiveExcludedIDs, requireCompact) + selection, err := s.selectAccountWithLoadAwareness(ctx, groupID, sessionHash, requestedModel, effectiveExcludedIDs, requireCompact, requiredCapability) if err != nil { return nil, decision, err } if selection == nil || selection.Account == nil { return selection, decision, nil } - if s.isOpenAIAccountTransportCompatible(selection.Account, requiredTransport) { + if s.isOpenAIAccountTransportCompatible(selection.Account, requiredTransport) && + accountSupportsOpenAICapabilities(selection.Account, requiredCapability, requiredImageCapability) { return selection, decision, nil } if selection.ReleaseFunc != nil { @@ -1213,12 +1231,21 @@ func (s *OpenAIGatewayService) selectAccountWithScheduler( PreviousResponseID: previousResponseID, RequestedModel: requestedModel, RequiredTransport: requiredTransport, + RequiredCapability: requiredCapability, RequiredImageCapability: requiredImageCapability, RequireCompact: requireCompact, ExcludedIDs: excludedIDs, }) } +func accountSupportsOpenAICapabilities(account *Account, requiredCapability OpenAIEndpointCapability, requiredImageCapability OpenAIImagesCapability) bool { + if account == nil { + return false + } + return account.SupportsOpenAIEndpointCapability(requiredCapability) && + account.SupportsOpenAIImageCapability(requiredImageCapability) +} + func cloneExcludedAccountIDs(excludedIDs map[int64]struct{}) map[int64]struct{} { if len(excludedIDs) == 0 { return nil diff --git a/backend/internal/service/openai_account_scheduler_test.go b/backend/internal/service/openai_account_scheduler_test.go index ba20ee5f..fedf7e9c 100644 --- a/backend/internal/service/openai_account_scheduler_test.go +++ b/backend/internal/service/openai_account_scheduler_test.go @@ -393,6 +393,64 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabled_Require require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) } +func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabled_EmbeddingsSkipsChatOnlyAccount(t *testing.T) { + resetOpenAIAdvancedSchedulerSettingCacheForTest() + + ctx := context.Background() + groupID := int64(10110) + accounts := []Account{ + { + ID: 36031, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 0, + Credentials: map[string]any{ + "openai_capabilities": []any{"chat_completions"}, + }, + }, + { + ID: 36032, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 5, + Credentials: map[string]any{ + "openai_capabilities": []any{"chat_completions", "embeddings"}, + }, + }, + } + cfg := &config.Config{} + cfg.Gateway.Scheduling.LoadBatchEnabled = false + svc := &OpenAIGatewayService{ + accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts}, + cache: &schedulerTestGatewayCache{}, + cfg: cfg, + concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{}), + } + + selection, decision, err := svc.SelectAccountWithSchedulerForCapability( + ctx, + &groupID, + "", + "", + "text-embedding-3-small", + nil, + OpenAIUpstreamTransportHTTPSSE, + OpenAIEndpointCapabilityEmbeddings, + false, + ) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, int64(36032), selection.Account.ID) + require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) +} + func TestOpenAIGatewayService_SelectAccountWithScheduler_EnabledUsesAdvancedPreviousResponseRouting(t *testing.T) { resetOpenAIAdvancedSchedulerSettingCacheForTest() @@ -458,6 +516,141 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_EnabledUsesAdvancedPrev require.True(t, decision.StickyPreviousHit) } +func TestOpenAIGatewayService_SelectAccountWithScheduler_Enabled_EmbeddingsSkipsChatOnlyAccount(t *testing.T) { + resetOpenAIAdvancedSchedulerSettingCacheForTest() + + ctx := context.Background() + groupID := int64(10111) + accounts := []Account{ + { + ID: 37011, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 0, + Credentials: map[string]any{ + "openai_capabilities": []any{"chat_completions"}, + }, + }, + { + ID: 37012, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 5, + Credentials: map[string]any{ + "openai_capabilities": []any{"chat_completions", "embeddings"}, + }, + }, + } + cfg := &config.Config{} + cfg.Gateway.Scheduling.LoadBatchEnabled = false + svc := &OpenAIGatewayService{ + accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts}, + cache: &schedulerTestGatewayCache{}, + cfg: cfg, + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"), + concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{}), + } + + selection, decision, err := svc.SelectAccountWithSchedulerForCapability( + ctx, + &groupID, + "", + "", + "text-embedding-3-small", + nil, + OpenAIUpstreamTransportHTTPSSE, + OpenAIEndpointCapabilityEmbeddings, + false, + ) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, int64(37012), selection.Account.ID) + require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) + require.Equal(t, 1, decision.CandidateCount) +} + +func TestOpenAIGatewayService_SelectAccountWithScheduler_Enabled_EmbeddingsSkipsChatOnlyStickyBindings(t *testing.T) { + resetOpenAIAdvancedSchedulerSettingCacheForTest() + + ctx := context.Background() + groupID := int64(10112) + accounts := []Account{ + { + ID: 37021, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 0, + Credentials: map[string]any{ + "openai_capabilities": []any{"chat_completions"}, + }, + Extra: map[string]any{ + "openai_apikey_responses_websockets_v2_enabled": true, + }, + }, + { + ID: 37022, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 5, + Credentials: map[string]any{ + "openai_capabilities": []any{"chat_completions", "embeddings"}, + }, + Extra: map[string]any{ + "openai_apikey_responses_websockets_v2_enabled": true, + }, + }, + } + cfg := newSchedulerTestOpenAIWSV2Config() + cfg.Gateway.Scheduling.LoadBatchEnabled = false + cache := &schedulerTestGatewayCache{ + sessionBindings: map[string]int64{ + "openai:session_hash_embeddings": 37021, + }, + } + svc := &OpenAIGatewayService{ + accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts}, + cache: cache, + cfg: cfg, + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"), + concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{}), + } + store := svc.getOpenAIWSStateStore() + require.NoError(t, store.BindResponseAccount(ctx, groupID, "resp_embeddings_chat_only", 37021, time.Hour)) + + selection, decision, err := svc.SelectAccountWithSchedulerForCapability( + ctx, + &groupID, + "resp_embeddings_chat_only", + "session_hash_embeddings", + "text-embedding-3-small", + nil, + OpenAIUpstreamTransportHTTPSSE, + OpenAIEndpointCapabilityEmbeddings, + false, + ) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, int64(37022), selection.Account.ID) + require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) + require.False(t, decision.StickyPreviousHit) + require.False(t, decision.StickySessionHit) + require.Equal(t, int64(37022), cache.sessionBindings["openai:session_hash_embeddings"]) +} + func TestOpenAIGatewayService_OpenAIAccountSchedulerMetrics_DisabledNoOp(t *testing.T) { resetOpenAIAdvancedSchedulerSettingCacheForTest() diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index f93cc221..77587f69 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -1279,7 +1279,7 @@ func (s *OpenAIGatewayService) SelectAccountForModel(ctx context.Context, groupI // SelectAccountForModelWithExclusions selects an account supporting the requested model while excluding specified accounts. // SelectAccountForModelWithExclusions 选择支持指定模型的账号,同时排除指定的账号。 func (s *OpenAIGatewayService) SelectAccountForModelWithExclusions(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}) (*Account, error) { - return s.selectAccountForModelWithExclusions(ctx, groupID, sessionHash, requestedModel, excludedIDs, false, 0) + return s.selectAccountForModelWithExclusions(ctx, groupID, sessionHash, requestedModel, excludedIDs, false, 0, "") } // noAvailableOpenAISelectionError builds the standard "no account available" error @@ -1312,13 +1312,16 @@ func openAICompactSupportTier(account *Account) int { // isOpenAIAccountEligibleForRequest centralises the schedulable / OpenAI / model / // compact-support checks used during account selection. -func isOpenAIAccountEligibleForRequest(ctx context.Context, account *Account, requestedModel string, requireCompact bool) bool { +func isOpenAIAccountEligibleForRequest(ctx context.Context, account *Account, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) bool { if account == nil || !account.IsOpenAI() || !account.IsSchedulableForModelWithContext(ctx, requestedModel) { return false } if requestedModel != "" && !account.IsModelSupported(requestedModel) { return false } + if !account.SupportsOpenAIEndpointCapability(requiredCapability) { + return false + } if requireCompact && openAICompactSupportTier(account) == 0 { return false } @@ -1366,7 +1369,7 @@ func resolveOpenAIAccountUpstreamModelForRequest(account *Account, requestedMode return upstreamModel } -func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, stickyAccountID int64) (*Account, error) { +func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, stickyAccountID int64, requiredCapability OpenAIEndpointCapability) (*Account, error) { if s.checkChannelPricingRestriction(ctx, groupID, requestedModel) { slog.Warn("channel pricing restriction blocked request", "group_id", derefGroupID(groupID), @@ -1376,7 +1379,7 @@ func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.C // 1. 尝试粘性会话命中 // Try sticky session hit - if account := s.tryStickySessionHit(ctx, groupID, sessionHash, requestedModel, excludedIDs, requireCompact, stickyAccountID); account != nil { + if account := s.tryStickySessionHit(ctx, groupID, sessionHash, requestedModel, excludedIDs, requireCompact, stickyAccountID, requiredCapability); account != nil { return account, nil } @@ -1389,7 +1392,7 @@ func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.C // 3. 按优先级 + LRU 选择最佳账号 // Select by priority + LRU - selected, compactBlocked := s.selectBestAccount(ctx, groupID, accounts, requestedModel, excludedIDs, requireCompact) + selected, compactBlocked := s.selectBestAccount(ctx, groupID, accounts, requestedModel, excludedIDs, requireCompact, requiredCapability) if selected == nil { return nil, noAvailableOpenAISelectionError(requestedModel, compactBlocked) @@ -1414,7 +1417,7 @@ func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.C // // tryStickySessionHit attempts to get account from sticky session. // Returns account if hit and usable; clears session and returns nil if account is unavailable. -func (s *OpenAIGatewayService) tryStickySessionHit(ctx context.Context, groupID *int64, sessionHash, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, stickyAccountID int64) *Account { +func (s *OpenAIGatewayService) tryStickySessionHit(ctx context.Context, groupID *int64, sessionHash, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, stickyAccountID int64, requiredCapability OpenAIEndpointCapability) *Account { if sessionHash == "" { return nil } @@ -1446,14 +1449,14 @@ func (s *OpenAIGatewayService) tryStickySessionHit(ctx context.Context, groupID // 验证账号是否可用于当前请求 // Verify account is usable for current request - if !isOpenAIAccountEligibleForRequest(ctx, account, requestedModel, false) { + if !isOpenAIAccountEligibleForRequest(ctx, account, requestedModel, false, requiredCapability) { return nil } if s.isOpenAIAccountRuntimeBlocked(account) { _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) return nil } - account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, requestedModel, requireCompact) + account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, requestedModel, requireCompact, requiredCapability) if account == nil { _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) return nil @@ -1477,7 +1480,7 @@ func (s *OpenAIGatewayService) tryStickySessionHit(ctx context.Context, groupID // Returns nil if no available account. The second return reports whether at // least one candidate was filtered out solely because it lacks compact support // (only meaningful when requireCompact=true). -func (s *OpenAIGatewayService) selectBestAccount(ctx context.Context, groupID *int64, accounts []Account, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool) (*Account, bool) { +func (s *OpenAIGatewayService) selectBestAccount(ctx context.Context, groupID *int64, accounts []Account, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, requiredCapability OpenAIEndpointCapability) (*Account, bool) { var selected *Account selectedCompactTier := -1 compactBlocked := false @@ -1492,11 +1495,11 @@ func (s *OpenAIGatewayService) selectBestAccount(ctx context.Context, groupID *i continue } - fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, requestedModel, false) + fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, requestedModel, false, requiredCapability) if fresh == nil { continue } - fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, false) + fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, false, requiredCapability) if fresh == nil { continue } @@ -1573,10 +1576,10 @@ func (s *OpenAIGatewayService) isBetterAccount(candidate, current *Account) bool // SelectAccountWithLoadAwareness selects an account with load-awareness and wait plan. func (s *OpenAIGatewayService) SelectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}) (*AccountSelectionResult, error) { - return s.selectAccountWithLoadAwareness(ctx, groupID, sessionHash, requestedModel, excludedIDs, false) + return s.selectAccountWithLoadAwareness(ctx, groupID, sessionHash, requestedModel, excludedIDs, false, "") } -func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool) (*AccountSelectionResult, error) { +func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, requiredCapability OpenAIEndpointCapability) (*AccountSelectionResult, error) { if s.checkChannelPricingRestriction(ctx, groupID, requestedModel) { slog.Warn("channel pricing restriction blocked request", "group_id", derefGroupID(groupID), @@ -1593,7 +1596,7 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex } } if s.concurrencyService == nil || !cfg.LoadBatchEnabled { - account, err := s.selectAccountForModelWithExclusions(ctx, groupID, sessionHash, requestedModel, excludedIDs, requireCompact, stickyAccountID) + account, err := s.selectAccountForModelWithExclusions(ctx, groupID, sessionHash, requestedModel, excludedIDs, requireCompact, stickyAccountID, requiredCapability) if err != nil { return nil, err } @@ -1646,8 +1649,8 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex if clearSticky { _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) } - if !clearSticky && isOpenAIAccountEligibleForRequest(ctx, account, requestedModel, false) { - account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, requestedModel, requireCompact) + if !clearSticky && isOpenAIAccountEligibleForRequest(ctx, account, requestedModel, false, requiredCapability) { + account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, requestedModel, requireCompact, requiredCapability) if account == nil { _ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash) } else if s.isOpenAIAccountRuntimeBlocked(account) { @@ -1691,15 +1694,12 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex // Scheduler snapshots can be temporarily stale (bucket rebuild is throttled); // re-check schedulability here so recently rate-limited/overloaded accounts // are not selected again before the bucket is rebuilt. - if !acc.IsSchedulable() { + if !isOpenAIAccountEligibleForRequest(ctx, acc, requestedModel, false, requiredCapability) { continue } if s.isOpenAIAccountRuntimeBlocked(acc) { continue } - if requestedModel != "" && !acc.IsModelSupported(requestedModel) { - continue - } if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, acc, requestedModel, requireCompact) { continue } @@ -1779,11 +1779,11 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex } for _, item := range selectionOrder { - fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, item.account, requestedModel, false) + fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, item.account, requestedModel, false, requiredCapability) if fresh == nil { continue } - fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, requireCompact) + fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, requireCompact, requiredCapability) if fresh == nil { continue } @@ -1813,11 +1813,11 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex ordered = prioritizeOpenAICompactAccounts(ordered) } for _, acc := range ordered { - fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, requestedModel, false) + fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, requestedModel, false, requiredCapability) if fresh == nil { continue } - fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, requireCompact) + fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, requireCompact, requiredCapability) if fresh == nil { continue } @@ -1858,11 +1858,11 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex candidates = prioritizeOpenAICompactAccounts(candidates) } for _, acc := range candidates { - fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, requestedModel, false) + fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, requestedModel, false, requiredCapability) if fresh == nil { continue } - fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, requireCompact) + fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, requireCompact, requiredCapability) if fresh == nil { continue } @@ -1910,7 +1910,7 @@ func (s *OpenAIGatewayService) tryAcquireAccountSlot(ctx context.Context, accoun return s.concurrencyService.AcquireAccountSlot(ctx, accountID, maxConcurrency) } -func (s *OpenAIGatewayService) resolveFreshSchedulableOpenAIAccount(ctx context.Context, account *Account, requestedModel string, requireCompact bool) *Account { +func (s *OpenAIGatewayService) resolveFreshSchedulableOpenAIAccount(ctx context.Context, account *Account, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) *Account { if account == nil { return nil } @@ -1924,7 +1924,7 @@ func (s *OpenAIGatewayService) resolveFreshSchedulableOpenAIAccount(ctx context. fresh = current } - if !isOpenAIAccountEligibleForRequest(ctx, fresh, requestedModel, requireCompact) { + if !isOpenAIAccountEligibleForRequest(ctx, fresh, requestedModel, requireCompact, requiredCapability) { return nil } if s.isOpenAIAccountRuntimeBlocked(fresh) { @@ -1933,12 +1933,12 @@ func (s *OpenAIGatewayService) resolveFreshSchedulableOpenAIAccount(ctx context. return fresh } -func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Context, account *Account, requestedModel string, requireCompact bool) *Account { +func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Context, account *Account, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) *Account { if account == nil { return nil } if s.schedulerSnapshot == nil || s.accountRepo == nil { - if !isOpenAIAccountEligibleForRequest(ctx, account, requestedModel, requireCompact) { + if !isOpenAIAccountEligibleForRequest(ctx, account, requestedModel, requireCompact, requiredCapability) { return nil } return account @@ -1948,7 +1948,7 @@ func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Co if err != nil || latest == nil { return nil } - if !isOpenAIAccountEligibleForRequest(ctx, latest, requestedModel, requireCompact) { + if !isOpenAIAccountEligibleForRequest(ctx, latest, requestedModel, requireCompact, requiredCapability) { return nil } if s.isOpenAIAccountRuntimeBlocked(latest) { diff --git a/backend/internal/service/openai_images_test.go b/backend/internal/service/openai_images_test.go index 854e9f6d..a87e96c1 100644 --- a/backend/internal/service/openai_images_test.go +++ b/backend/internal/service/openai_images_test.go @@ -413,6 +413,79 @@ func TestAccountSupportsOpenAIImageCapability_OAuthSupportsNative(t *testing.T) require.True(t, account.SupportsOpenAIImageCapability(OpenAIImagesCapabilityNative)) } +func TestAccountSupportsOpenAIEndpointCapability(t *testing.T) { + t.Run("OpenAI APIKey 默认兼容 chat 和 embeddings", func(t *testing.T) { + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + } + + require.True(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityChatCompletions)) + require.True(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityEmbeddings)) + }) + + t.Run("OpenAI OAuth 默认仅兼容 chat", func(t *testing.T) { + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + } + + require.True(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityChatCompletions)) + require.False(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityEmbeddings)) + }) + + t.Run("显式列表支持同时声明 chat 和 embeddings", func(t *testing.T) { + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "openai_capabilities": []any{"chat_completions", "embeddings"}, + }, + } + + require.True(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityChatCompletions)) + require.True(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityEmbeddings)) + }) + + t.Run("显式列表只声明 chat 时不支持 embeddings", func(t *testing.T) { + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "openai_capabilities": []any{"chat_completions"}, + }, + } + + require.True(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityChatCompletions)) + require.False(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityEmbeddings)) + }) + + t.Run("显式 map 支持单独关闭 chat 并开启 embeddings", func(t *testing.T) { + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "openai_capabilities": map[string]any{ + "chat_completions": false, + "embeddings": true, + }, + }, + } + + require.False(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityChatCompletions)) + require.True(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityEmbeddings)) + }) + + t.Run("未知能力不应默认放行", func(t *testing.T) { + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + } + + require.False(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapability("unknown"))) + }) +} + func TestBuildOpenAIImagesURL_HandlesVersionedBaseURL(t *testing.T) { require.Equal(t, "https://image-upstream.example/v1/images/generations", diff --git a/backend/internal/service/openai_ws_account_sticky_test.go b/backend/internal/service/openai_ws_account_sticky_test.go index 4005a921..c8b28a46 100644 --- a/backend/internal/service/openai_ws_account_sticky_test.go +++ b/backend/internal/service/openai_ws_account_sticky_test.go @@ -268,6 +268,52 @@ func TestOpenAIGatewayService_SelectAccountByPreviousResponseID_BusyKeepsSticky( require.Equal(t, int64(21), selection.WaitPlan.AccountID) } +func TestOpenAIGatewayService_SelectAccountByPreviousResponseID_CapabilityMismatchKeepsSticky(t *testing.T) { + ctx := context.Background() + groupID := int64(25) + account := Account{ + ID: 31, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{ + "openai_capabilities": []any{"chat_completions"}, + }, + Extra: map[string]any{ + "openai_apikey_responses_websockets_v2_enabled": true, + }, + } + cache := &stubGatewayCache{} + store := NewOpenAIWSStateStore(cache) + cfg := newOpenAIWSV2TestConfig() + svc := &OpenAIGatewayService{ + accountRepo: stubOpenAIAccountRepo{accounts: []Account{account}}, + cache: cache, + cfg: cfg, + concurrencyService: NewConcurrencyService(stubConcurrencyCache{}), + openaiWSStateStore: store, + } + + require.NoError(t, store.BindResponseAccount(ctx, groupID, "resp_prev_capability", account.ID, time.Hour)) + + selection, err := svc.selectAccountByPreviousResponseIDForCapability( + ctx, + &groupID, + "resp_prev_capability", + "text-embedding-3-small", + nil, + OpenAIEndpointCapabilityEmbeddings, + false, + ) + require.NoError(t, err) + require.Nil(t, selection) + boundAccountID, getErr := store.GetResponseAccount(ctx, groupID, "resp_prev_capability") + require.NoError(t, getErr) + require.Equal(t, account.ID, boundAccountID) +} + func newOpenAIWSV2TestConfig() *config.Config { cfg := &config.Config{} cfg.Gateway.OpenAIWS.Enabled = true diff --git a/backend/internal/service/openai_ws_forwarder.go b/backend/internal/service/openai_ws_forwarder.go index b8e558ae..5fd5cffc 100644 --- a/backend/internal/service/openai_ws_forwarder.go +++ b/backend/internal/service/openai_ws_forwarder.go @@ -3987,6 +3987,18 @@ func (s *OpenAIGatewayService) SelectAccountByPreviousResponseID( requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, +) (*AccountSelectionResult, error) { + return s.selectAccountByPreviousResponseIDForCapability(ctx, groupID, previousResponseID, requestedModel, excludedIDs, "", requireCompact) +} + +func (s *OpenAIGatewayService) selectAccountByPreviousResponseIDForCapability( + ctx context.Context, + groupID *int64, + previousResponseID string, + requestedModel string, + excludedIDs map[int64]struct{}, + requiredCapability OpenAIEndpointCapability, + requireCompact bool, ) (*AccountSelectionResult, error) { if s == nil { return nil, nil @@ -4027,12 +4039,31 @@ func (s *OpenAIGatewayService) SelectAccountByPreviousResponseID( if requestedModel != "" && !account.IsModelSupported(requestedModel) { return nil, nil } - account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, requestedModel, requireCompact) - if account == nil { - _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) + if !account.SupportsOpenAIEndpointCapability(requiredCapability) { return nil, nil } - // 兜底:若上游 compact 能力刚被探测为不支持,但 sticky 还在,需要主动放弃。 + if s.schedulerSnapshot != nil && s.accountRepo != nil { + latest, latestErr := s.accountRepo.GetByID(ctx, account.ID) + if latestErr != nil || latest == nil { + _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) + return nil, nil + } + if shouldClearStickySession(latest, requestedModel) || !latest.IsOpenAI() || !latest.IsSchedulable() { + _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) + return nil, nil + } + if requestedModel != "" && !latest.IsModelSupported(requestedModel) { + return nil, nil + } + if !latest.SupportsOpenAIEndpointCapability(requiredCapability) { + return nil, nil + } + if s.isOpenAIAccountRuntimeBlocked(latest) { + _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) + return nil, nil + } + account = latest + } if requireCompact && openAICompactSupportTier(account) == 0 { _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) return nil, nil diff --git a/frontend/src/components/account/CreateAccountModal.vue b/frontend/src/components/account/CreateAccountModal.vue index 331295f7..665c4695 100644 --- a/frontend/src/components/account/CreateAccountModal.vue +++ b/frontend/src/components/account/CreateAccountModal.vue @@ -2679,7 +2679,7 @@
@@ -2696,6 +2696,26 @@ />
+
+ +
+ +
+

{{ t('admin.accounts.openai.endpointCapabilitiesDesc') }}

+
@@ -3172,7 +3192,8 @@ import type { CreateAccountRequest, CodexSessionImportMessage, OpenAICompactMode, - OpenAIResponsesMode + OpenAIResponsesMode, + OpenAIEndpointCapability } from '@/types' import BaseDialog from '@/components/common/BaseDialog.vue' import ConfirmDialog from '@/components/common/ConfirmDialog.vue' @@ -3350,6 +3371,7 @@ const autoPauseOnExpired = ref(true) const openaiPassthroughEnabled = ref(false) const openAICompactMode = ref('auto') const openAIResponsesMode = ref('auto') +const openAIEndpointCapabilities = ref(['chat_completions', 'embeddings']) const openaiOAuthResponsesWebSocketV2Mode = ref(OPENAI_WS_MODE_OFF) const openaiAPIKeyResponsesWebSocketV2Mode = ref(OPENAI_WS_MODE_OFF) const codexCLIOnlyEnabled = ref(false) @@ -3412,6 +3434,43 @@ const openAIResponsesModeOptions = computed(() => [ { value: 'force_responses', label: t('admin.accounts.openai.responsesModeForceResponses') }, { value: 'force_chat_completions', label: t('admin.accounts.openai.responsesModeForceChatCompletions') } ]) +const openAIEndpointCapabilityOptions = computed<{ value: OpenAIEndpointCapability; label: string }[]>(() => [ + { value: 'chat_completions', label: t('admin.accounts.openai.capabilityChatCompletions') }, + { value: 'embeddings', label: t('admin.accounts.openai.capabilityEmbeddings') } +]) + +const normalizeOpenAIEndpointCapabilities = (values: OpenAIEndpointCapability[]) => { + const allowed: OpenAIEndpointCapability[] = ['chat_completions', 'embeddings'] + const selected = allowed.filter((value) => values.includes(value)) + return selected.length > 0 ? selected : allowed +} + +const toggleOpenAIEndpointCapability = (capability: OpenAIEndpointCapability, event?: Event) => { + if (openAIEndpointCapabilities.value.includes(capability)) { + if (openAIEndpointCapabilities.value.length <= 1) { + const input = event?.target as HTMLInputElement | null + if (input) input.checked = true + return + } + openAIEndpointCapabilities.value = openAIEndpointCapabilities.value.filter( + (value) => value !== capability + ) + return + } + openAIEndpointCapabilities.value = normalizeOpenAIEndpointCapabilities([ + ...openAIEndpointCapabilities.value, + capability + ]) +} + +const applyOpenAIEndpointCapabilities = (credentials: Record) => { + const capabilities = normalizeOpenAIEndpointCapabilities(openAIEndpointCapabilities.value) + if (capabilities.length === 2) { + delete credentials.openai_capabilities + return + } + credentials.openai_capabilities = capabilities +} function buildAntigravityExtra(): Record | undefined { const extra: Record = {} @@ -3721,6 +3780,7 @@ watch( } if (newPlatform !== 'openai') { openaiPassthroughEnabled.value = false + openAIEndpointCapabilities.value = ['chat_completions', 'embeddings'] openaiOAuthResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF openaiAPIKeyResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF codexCLIOnlyEnabled.value = false @@ -4120,6 +4180,7 @@ const resetForm = () => { openaiPassthroughEnabled.value = false openAICompactMode.value = 'auto' openAIResponsesMode.value = 'auto' + openAIEndpointCapabilities.value = ['chat_completions', 'embeddings'] openaiOAuthResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF openaiAPIKeyResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF codexCLIOnlyEnabled.value = false @@ -4498,6 +4559,7 @@ const handleSubmit = async () => { } } if (form.platform === 'openai') { + applyOpenAIEndpointCapabilities(credentials) const compactModelMapping = buildOpenAICompactModelMapping() if (compactModelMapping) { credentials.compact_model_mapping = compactModelMapping @@ -4620,6 +4682,9 @@ const createAccountAndFinish = async ( } } if (platform === 'openai') { + if (type === 'apikey') { + applyOpenAIEndpointCapabilities(credentials) + } const compactModelMapping = buildOpenAICompactModelMapping() if (compactModelMapping) { credentials.compact_model_mapping = compactModelMapping diff --git a/frontend/src/components/account/EditAccountModal.vue b/frontend/src/components/account/EditAccountModal.vue index f44b5d38..3cb10591 100644 --- a/frontend/src/components/account/EditAccountModal.vue +++ b/frontend/src/components/account/EditAccountModal.vue @@ -1439,7 +1439,7 @@
@@ -1459,6 +1459,26 @@
{{ t(openAIResponsesStatusKey) }}
+
+ +
+ +
+

{{ t('admin.accounts.openai.endpointCapabilitiesDesc') }}

+
@@ -2245,7 +2265,15 @@ import { useAppStore } from '@/stores/app' import { useAuthStore } from '@/stores/auth' import { adminAPI } from '@/api/admin' import { useQuotaNotifyState } from '@/composables/useQuotaNotifyState' -import type { Account, Proxy, AdminGroup, CheckMixedChannelResponse, OpenAICompactMode, OpenAIResponsesMode } from '@/types' +import type { + Account, + Proxy, + AdminGroup, + CheckMixedChannelResponse, + OpenAICompactMode, + OpenAIResponsesMode, + OpenAIEndpointCapability +} from '@/types' import BaseDialog from '@/components/common/BaseDialog.vue' import ConfirmDialog from '@/components/common/ConfirmDialog.vue' import Select from '@/components/common/Select.vue' @@ -2433,6 +2461,7 @@ const customBaseUrl = ref('') const openaiPassthroughEnabled = ref(false) const openAICompactMode = ref('auto') const openAIResponsesMode = ref('auto') +const openAIEndpointCapabilities = ref(['chat_completions', 'embeddings']) const openaiOAuthResponsesWebSocketV2Mode = ref(OPENAI_WS_MODE_OFF) const openaiAPIKeyResponsesWebSocketV2Mode = ref(OPENAI_WS_MODE_OFF) const codexCLIOnlyEnabled = ref(false) @@ -2539,6 +2568,63 @@ const openAIResponsesModeOptions = computed(() => [ { value: 'force_responses', label: t('admin.accounts.openai.responsesModeForceResponses') }, { value: 'force_chat_completions', label: t('admin.accounts.openai.responsesModeForceChatCompletions') } ]) +const openAIEndpointCapabilityOptions = computed<{ value: OpenAIEndpointCapability; label: string }[]>(() => [ + { value: 'chat_completions', label: t('admin.accounts.openai.capabilityChatCompletions') }, + { value: 'embeddings', label: t('admin.accounts.openai.capabilityEmbeddings') } +]) + +const normalizeOpenAIEndpointCapabilities = (values: OpenAIEndpointCapability[]) => { + const allowed: OpenAIEndpointCapability[] = ['chat_completions', 'embeddings'] + const selected = allowed.filter((value) => values.includes(value)) + return selected.length > 0 ? selected : allowed +} + +const readOpenAIEndpointCapabilities = (credentials?: Record): OpenAIEndpointCapability[] => { + const raw = credentials?.openai_capabilities + if (Array.isArray(raw)) { + return normalizeOpenAIEndpointCapabilities( + raw.filter((value): value is OpenAIEndpointCapability => + value === 'chat_completions' || value === 'embeddings' + ) + ) + } + if (raw !== null && typeof raw === 'object') { + const capabilityMap = raw as Record + return normalizeOpenAIEndpointCapabilities( + openAIEndpointCapabilityOptions.value + .map((option) => option.value) + .filter((value) => capabilityMap[value] === true) + ) + } + return ['chat_completions', 'embeddings'] +} + +const toggleOpenAIEndpointCapability = (capability: OpenAIEndpointCapability, event?: Event) => { + if (openAIEndpointCapabilities.value.includes(capability)) { + if (openAIEndpointCapabilities.value.length <= 1) { + const input = event?.target as HTMLInputElement | null + if (input) input.checked = true + return + } + openAIEndpointCapabilities.value = openAIEndpointCapabilities.value.filter( + (value) => value !== capability + ) + return + } + openAIEndpointCapabilities.value = normalizeOpenAIEndpointCapabilities([ + ...openAIEndpointCapabilities.value, + capability + ]) +} + +const applyOpenAIEndpointCapabilities = (credentials: Record) => { + const capabilities = normalizeOpenAIEndpointCapabilities(openAIEndpointCapabilities.value) + if (capabilities.length === 2) { + delete credentials.openai_capabilities + return + } + credentials.openai_capabilities = capabilities +} const normalizeOpenAIResponsesMode = (mode: unknown): OpenAIResponsesMode => { if (mode === 'force_responses' || mode === 'force_chat_completions') { return mode @@ -2724,6 +2810,7 @@ const syncFormFromAccount = (newAccount: Account | null) => { openaiPassthroughEnabled.value = false openAICompactMode.value = 'auto' openAIResponsesMode.value = 'auto' + openAIEndpointCapabilities.value = ['chat_completions', 'embeddings'] openAICompactModelMappings.value = [] openaiOAuthResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF openaiAPIKeyResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF @@ -2736,6 +2823,9 @@ const syncFormFromAccount = (newAccount: Account | null) => { openAICompactMode.value = (extra?.openai_compact_mode as OpenAICompactMode) || 'auto' if (newAccount.type === 'apikey') { openAIResponsesMode.value = normalizeOpenAIResponsesMode(extra?.openai_responses_mode) + openAIEndpointCapabilities.value = readOpenAIEndpointCapabilities( + newAccount.credentials as Record | undefined + ) } const codexImageGenerationBridgeValue = typeof extra?.codex_image_generation_bridge === 'boolean' ? extra.codex_image_generation_bridge @@ -3476,6 +3566,7 @@ const handleSubmit = async () => { newCredentials.model_mapping = currentCredentials.model_mapping } if (props.account.platform === 'openai') { + applyOpenAIEndpointCapabilities(newCredentials) const compactModelMapping = buildModelMappingObject('mapping', [], openAICompactModelMappings.value) if (compactModelMapping) { newCredentials.compact_model_mapping = compactModelMapping diff --git a/frontend/src/components/account/__tests__/EditAccountModal.spec.ts b/frontend/src/components/account/__tests__/EditAccountModal.spec.ts index 0b8e939c..db012a30 100644 --- a/frontend/src/components/account/__tests__/EditAccountModal.spec.ts +++ b/frontend/src/components/account/__tests__/EditAccountModal.spec.ts @@ -310,6 +310,63 @@ describe('EditAccountModal', () => { expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.openai_responses_supported).toBe(true) }) + it('submits OpenAI APIKey endpoint capabilities from credentials', async () => { + const account = buildAccount() + account.credentials.openai_capabilities = ['chat_completions'] + updateAccountMock.mockReset() + checkMixedChannelRiskMock.mockReset() + checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false }) + updateAccountMock.mockResolvedValue(account) + + const wrapper = mountModal(account) + + expect(wrapper.findAll('input[type="checkbox"]').some((input) => (input.element as HTMLInputElement).checked)).toBe(true) + + await wrapper.get('form#edit-account-form').trigger('submit.prevent') + + expect(updateAccountMock).toHaveBeenCalledTimes(1) + expect(updateAccountMock.mock.calls[0]?.[1]?.credentials?.openai_capabilities).toEqual([ + 'chat_completions' + ]) + }) + + it('keeps at least one OpenAI APIKey endpoint capability selected', async () => { + const account = buildAccount() + updateAccountMock.mockReset() + checkMixedChannelRiskMock.mockReset() + checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false }) + updateAccountMock.mockResolvedValue(account) + + const wrapper = mountModal(account) + + const chatCheckbox = wrapper.get( + '[data-testid="openai-endpoint-capability-chat_completions"]' + ) + const embeddingsCheckbox = wrapper.get( + '[data-testid="openai-endpoint-capability-embeddings"]' + ) + + expect(chatCheckbox.element.checked).toBe(true) + expect(embeddingsCheckbox.element.checked).toBe(true) + + await embeddingsCheckbox.setValue(false) + + expect(chatCheckbox.element.checked).toBe(true) + expect(embeddingsCheckbox.element.checked).toBe(false) + + await chatCheckbox.setValue(false) + + expect(chatCheckbox.element.checked).toBe(true) + expect(embeddingsCheckbox.element.checked).toBe(false) + + await wrapper.get('form#edit-account-form').trigger('submit.prevent') + + expect(updateAccountMock).toHaveBeenCalledTimes(1) + expect(updateAccountMock.mock.calls[0]?.[1]?.credentials?.openai_capabilities).toEqual([ + 'chat_completions' + ]) + }) + it('submits account-level Codex image generation bridge override', async () => { const account = buildAccount() account.extra = { diff --git a/frontend/src/i18n/locales/en.ts b/frontend/src/i18n/locales/en.ts index 41c3c495..ec659dd4 100644 --- a/frontend/src/i18n/locales/en.ts +++ b/frontend/src/i18n/locales/en.ts @@ -3353,6 +3353,11 @@ export default { responsesModeAuto: 'Auto', responsesModeForceResponses: 'Force Responses', responsesModeForceChatCompletions: 'Force Chat Completions', + endpointCapabilities: 'Endpoint capabilities', + endpointCapabilitiesDesc: + 'Used by account routing. Both endpoints are allowed by default; if the upstream only supports one, select only the supported endpoint.', + capabilityChatCompletions: 'Chat Completions', + capabilityEmbeddings: 'Embeddings', responsesStatusAutoSupported: 'Auto probe: Responses', responsesStatusAutoUnsupported: 'Auto probe: Chat Completions', responsesStatusAutoUnknown: 'Auto probe: unknown', diff --git a/frontend/src/i18n/locales/zh.ts b/frontend/src/i18n/locales/zh.ts index 8ff8ea80..36b0d8c3 100644 --- a/frontend/src/i18n/locales/zh.ts +++ b/frontend/src/i18n/locales/zh.ts @@ -3499,6 +3499,11 @@ export default { responsesModeAuto: '自动', responsesModeForceResponses: '强制 Responses', responsesModeForceChatCompletions: '强制 Chat Completions', + endpointCapabilities: '端点能力', + endpointCapabilitiesDesc: + '用于调度筛选。默认两个端点都可用;如果上游只支持其中一个,请只勾选实际支持的端点。', + capabilityChatCompletions: 'Chat Completions', + capabilityEmbeddings: 'Embeddings', responsesStatusAutoSupported: '自动探测:Responses', responsesStatusAutoUnsupported: '自动探测:Chat Completions', responsesStatusAutoUnknown: '自动探测:未探测', diff --git a/frontend/src/types/index.ts b/frontend/src/types/index.ts index eae5e455..c2136169 100644 --- a/frontend/src/types/index.ts +++ b/frontend/src/types/index.ts @@ -997,6 +997,7 @@ export interface CodexUsageSnapshot { export type OpenAICompactMode = 'auto' | 'force_on' | 'force_off' export type OpenAIResponsesMode = 'auto' | 'force_responses' | 'force_chat_completions' +export type OpenAIEndpointCapability = 'chat_completions' | 'embeddings' export interface OpenAICompactState { openai_compact_mode?: OpenAICompactMode From 37044b83eb67ed8680c3edc7249c9c2a75dd56fa Mon Sep 17 00:00:00 2001 From: shaw Date: Fri, 29 May 2026 09:23:06 +0800 Subject: [PATCH 29/79] fix(openai): clarify endpoint capability UI --- .../components/account/CreateAccountModal.vue | 31 +++++++++++++- .../components/account/EditAccountModal.vue | 42 +++++++++++++++++-- .../__tests__/EditAccountModal.spec.ts | 31 ++++++++++++++ frontend/src/i18n/locales/en.ts | 10 ++++- frontend/src/i18n/locales/zh.ts | 9 +++- 5 files changed, 114 insertions(+), 9 deletions(-) diff --git a/frontend/src/components/account/CreateAccountModal.vue b/frontend/src/components/account/CreateAccountModal.vue index 665c4695..5f5a11d7 100644 --- a/frontend/src/components/account/CreateAccountModal.vue +++ b/frontend/src/components/account/CreateAccountModal.vue @@ -2692,10 +2692,18 @@ +

{{ t('admin.accounts.autoPauseThresholdHint') }}

+
+
+ + +

{{ t('admin.accounts.autoPauseThresholdHint') }}

+
+
+
([]) const customErrorCodeInput = ref(null) const interceptWarmupRequests = ref(false) const autoPauseOnExpired = ref(false) +const autoPause5hThreshold = ref(null) +const autoPause7dThreshold = ref(null) const mixedScheduling = ref(false) // For antigravity accounts: enable mixed scheduling const allowOverages = ref(false) // For antigravity accounts: enable AI Credits overages const antigravityModelRestrictionMode = ref<'whitelist' | 'mapping'>('whitelist') @@ -2862,9 +2896,11 @@ const syncFormFromAccount = (newAccount: Account | null) => { // Load mixed scheduling setting (only for antigravity accounts) mixedScheduling.value = false allowOverages.value = false - const extra = newAccount.extra as Record | undefined - mixedScheduling.value = extra?.mixed_scheduling === true - allowOverages.value = extra?.allow_overages === true + const extra = newAccount.extra as Record | undefined + mixedScheduling.value = extra?.mixed_scheduling === true + allowOverages.value = extra?.allow_overages === true + autoPause5hThreshold.value = typeof extra?.auto_pause_5h_threshold === 'number' ? extra.auto_pause_5h_threshold * 100 : null + autoPause7dThreshold.value = typeof extra?.auto_pause_7d_threshold === 'number' ? extra.auto_pause_7d_threshold * 100 : null // Load OpenAI passthrough toggle (OpenAI OAuth/API Key) openaiPassthroughEnabled.value = false @@ -3987,9 +4023,9 @@ const handleSubmit = async () => { } // For OpenAI OAuth/API Key accounts, handle passthrough mode in extra - if (props.account.platform === 'openai' && (props.account.type === 'oauth' || props.account.type === 'apikey')) { - const currentExtra = (props.account.extra as Record) || {} - const newExtra: Record = { ...currentExtra } + if (props.account.platform === 'openai' && (props.account.type === 'oauth' || props.account.type === 'apikey')) { + const currentExtra = (props.account.extra as Record) || {} + const newExtra: Record = { ...currentExtra } const hadCodexCLIOnlyEnabled = currentExtra.codex_cli_only === true if (props.account.type === 'oauth') { newExtra.openai_oauth_responses_websockets_v2_mode = openaiOAuthResponsesWebSocketV2Mode.value @@ -4011,15 +4047,25 @@ const handleSubmit = async () => { } else { newExtra.openai_compact_mode = openAICompactMode.value } - if (props.account.type === 'apikey') { + if (props.account.type === 'apikey') { if (!openAITextGenerationCapabilityEnabled.value || openAIResponsesMode.value === 'auto') { delete newExtra.openai_responses_mode } else { newExtra.openai_responses_mode = openAIResponsesMode.value } - } + } + if (autoPause5hThreshold.value != null && autoPause5hThreshold.value > 0) { + newExtra.auto_pause_5h_threshold = autoPause5hThreshold.value / 100 + } else { + delete newExtra.auto_pause_5h_threshold + } + if (autoPause7dThreshold.value != null && autoPause7dThreshold.value > 0) { + newExtra.auto_pause_7d_threshold = autoPause7dThreshold.value / 100 + } else { + delete newExtra.auto_pause_7d_threshold + } - delete newExtra.codex_image_generation_bridge_enabled + delete newExtra.codex_image_generation_bridge_enabled if (codexImageGenerationBridgeMode.value === 'inherit') { delete newExtra.codex_image_generation_bridge } else { diff --git a/frontend/src/components/account/__tests__/EditAccountModal.spec.ts b/frontend/src/components/account/__tests__/EditAccountModal.spec.ts index 4561924f..6db63831 100644 --- a/frontend/src/components/account/__tests__/EditAccountModal.spec.ts +++ b/frontend/src/components/account/__tests__/EditAccountModal.spec.ts @@ -330,6 +330,28 @@ describe('EditAccountModal', () => { ]) }) + it('submits OpenAI quota auto-pause thresholds in extra', async () => { + const account = buildAccount() + account.extra = { + auto_pause_5h_threshold: 0.9, + auto_pause_7d_threshold: 0.8 + } + updateAccountMock.mockReset() + checkMixedChannelRiskMock.mockReset() + checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false }) + updateAccountMock.mockResolvedValue(account) + + const wrapper = mountModal(account) + + await wrapper.get('[data-testid="auto-pause-5h-threshold"]').setValue('95') + await wrapper.get('[data-testid="auto-pause-7d-threshold"]').setValue('96') + await wrapper.get('form#edit-account-form').trigger('submit.prevent') + + expect(updateAccountMock).toHaveBeenCalledTimes(1) + expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.auto_pause_5h_threshold).toBe(0.95) + expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.auto_pause_7d_threshold).toBe(0.96) + }) + it('keeps at least one OpenAI APIKey endpoint capability selected', async () => { const account = buildAccount() updateAccountMock.mockReset() diff --git a/frontend/src/i18n/locales/en.ts b/frontend/src/i18n/locales/en.ts index b2aeb2f8..fa2e5a92 100644 --- a/frontend/src/i18n/locales/en.ts +++ b/frontend/src/i18n/locales/en.ts @@ -3475,6 +3475,9 @@ export default { 'When enabled, warmup requests like title generation will return mock responses without consuming upstream tokens', autoPauseOnExpired: 'Auto Pause On Expired', autoPauseOnExpiredDesc: 'When enabled, the account will auto pause scheduling after it expires', + autoPause5hThreshold: '5h Usage Threshold (%)', + autoPause7dThreshold: '7d Usage Threshold (%)', + autoPauseThresholdHint: 'Leave empty or set 0 to disable. Reaching the threshold only skips the account during scheduling and does not modify schedulable.', // Quota control (Anthropic OAuth/SetupToken only) quotaControl: { title: 'Quota Control', diff --git a/frontend/src/i18n/locales/zh.ts b/frontend/src/i18n/locales/zh.ts index 85d1feee..2364f9c4 100644 --- a/frontend/src/i18n/locales/zh.ts +++ b/frontend/src/i18n/locales/zh.ts @@ -3613,6 +3613,9 @@ export default { interceptWarmupRequestsDesc: '启用后,标题生成等预热请求将返回 mock 响应,不消耗上游 token', autoPauseOnExpired: '过期自动暂停调度', autoPauseOnExpiredDesc: '启用后,账号过期将自动暂停调度', + autoPause5hThreshold: '5h 用量阈值(%)', + autoPause7dThreshold: '7d 用量阈值(%)', + autoPauseThresholdHint: '填 0 或留空表示不启用;达到阈值后仅在调度时跳过账号,不修改 schedulable。', // Quota control (Anthropic OAuth/SetupToken only) quotaControl: { title: '配额控制', From 8b7a8227060e3eac20b5c0e331b50b2dc99f5e13 Mon Sep 17 00:00:00 2001 From: wucm667 Date: Fri, 29 May 2026 12:20:30 +0800 Subject: [PATCH 35/79] fix(account): address review on OpenAI quota auto-pause - gate previous_response_id sticky path with quota auto-pause check at both the snapshot and DB-recheck stages (previously bypassed, #1) - skip pausing when the usage window already reset to avoid a stale stuck-pause; carry codex_*_reset_at / reset_after_seconds / codex_usage_updated_at through the scheduler snapshot whitelist (#2) - remove the incomplete limit mode; percentage threshold only (#3) - add global default 5h/7d threshold inputs to the Ops settings dialog with validation and en/zh i18n (#4) - downgrade account_auto_paused_by_quota log from Info to Debug; it fires per-candidate on the scheduling hot path (#5) Co-Authored-By: Claude Opus 4.8 --- .../internal/repository/scheduler_cache.go | 7 +- .../repository/scheduler_cache_unit_test.go | 18 +++-- .../service/openai_account_scheduler_test.go | 58 +++++++++++++-- .../service/openai_gateway_service.go | 72 +++++++++++++------ .../service/openai_ws_account_sticky_test.go | 40 +++++++++++ .../internal/service/openai_ws_forwarder.go | 10 +++ frontend/src/api/admin/ops.ts | 6 ++ frontend/src/i18n/locales/en.ts | 8 ++- frontend/src/i18n/locales/zh.ts | 8 ++- .../ops/components/OpsSettingsDialog.vue | 65 +++++++++++++++++ 10 files changed, 257 insertions(+), 35 deletions(-) diff --git a/backend/internal/repository/scheduler_cache.go b/backend/internal/repository/scheduler_cache.go index ec8c72dc..ff3c4301 100644 --- a/backend/internal/repository/scheduler_cache.go +++ b/backend/internal/repository/scheduler_cache.go @@ -550,10 +550,13 @@ func filterSchedulerExtra(extra map[string]any) map[string]any { "openai_responses_supported", "codex_5h_used_percent", "codex_7d_used_percent", + "codex_5h_reset_at", + "codex_7d_reset_at", + "codex_5h_reset_after_seconds", + "codex_7d_reset_after_seconds", + "codex_usage_updated_at", "auto_pause_5h_threshold", "auto_pause_7d_threshold", - "auto_pause_5h_limit", - "auto_pause_7d_limit", } 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 fabc6bad..9e4ec23e 100644 --- a/backend/internal/repository/scheduler_cache_unit_test.go +++ b/backend/internal/repository/scheduler_cache_unit_test.go @@ -80,10 +80,15 @@ func TestBuildSchedulerMetadataAccount_KeepsQuotaAutoPauseFields(t *testing.T) { account := service.Account{ ID: 88, Extra: map[string]any{ - "codex_5h_used_percent": 12.34, - "codex_7d_used_percent": 56.78, - "auto_pause_5h_threshold": 0.95, - "auto_pause_7d_threshold": 0.96, + "codex_5h_used_percent": 12.34, + "codex_7d_used_percent": 56.78, + "codex_5h_reset_at": "2026-05-29T10:00:00Z", + "codex_7d_reset_at": "2026-06-01T10:00:00Z", + "codex_5h_reset_after_seconds": 300, + "codex_7d_reset_after_seconds": 600, + "codex_usage_updated_at": "2026-05-29T09:00:00Z", + "auto_pause_5h_threshold": 0.95, + "auto_pause_7d_threshold": 0.96, }, } @@ -91,6 +96,11 @@ func TestBuildSchedulerMetadataAccount_KeepsQuotaAutoPauseFields(t *testing.T) { require.Equal(t, 12.34, got.Extra["codex_5h_used_percent"]) require.Equal(t, 56.78, got.Extra["codex_7d_used_percent"]) + require.Equal(t, "2026-05-29T10:00:00Z", got.Extra["codex_5h_reset_at"]) + require.Equal(t, "2026-06-01T10:00:00Z", got.Extra["codex_7d_reset_at"]) + require.Equal(t, 300, got.Extra["codex_5h_reset_after_seconds"]) + require.Equal(t, 600, got.Extra["codex_7d_reset_after_seconds"]) + require.Equal(t, "2026-05-29T09:00:00Z", got.Extra["codex_usage_updated_at"]) require.Equal(t, 0.95, got.Extra["auto_pause_5h_threshold"]) require.Equal(t, 0.96, got.Extra["auto_pause_7d_threshold"]) } diff --git a/backend/internal/service/openai_account_scheduler_test.go b/backend/internal/service/openai_account_scheduler_test.go index 37810870..531769a7 100644 --- a/backend/internal/service/openai_account_scheduler_test.go +++ b/backend/internal/service/openai_account_scheduler_test.go @@ -704,7 +704,6 @@ func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_AutoPauseBy5hT Extra: map[string]any{ "codex_5h_used_percent": 95.0, "auto_pause_5h_threshold": 0.95, - "auto_pause_5h_limit": 100, }, } secondary := Account{ID: 35002, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5} @@ -729,7 +728,6 @@ func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_AllowsBelow5hT Extra: map[string]any{ "codex_5h_used_percent": 80.0, "auto_pause_5h_threshold": 0.95, - "auto_pause_5h_limit": 100, }, } secondary := Account{ID: 35102, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5} @@ -754,7 +752,6 @@ func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_AutoPauseBy7dT Extra: map[string]any{ "codex_7d_used_percent": 95.0, "auto_pause_7d_threshold": 0.95, - "auto_pause_7d_limit": 200, }, } secondary := Account{ID: 35202, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5} @@ -790,7 +787,6 @@ func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_UsesGlobalDefa Priority: 0, Extra: map[string]any{ "codex_5h_used_percent": 95.0, - "auto_pause_5h_limit": 100, }, } secondary := Account{ID: 35402, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5} @@ -802,6 +798,60 @@ func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_UsesGlobalDefa require.Equal(t, int64(35402), account.ID) } +func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_StaleUsageWindowResetSkipsPause(t *testing.T) { + ctx := context.Background() + // Usage is over threshold but the window's reset time has already passed, so the + // cached percentage is stale (the real window rolled over) and the account must NOT + // stay paused — otherwise it could be skipped forever with no traffic to refresh it. + primary := Account{ + ID: 35501, + 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.Minute).Format(time.RFC3339), + }, + } + secondary := Account{ID: 35502, 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(35501), account.ID) +} + +func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_FreshUsageWindowStillPauses(t *testing.T) { + ctx := context.Background() + // Same as above but the window has not reset yet, so the account stays paused. + primary := Account{ + ID: 35601, + 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), + }, + } + secondary := Account{ID: 35602, 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(35602), 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 268f985c..e2534cc2 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -1328,13 +1328,12 @@ func isOpenAIAccountEligibleForRequest(ctx context.Context, account *Account, re return false } if paused, reason := shouldAutoPauseOpenAIAccountByQuota(ctx, account); paused { - slog.Info("account_auto_paused_by_quota", + // Debug level: this fires per-candidate on the scheduling hot path, so Info + // would amplify into log spam once several accounts cross the threshold. + slog.Debug("account_auto_paused_by_quota", "account_id", account.ID, - "usage_5h_percent", readOpenAIQuotaUsedPercent(account.Extra, "5h"), - "usage_7d_percent", readOpenAIQuotaUsedPercent(account.Extra, "7d"), - "threshold_type", reason.window, + "window", reason.window, "threshold", reason.threshold, - "limit", reason.limit, "utilization", reason.utilization, ) return false @@ -1354,7 +1353,6 @@ func isOpenAIAccountEligibleForRequest(ctx context.Context, account *Account, re type openAIQuotaAutoPauseDecision struct { window string threshold float64 - limit float64 utilization float64 } @@ -1363,18 +1361,15 @@ func shouldAutoPauseOpenAIAccountByQuota(ctx context.Context, account *Account) return false, openAIQuotaAutoPauseDecision{} } threshold5h, threshold7d := resolveOpenAIQuotaAutoPauseThresholds(ctx, account) + now := time.Now() if threshold5h > 0 { - if utilization, limit, ok := resolveOpenAIQuotaUtilization(account.Extra, "5h"); ok { - if utilization >= threshold5h { - return true, openAIQuotaAutoPauseDecision{window: "5h", threshold: threshold5h, limit: limit, utilization: utilization} - } + if utilization, ok := resolveOpenAIQuotaUtilization(account.Extra, "5h", now); ok && utilization >= threshold5h { + return true, openAIQuotaAutoPauseDecision{window: "5h", threshold: threshold5h, utilization: utilization} } } if threshold7d > 0 { - if utilization, limit, ok := resolveOpenAIQuotaUtilization(account.Extra, "7d"); ok { - if utilization >= threshold7d { - return true, openAIQuotaAutoPauseDecision{window: "7d", threshold: threshold7d, limit: limit, utilization: utilization} - } + if utilization, ok := resolveOpenAIQuotaUtilization(account.Extra, "7d", now); ok && utilization >= threshold7d { + return true, openAIQuotaAutoPauseDecision{window: "7d", threshold: threshold7d, utilization: utilization} } } return false, openAIQuotaAutoPauseDecision{} @@ -1431,18 +1426,49 @@ func resolveAccountExtraNumber(extra map[string]any, keys ...string) (float64, b return 0, false } -func resolveOpenAIQuotaUtilization(extra map[string]any, window string) (float64, float64, bool) { - limitKeys := []string{"auto_pause_" + window + "_limit", "quota_" + window + "_limit", window + "_limit"} - if limit, ok := resolveAccountExtraNumber(extra, limitKeys...); ok && limit > 0 { - if usage, ok := resolveAccountExtraNumber(extra, "usage_"+window); ok && usage >= 0 { - return usage / limit, limit, true - } - } +// resolveOpenAIQuotaUtilization returns the current utilization ratio (0..1) for the +// given Codex usage window. ok=false means there is no usable signal to pause on: +// either no snapshot exists, or the window has already rolled over so the cached +// percentage is stale. The stale guard matters because a paused account stops +// receiving requests, so its snapshot is never refreshed from upstream headers — +// without this check an old used_percent would keep the account paused forever even +// after the real window reset. +func resolveOpenAIQuotaUtilization(extra map[string]any, window string, now time.Time) (float64, bool) { usedPercent := readOpenAIQuotaUsedPercent(extra, window) if usedPercent <= 0 { - return 0, 0, false + return 0, false } - return usedPercent / 100, 100, true + if openAIQuotaWindowReset(extra, window, now) { + return 0, false + } + return usedPercent / 100, true +} + +// 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 +// codex_usage_updated_at, mirroring AccountUsageService's window-progress logic. +func openAIQuotaWindowReset(extra map[string]any, window string, now time.Time) bool { + if len(extra) == 0 { + return false + } + if resetAtRaw, ok := extra["codex_"+window+"_reset_at"]; ok { + if resetAt, err := parseTime(fmt.Sprint(resetAtRaw)); err == nil { + return !now.Before(resetAt) + } + } + resetAfter := parseExtraInt(extra["codex_"+window+"_reset_after_seconds"]) + if resetAfter <= 0 { + return false + } + base := now + if updatedRaw, ok := extra["codex_usage_updated_at"]; ok { + if updatedAt, err := parseTime(fmt.Sprint(updatedRaw)); err == nil { + base = updatedAt + } + } + resetAt := base.Add(time.Duration(resetAfter) * time.Second) + return !now.Before(resetAt) } func readOpenAIQuotaUsedPercent(extra map[string]any, window string) float64 { diff --git a/backend/internal/service/openai_ws_account_sticky_test.go b/backend/internal/service/openai_ws_account_sticky_test.go index c8b28a46..6fc44298 100644 --- a/backend/internal/service/openai_ws_account_sticky_test.go +++ b/backend/internal/service/openai_ws_account_sticky_test.go @@ -48,6 +48,46 @@ func TestOpenAIGatewayService_SelectAccountByPreviousResponseID_Hit(t *testing.T } } +func TestOpenAIGatewayService_SelectAccountByPreviousResponseID_QuotaAutoPausedMiss(t *testing.T) { + ctx := context.Background() + groupID := int64(23) + account := Account{ + ID: 77, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 2, + Extra: map[string]any{ + "openai_apikey_responses_websockets_v2_enabled": true, + "codex_5h_used_percent": 96.0, + "auto_pause_5h_threshold": 0.95, + }, + } + cache := &stubGatewayCache{} + store := NewOpenAIWSStateStore(cache) + cfg := newOpenAIWSV2TestConfig() + svc := &OpenAIGatewayService{ + accountRepo: stubOpenAIAccountRepo{accounts: []Account{account}}, + cache: cache, + cfg: cfg, + concurrencyService: NewConcurrencyService(stubConcurrencyCache{}), + openaiWSStateStore: store, + } + + require.NoError(t, store.BindResponseAccount(ctx, groupID, "resp_prev_quota", account.ID, time.Hour)) + + selection, err := svc.SelectAccountByPreviousResponseID(ctx, &groupID, "resp_prev_quota", "gpt-5.1", nil, false) + require.NoError(t, err) + require.Nil(t, selection, "超过 5h 配额阈值的账号不应继续命中 previous_response_id 粘连") + + // Auto-pause is transient, so the binding is preserved: the chain can resume on the + // same account once the quota window resets. + boundAccountID, getErr := store.GetResponseAccount(ctx, groupID, "resp_prev_quota") + require.NoError(t, getErr) + require.Equal(t, account.ID, boundAccountID) +} + func TestOpenAIGatewayService_SelectAccountByPreviousResponseID_RateLimitedMiss(t *testing.T) { ctx := context.Background() groupID := int64(23) diff --git a/backend/internal/service/openai_ws_forwarder.go b/backend/internal/service/openai_ws_forwarder.go index 6eea0191..878ff486 100644 --- a/backend/internal/service/openai_ws_forwarder.go +++ b/backend/internal/service/openai_ws_forwarder.go @@ -4045,6 +4045,13 @@ func (s *OpenAIGatewayService) selectAccountByPreviousResponseIDForCapability( if !account.SupportsOpenAIEndpointCapability(requiredCapability) { return nil, nil } + // Quota auto-pause must also gate the previous_response_id sticky path; otherwise an + // account over its 5h/7d threshold keeps serving the same response chain even though + // normal scheduling skips it. Pause is transient, so fall through to normal scheduling + // without deleting the binding (the window may reset before the next turn). + if paused, _ := shouldAutoPauseOpenAIAccountByQuota(ctx, account); paused { + return nil, nil + } if s.schedulerSnapshot != nil && s.accountRepo != nil { latest, latestErr := s.accountRepo.GetByID(ctx, account.ID) if latestErr != nil || latest == nil { @@ -4061,6 +4068,9 @@ func (s *OpenAIGatewayService) selectAccountByPreviousResponseIDForCapability( if !latest.SupportsOpenAIEndpointCapability(requiredCapability) { return nil, nil } + if paused, _ := shouldAutoPauseOpenAIAccountByQuota(ctx, latest); paused { + return nil, nil + } if s.isOpenAIAccountRuntimeBlocked(latest) { _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID) return nil, nil diff --git a/frontend/src/api/admin/ops.ts b/frontend/src/api/admin/ops.ts index 69235668..847fc8c9 100644 --- a/frontend/src/api/admin/ops.ts +++ b/frontend/src/api/admin/ops.ts @@ -778,9 +778,15 @@ export interface OpsAlertRuntimeSettings { thresholds: OpsMetricThresholds // 指标阈值配置 } +export interface OpsOpenAIAccountQuotaAutoPauseSettings { + default_threshold_5h: number // 0~1,0 表示不启用全局默认 5h 阈值 + default_threshold_7d: number // 0~1,0 表示不启用全局默认 7d 阈值 +} + export interface OpsAdvancedSettings { data_retention: OpsDataRetentionSettings aggregation: OpsAggregationSettings + openai_account_quota_auto_pause: OpsOpenAIAccountQuotaAutoPauseSettings ignore_count_tokens_errors: boolean ignore_context_canceled: boolean ignore_no_available_accounts: boolean diff --git a/frontend/src/i18n/locales/en.ts b/frontend/src/i18n/locales/en.ts index fa2e5a92..8ab90961 100644 --- a/frontend/src/i18n/locales/en.ts +++ b/frontend/src/i18n/locales/en.ts @@ -5193,6 +5193,11 @@ export default { aggregation: 'Pre-aggregation Tasks', enableAggregation: 'Enable Pre-aggregation', aggregationHint: 'Pre-aggregation improves query performance for long time windows', + openaiQuotaAutoPause: 'OpenAI Account Quota Auto-pause', + openaiQuotaAutoPauseHint: 'When an OpenAI account reaches its 5h / 7d usage threshold, the scheduler skips it automatically and resumes once the window rolls over. Per-account thresholds take precedence over this global default.', + openaiQuotaAutoPauseDefault5h: 'Default 5h usage threshold (%)', + openaiQuotaAutoPauseDefault7d: 'Default 7d usage threshold (%)', + openaiQuotaAutoPauseThresholdHint: 'Value 0-100; leave blank or 0 to disable the global default threshold.', errorFiltering: 'Error Filtering', ignoreCountTokensErrors: 'Ignore count_tokens errors', ignoreCountTokensErrorsHint: 'When enabled, errors from count_tokens requests will not be written to the error log.', @@ -5223,7 +5228,8 @@ export default { slaMinPercentRange: 'SLA minimum percentage must be between 0 and 100', ttftP99MaxRange: 'TTFT P99 maximum must be a number ≥ 0', requestErrorRateMaxRange: 'Request error rate maximum must be between 0 and 100', - upstreamErrorRateMaxRange: 'Upstream error rate maximum must be between 0 and 100' + upstreamErrorRateMaxRange: 'Upstream error rate maximum must be between 0 and 100', + openaiQuotaAutoPauseRange: 'OpenAI quota auto-pause threshold must be between 0 and 100' } }, concurrency: { diff --git a/frontend/src/i18n/locales/zh.ts b/frontend/src/i18n/locales/zh.ts index 2364f9c4..4f1d1f13 100644 --- a/frontend/src/i18n/locales/zh.ts +++ b/frontend/src/i18n/locales/zh.ts @@ -5352,6 +5352,11 @@ export default { aggregation: '预聚合任务', enableAggregation: '启用预聚合任务', aggregationHint: '预聚合可提升长时间窗口查询性能', + openaiQuotaAutoPause: 'OpenAI 账号配额自动暂停', + openaiQuotaAutoPauseHint: '当 OpenAI 账号 5h / 7d 用量达到阈值时,调度会自动跳过该账号;窗口滚动后自动恢复。账号级阈值优先于此全局默认值。', + openaiQuotaAutoPauseDefault5h: '默认 5h 用量阈值 (%)', + openaiQuotaAutoPauseDefault7d: '默认 7d 用量阈值 (%)', + openaiQuotaAutoPauseThresholdHint: '取值 0-100,留空或 0 表示不启用全局默认阈值。', errorFiltering: '错误过滤', ignoreCountTokensErrors: '忽略 count_tokens 错误', ignoreCountTokensErrorsHint: '启用后,count_tokens 请求的错误将不会写入错误日志。', @@ -5383,7 +5388,8 @@ export default { slaMinPercentRange: 'SLA最低百分比必须在0-100之间', ttftP99MaxRange: 'TTFT P99最大值必须大于等于0', requestErrorRateMaxRange: '请求错误率最大值必须在0-100之间', - upstreamErrorRateMaxRange: '上游错误率最大值必须在0-100之间' + upstreamErrorRateMaxRange: '上游错误率最大值必须在0-100之间', + openaiQuotaAutoPauseRange: 'OpenAI 配额自动暂停阈值必须在 0-100 之间' } }, concurrency: { diff --git a/frontend/src/views/admin/ops/components/OpsSettingsDialog.vue b/frontend/src/views/admin/ops/components/OpsSettingsDialog.vue index 5dba5b1d..bfb7a65f 100644 --- a/frontend/src/views/admin/ops/components/OpsSettingsDialog.vue +++ b/frontend/src/views/admin/ops/components/OpsSettingsDialog.vue @@ -50,6 +50,10 @@ async function loadAllSettings() { runtimeSettings.value = runtime emailConfig.value = email advancedSettings.value = advanced + // 兼容旧 payload:后端未返回该字段时补默认值,保证表单可绑定 + if (advancedSettings.value && !advancedSettings.value.openai_account_quota_auto_pause) { + advancedSettings.value.openai_account_quota_auto_pause = { default_threshold_5h: 0, default_threshold_7d: 0 } + } // 如果后端返回了阈值,使用后端的值;否则保持默认值 if (thresholds && Object.keys(thresholds).length > 0) { metricThresholds.value = { @@ -119,6 +123,28 @@ function removeRecipient(target: 'alert' | 'report', email: string) { if (idx >= 0) list.splice(idx, 1) } +// OpenAI 账号配额自动暂停:后端按 0~1 分数存储,UI 按百分比(0~100)展示 +const quotaAutoPause5hPercent = computed({ + get() { + const v = advancedSettings.value?.openai_account_quota_auto_pause?.default_threshold_5h + return v && v > 0 ? Math.round(v * 1000) / 10 : null + }, + set(val) { + if (!advancedSettings.value?.openai_account_quota_auto_pause) return + advancedSettings.value.openai_account_quota_auto_pause.default_threshold_5h = val != null && val > 0 ? val / 100 : 0 + } +}) +const quotaAutoPause7dPercent = computed({ + get() { + const v = advancedSettings.value?.openai_account_quota_auto_pause?.default_threshold_7d + return v && v > 0 ? Math.round(v * 1000) / 10 : null + }, + set(val) { + if (!advancedSettings.value?.openai_account_quota_auto_pause) return + advancedSettings.value.openai_account_quota_auto_pause.default_threshold_7d = val != null && val > 0 ? val / 100 : 0 + } +}) + // 验证 const validation = computed(() => { const errors: string[] = [] @@ -145,6 +171,11 @@ const validation = computed(() => { if (hourly_metrics_retention_days < 0 || hourly_metrics_retention_days > 365) { errors.push(t('admin.ops.settings.validation.retentionDaysRange')) } + + const { default_threshold_5h, default_threshold_7d } = advancedSettings.value.openai_account_quota_auto_pause + if (default_threshold_5h < 0 || default_threshold_5h > 1 || default_threshold_7d < 0 || default_threshold_7d > 1) { + errors.push(t('admin.ops.settings.validation.openaiQuotaAutoPauseRange')) + } } // 验证指标阈值 @@ -473,6 +504,40 @@ async function saveAllSettings() {
+ +
+
{{ t('admin.ops.settings.openaiQuotaAutoPause') }}
+

{{ t('admin.ops.settings.openaiQuotaAutoPauseHint') }}

+ +
+
+ + +
+
+ + +
+
+

{{ t('admin.ops.settings.openaiQuotaAutoPauseThresholdHint') }}

+
+
{{ t('admin.ops.settings.errorFiltering') }}
From eba2046320ec643e7976296aabbc283f05aeaa52 Mon Sep 17 00:00:00 2001 From: xiaoyiluck666 <83876597+xiaoyiluck666@users.noreply.github.com> Date: Fri, 29 May 2026 13:34:10 +0800 Subject: [PATCH 36/79] fix: enrich OpenAI OAuth token refresh --- backend/cmd/server/wire_gen.go | 2 +- .../internal/service/openai_oauth_service.go | 68 ++++++++++++------- .../openai_oauth_service_refresh_test.go | 7 ++ .../service/openai_privacy_service.go | 58 ++++++++++++++++ .../service/openai_subscription_test.go | 42 ++++++++++++ backend/internal/service/wire.go | 13 +++- 6 files changed, 164 insertions(+), 26 deletions(-) create mode 100644 backend/internal/service/openai_subscription_test.go diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index 465f5e25..441bcd67 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -137,7 +137,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { httpUpstream := repository.NewHTTPUpstream(configConfig) deferredService := service.ProvideDeferredService(accountRepository, timingWheelService) openAIOAuthClient := repository.NewOpenAIOAuthClient() - openAIOAuthService := service.NewOpenAIOAuthService(proxyRepository, openAIOAuthClient) + openAIOAuthService := service.ProvideOpenAIOAuthService(proxyRepository, openAIOAuthClient, privacyClientFactory) oAuthRefreshAPI := service.ProvideOAuthRefreshAPI(accountRepository, geminiTokenCache) openAITokenProvider := service.ProvideOpenAITokenProvider(accountRepository, geminiTokenCache, openAIOAuthService, oAuthRefreshAPI) channelRepository := repository.NewChannelRepository(db) diff --git a/backend/internal/service/openai_oauth_service.go b/backend/internal/service/openai_oauth_service.go index dc094d43..0ee357a9 100644 --- a/backend/internal/service/openai_oauth_service.go +++ b/backend/internal/service/openai_oauth_service.go @@ -278,11 +278,29 @@ func (s *OpenAIOAuthService) enrichTokenInfo(ctx context.Context, tokenInfo *Ope tokenInfo.Email = info.Email } } + if strings.TrimSpace(tokenInfo.SubscriptionExpiresAt) == "" { + if expiresAt := fetchChatGPTSubscriptionExpiresAt(ctx, s.privacyClientFactory, tokenInfo.AccessToken, proxyURL, resolveChatGPTSubscriptionAccountID(tokenInfo, orgID)); expiresAt != "" { + tokenInfo.SubscriptionExpiresAt = expiresAt + } + } // 尝试设置隐私(关闭训练数据共享),best-effort tokenInfo.PrivacyMode = disableOpenAITraining(ctx, s.privacyClientFactory, tokenInfo.AccessToken, proxyURL) } +func resolveChatGPTSubscriptionAccountID(tokenInfo *OpenAITokenInfo, orgID string) string { + for _, candidate := range []string{ + tokenInfo.ChatGPTAccountID, + tokenInfo.OrganizationID, + orgID, + } { + if trimmed := strings.TrimSpace(candidate); trimmed != "" { + return trimmed + } + } + return "" +} + // RefreshAccountToken refreshes token for an OpenAI OAuth account func (s *OpenAIOAuthService) RefreshAccountToken(ctx context.Context, account *Account) (*OpenAITokenInfo, error) { if account.Platform != PlatformOpenAI { @@ -292,30 +310,6 @@ func (s *OpenAIOAuthService) RefreshAccountToken(ctx context.Context, account *A return nil, infraerrors.New(http.StatusBadRequest, "OPENAI_OAUTH_INVALID_ACCOUNT_TYPE", "account is not an OAuth account") } - refreshToken := account.GetCredential("refresh_token") - if refreshToken == "" { - accessToken := account.GetCredential("access_token") - if accessToken != "" { - tokenInfo := &OpenAITokenInfo{ - AccessToken: accessToken, - RefreshToken: "", - IDToken: account.GetCredential("id_token"), - ClientID: account.GetCredential("client_id"), - Email: account.GetCredential("email"), - ChatGPTAccountID: account.GetCredential("chatgpt_account_id"), - ChatGPTUserID: account.GetCredential("chatgpt_user_id"), - OrganizationID: account.GetCredential("organization_id"), - PlanType: account.GetCredential("plan_type"), - } - if expiresAt := account.GetCredentialAsTime("expires_at"); expiresAt != nil { - tokenInfo.ExpiresAt = expiresAt.Unix() - tokenInfo.ExpiresIn = int64(time.Until(*expiresAt).Seconds()) - } - return tokenInfo, nil - } - return nil, infraerrors.New(http.StatusBadRequest, "OPENAI_OAUTH_NO_REFRESH_TOKEN", "no refresh token available") - } - var proxyURL string if account.ProxyID != nil { proxy, err := s.proxyRepo.GetByID(ctx, *account.ProxyID) @@ -324,6 +318,32 @@ func (s *OpenAIOAuthService) RefreshAccountToken(ctx context.Context, account *A } } + refreshToken := account.GetCredential("refresh_token") + if refreshToken == "" { + accessToken := account.GetCredential("access_token") + if accessToken != "" { + tokenInfo := &OpenAITokenInfo{ + AccessToken: accessToken, + RefreshToken: "", + IDToken: account.GetCredential("id_token"), + ClientID: account.GetCredential("client_id"), + Email: account.GetCredential("email"), + ChatGPTAccountID: account.GetCredential("chatgpt_account_id"), + ChatGPTUserID: account.GetCredential("chatgpt_user_id"), + OrganizationID: account.GetCredential("organization_id"), + PlanType: account.GetCredential("plan_type"), + SubscriptionExpiresAt: account.GetCredential("subscription_expires_at"), + } + if expiresAt := account.GetCredentialAsTime("expires_at"); expiresAt != nil { + tokenInfo.ExpiresAt = expiresAt.Unix() + tokenInfo.ExpiresIn = int64(time.Until(*expiresAt).Seconds()) + } + s.enrichTokenInfo(ctx, tokenInfo, proxyURL) + return tokenInfo, nil + } + return nil, infraerrors.New(http.StatusBadRequest, "OPENAI_OAUTH_NO_REFRESH_TOKEN", "no refresh token available") + } + clientID := account.GetCredential("client_id") return s.RefreshTokenWithClientID(ctx, refreshToken, proxyURL, clientID) } diff --git a/backend/internal/service/openai_oauth_service_refresh_test.go b/backend/internal/service/openai_oauth_service_refresh_test.go index 84b68ea6..75588c8d 100644 --- a/backend/internal/service/openai_oauth_service_refresh_test.go +++ b/backend/internal/service/openai_oauth_service_refresh_test.go @@ -8,6 +8,7 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/pkg/openai" + "github.com/imroc/req/v3" "github.com/stretchr/testify/require" ) @@ -32,6 +33,11 @@ func (s *openaiOAuthClientRefreshStub) RefreshTokenWithClientID(ctx context.Cont func TestOpenAIOAuthService_RefreshAccountToken_NoRefreshTokenUsesExistingAccessToken(t *testing.T) { client := &openaiOAuthClientRefreshStub{} svc := NewOpenAIOAuthService(nil, client) + var privacyClientCalls int32 + svc.SetPrivacyClientFactory(func(proxyURL string) (*req.Client, error) { + atomic.AddInt32(&privacyClientCalls, 1) + return nil, errors.New("stop before request") + }) expiresAt := time.Now().Add(30 * time.Minute).UTC().Format(time.RFC3339) account := &Account{ @@ -51,6 +57,7 @@ func TestOpenAIOAuthService_RefreshAccountToken_NoRefreshTokenUsesExistingAccess require.Equal(t, "existing-access-token", info.AccessToken) require.Equal(t, "client-id-1", info.ClientID) require.Zero(t, atomic.LoadInt32(&client.refreshCalls), "existing access token should be reused without calling refresh") + require.Positive(t, atomic.LoadInt32(&privacyClientCalls), "existing access token should still run enrichment") } func TestOpenAITokenRefresher_NeedsRefresh_SkipsAccountWithoutRefreshToken(t *testing.T) { diff --git a/backend/internal/service/openai_privacy_service.go b/backend/internal/service/openai_privacy_service.go index da6dbefc..99cbb726 100644 --- a/backend/internal/service/openai_privacy_service.go +++ b/backend/internal/service/openai_privacy_service.go @@ -95,6 +95,8 @@ type ChatGPTAccountInfo struct { const chatGPTAccountsCheckURL = "https://chatgpt.com/backend-api/accounts/check/v4-2023-04-27" +var chatGPTSubscriptionsURL = "https://chatgpt.com/backend-api/subscriptions" + // fetchChatGPTAccountInfo calls ChatGPT backend-api to get account info (plan_type, etc.). // Used as fallback when id_token doesn't contain these fields (e.g., Mobile RT). // orgID is used to match the correct account when multiple accounts exist (e.g., personal + team). @@ -199,6 +201,62 @@ func fetchChatGPTAccountInfo(ctx context.Context, clientFactory PrivacyClientFac return info } +// fetchChatGPTSubscriptionExpiresAt reads the lightweight subscription endpoint used by +// ChatGPT/Codex clients. Some Plus accounts no longer expose entitlement.expires_at in +// accounts/check, but this endpoint still returns active_until. +func fetchChatGPTSubscriptionExpiresAt(ctx context.Context, clientFactory PrivacyClientFactory, accessToken, proxyURL, accountID string) string { + accountID = strings.TrimSpace(accountID) + if accessToken == "" || accountID == "" || clientFactory == nil { + return "" + } + + ctx, cancel := context.WithTimeout(ctx, 15*time.Second) + defer cancel() + + client, err := clientFactory(proxyURL) + if err != nil { + slog.Debug("chatgpt_subscription_client_error", "error", err.Error()) + return "" + } + + var result struct { + PlanType string `json:"plan_type"` + ActiveUntil string `json:"active_until"` + WillRenew bool `json:"will_renew"` + ID string `json:"id"` + } + resp, err := client.R(). + SetContext(ctx). + SetHeader("Authorization", "Bearer "+accessToken). + SetHeader("Origin", "https://chatgpt.com"). + SetHeader("Referer", "https://chatgpt.com/"). + SetHeader("Accept", "application/json"). + SetSuccessResult(&result). + SetQueryParam("account_id", accountID). + Get(chatGPTSubscriptionsURL) + if err != nil { + slog.Debug("chatgpt_subscription_request_error", "error", err.Error()) + return "" + } + if !resp.IsSuccessState() { + slog.Debug("chatgpt_subscription_failed", "status", resp.StatusCode, "body", truncate(resp.String(), 200)) + return "" + } + + activeUntil := strings.TrimSpace(result.ActiveUntil) + if activeUntil == "" { + slog.Debug("chatgpt_subscription_no_active_until", "plan_type", result.PlanType, "has_subscription_id", strings.TrimSpace(result.ID) != "", "will_renew", result.WillRenew) + return "" + } + if _, err := time.Parse(time.RFC3339, activeUntil); err != nil { + slog.Debug("chatgpt_subscription_bad_active_until", "active_until", activeUntil, "error", err.Error()) + return "" + } + + slog.Info("chatgpt_subscription_success", "plan_type", result.PlanType, "subscription_expires_at", activeUntil, "account_id", accountID) + return activeUntil +} + // fillAccountInfo 从单个 account 对象中提取 plan_type 和 subscription_expires_at func fillAccountInfo(info *ChatGPTAccountInfo, acct map[string]any) { info.PlanType = extractPlanType(acct) diff --git a/backend/internal/service/openai_subscription_test.go b/backend/internal/service/openai_subscription_test.go new file mode 100644 index 00000000..89df54db --- /dev/null +++ b/backend/internal/service/openai_subscription_test.go @@ -0,0 +1,42 @@ +package service + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/imroc/req/v3" + "github.com/stretchr/testify/require" +) + +func TestFetchChatGPTSubscriptionExpiresAt(t *testing.T) { + const wantExpiresAt = "2026-06-10T02:52:15Z" + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + require.Equal(t, "/backend-api/subscriptions", r.URL.Path) + require.Equal(t, "acc_123", r.URL.Query().Get("account_id")) + require.Equal(t, "Bearer access-token", r.Header.Get("Authorization")) + + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{ + "plan_type": "plus", + "active_until": wantExpiresAt, + "will_renew": true, + "id": "sub_123", + }) + })) + defer server.Close() + + oldURL := chatGPTSubscriptionsURL + chatGPTSubscriptionsURL = server.URL + "/backend-api/subscriptions" + t.Cleanup(func() { chatGPTSubscriptionsURL = oldURL }) + + got := fetchChatGPTSubscriptionExpiresAt(context.Background(), func(proxyURL string) (*req.Client, error) { + return req.C().SetTimeout(5 * time.Second), nil + }, "access-token", "", "acc_123") + + require.Equal(t, wantExpiresAt, got) +} diff --git a/backend/internal/service/wire.go b/backend/internal/service/wire.go index b22e10ae..e0c9f591 100644 --- a/backend/internal/service/wire.go +++ b/backend/internal/service/wire.go @@ -45,6 +45,17 @@ func ProvideOAuthRefreshAPI(accountRepo AccountRepository, tokenCache GeminiToke return NewOAuthRefreshAPI(accountRepo, tokenCache) } +// ProvideOpenAIOAuthService creates OpenAIOAuthService with privacy/account enrichment support. +func ProvideOpenAIOAuthService( + proxyRepo ProxyRepository, + oauthClient OpenAIOAuthClient, + privacyClientFactory PrivacyClientFactory, +) *OpenAIOAuthService { + svc := NewOpenAIOAuthService(proxyRepo, oauthClient) + svc.SetPrivacyClientFactory(privacyClientFactory) + return svc +} + // ProvideTokenRefreshService creates and starts TokenRefreshService func ProvideTokenRefreshService( accountRepo AccountRepository, @@ -461,7 +472,7 @@ var ProviderSet = wire.NewSet( NewOpenAIGatewayService, wire.Bind(new(AccountRuntimeBlocker), new(*OpenAIGatewayService)), NewOAuthService, - NewOpenAIOAuthService, + ProvideOpenAIOAuthService, NewGeminiOAuthService, NewGeminiQuotaService, NewCompositeTokenCacheInvalidator, From c9caadb3782a4e658ce46789df092d0b15d58172 Mon Sep 17 00:00:00 2001 From: wucm667 Date: Fri, 29 May 2026 14:32:45 +0800 Subject: [PATCH 37/79] fix(account): address second-round review on quota auto-pause - TopK initial filter now drops quota-paused accounts: fold the quota check into isAccountRequestCompatible so session-hash, TopK pool, and per-candidate rechecks all skip paused accounts. Previously the candidate pool was built without the quota check, so paused accounts could fill TopK and leave the scheduler returning "no available accounts" even with healthy ones available. - Add per-account explicit disable flags auto_pause_5h_disabled / auto_pause_7d_disabled with toggles in EditAccountModal. Without these, leaving the account threshold blank silently falls back to the global default, so admins could not exempt a single account once a global default existed. Disable is per-window: an account can opt out of 5h auto-pause while still honoring 7d. Schedule snapshot whitelist includes the new fields, i18n EN/ZH updated, threshold-hint text revised to explain "blank = global default". - Move quota auto-pause settings off the request hot path: replace the per-repo TTL+singleflight sync DB read with a per-SettingService stale-while-revalidate in-memory snapshot. Get is non-blocking (atomic.Pointer load + async refresh on staleness); writes via UpdateOpsAdvancedSettings push directly into the cache through an injected sink; wire warms the cache at startup. Adds Warm (sync) for tests/init and SetOpenAIQuotaAutoPauseSettings (sink target). Co-Authored-By: Claude Opus 4.7 --- backend/cmd/server/wire_gen.go | 2 +- .../internal/repository/scheduler_cache.go | 2 + .../repository/scheduler_cache_unit_test.go | 4 + .../service/openai_account_scheduler.go | 7 + .../service/openai_account_scheduler_test.go | 136 ++++++++++++++ .../service/openai_gateway_service.go | 44 ++++- backend/internal/service/ops_service.go | 15 ++ backend/internal/service/ops_settings.go | 12 +- .../service/ops_settings_advanced_test.go | 50 ++++- backend/internal/service/setting_service.go | 176 +++++++++++------- backend/internal/service/wire.go | 42 ++++- .../components/account/EditAccountModal.vue | 60 ++++++ .../__tests__/EditAccountModal.spec.ts | 21 +++ frontend/src/i18n/locales/en.ts | 5 +- frontend/src/i18n/locales/zh.ts | 5 +- 15 files changed, 505 insertions(+), 76 deletions(-) diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index 465f5e25..6e8be8fc 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -195,7 +195,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { 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.NewOpsService(opsRepository, settingRepository, configConfig, accountRepository, userRepository, concurrencyService, gatewayService, openAIGatewayService, geminiMessagesCompatService, antigravityGatewayService, opsSystemLogSink) + 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/repository/scheduler_cache.go b/backend/internal/repository/scheduler_cache.go index ff3c4301..cf19deda 100644 --- a/backend/internal/repository/scheduler_cache.go +++ b/backend/internal/repository/scheduler_cache.go @@ -557,6 +557,8 @@ func filterSchedulerExtra(extra map[string]any) map[string]any { "codex_usage_updated_at", "auto_pause_5h_threshold", "auto_pause_7d_threshold", + "auto_pause_5h_disabled", + "auto_pause_7d_disabled", } 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 9e4ec23e..a4667591 100644 --- a/backend/internal/repository/scheduler_cache_unit_test.go +++ b/backend/internal/repository/scheduler_cache_unit_test.go @@ -89,6 +89,8 @@ func TestBuildSchedulerMetadataAccount_KeepsQuotaAutoPauseFields(t *testing.T) { "codex_usage_updated_at": "2026-05-29T09:00:00Z", "auto_pause_5h_threshold": 0.95, "auto_pause_7d_threshold": 0.96, + "auto_pause_5h_disabled": true, + "auto_pause_7d_disabled": false, }, } @@ -103,4 +105,6 @@ func TestBuildSchedulerMetadataAccount_KeepsQuotaAutoPauseFields(t *testing.T) { require.Equal(t, "2026-05-29T09:00:00Z", got.Extra["codex_usage_updated_at"]) require.Equal(t, 0.95, got.Extra["auto_pause_5h_threshold"]) require.Equal(t, 0.96, got.Extra["auto_pause_7d_threshold"]) + require.Equal(t, true, got.Extra["auto_pause_5h_disabled"]) + require.Equal(t, false, got.Extra["auto_pause_7d_disabled"]) } diff --git a/backend/internal/service/openai_account_scheduler.go b/backend/internal/service/openai_account_scheduler.go index fd28fa86..47a8142a 100644 --- a/backend/internal/service/openai_account_scheduler.go +++ b/backend/internal/service/openai_account_scheduler.go @@ -974,6 +974,13 @@ func (s *defaultOpenAIAccountScheduler) isAccountRequestCompatible(ctx context.C if s != nil && s.service != nil && s.service.isOpenAIAccountRuntimeBlocked(account) { return false } + // Quota auto-pause must be evaluated during the initial filter too. Without it the + // TopK candidate pool can be filled with paused accounts and the later fresh/DB + // rechecks won't reach healthy accounts that fell outside TopK — manifesting as + // "no available accounts" even though healthy ones exist. + if paused, _ := shouldAutoPauseOpenAIAccountByQuota(ctx, account); paused { + return false + } if req.RequestedModel != "" && !account.IsModelSupported(req.RequestedModel) { return false } diff --git a/backend/internal/service/openai_account_scheduler_test.go b/backend/internal/service/openai_account_scheduler_test.go index 531769a7..da5f0a66 100644 --- a/backend/internal/service/openai_account_scheduler_test.go +++ b/backend/internal/service/openai_account_scheduler_test.go @@ -798,6 +798,63 @@ func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_UsesGlobalDefa require.Equal(t, int64(35402), account.ID) } +// Regression: a per-account explicit-disable flag exempts the account from auto-pause +// even when a global default threshold is set. Without this, "leave threshold blank" +// silently falls back to global default and admins have no way to whitelist a single +// account. +func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_PerAccountDisableOverridesGlobalDefault(t *testing.T) { + ctx := withOpenAIQuotaAutoPauseSettings(context.Background(), OpsOpenAIAccountQuotaAutoPauseSettings{DefaultThreshold5h: 0.95}) + // Account has high usage AND no per-account threshold (would normally fall back to + // the global default and get paused), but the explicit disable flag is set. + 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_disabled": true, + }, + } + 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) +} + +// Disable is per-window: disabling only 5h must still allow 7d auto-pause to fire. +func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_PerWindowDisableScoped(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, + "codex_7d_used_percent": 99.0, + "auto_pause_5h_disabled": true, + "auto_pause_7d_threshold": 0.95, + }, + } + 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, "7d auto-pause must still fire even though 5h is disabled") +} + func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_StaleUsageWindowResetSkipsPause(t *testing.T) { ctx := context.Background() // Usage is over threshold but the window's reset time has already passed, so the @@ -1399,6 +1456,85 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_LoadBalanceTopKFallback } } +// Regression: TopK initial filter must drop quota-auto-paused accounts. Otherwise +// the candidate pool is filled with paused accounts, healthy accounts fall outside +// TopK, and the scheduler returns "no available accounts" even though healthy ones +// exist. +func TestOpenAIGatewayService_SelectAccountWithScheduler_LoadBalanceTopKExcludesQuotaPaused(t *testing.T) { + ctx := context.Background() + groupID := int64(110) + accounts := []Account{ + { + ID: 37001, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 0, + Extra: map[string]any{ + "codex_5h_used_percent": 96.0, + "auto_pause_5h_threshold": 0.95, + }, + }, + { + ID: 37002, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Priority: 5, + }, + } + + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.LBTopK = 1 // TopK=1 makes the bug fatal: paused account would crowd out the healthy one entirely + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Priority = 0.4 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Load = 1.0 + cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Queue = 1.0 + + concurrencyCache := schedulerTestConcurrencyCache{ + loadMap: map[int64]*AccountLoadInfo{ + 37001: {AccountID: 37001, LoadRate: 5, WaitingCount: 0}, + 37002: {AccountID: 37002, LoadRate: 5, WaitingCount: 0}, + }, + acquireResults: map[int64]bool{ + 37002: true, + }, + } + + svc := &OpenAIGatewayService{ + accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts}, + cache: &schedulerTestGatewayCache{}, + cfg: cfg, + rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"), + concurrencyService: NewConcurrencyService(concurrencyCache), + } + + selection, decision, err := svc.SelectAccountWithScheduler( + ctx, + &groupID, + "", + "", + "gpt-5.1", + nil, + OpenAIUpstreamTransportAny, + false, + ) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, int64(37002), selection.Account.ID) + require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) + // Only the healthy account should ever enter the candidate pool; the paused one + // must be filtered out at the initial-filter stage. + require.Equal(t, 1, decision.CandidateCount) + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } +} + func TestOpenAIGatewayService_OpenAIAccountSchedulerMetrics(t *testing.T) { ctx := context.Background() groupID := int64(12) diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index e2534cc2..b1ae6a9a 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -1360,14 +1360,21 @@ func shouldAutoPauseOpenAIAccountByQuota(ctx context.Context, account *Account) if account == nil || !account.IsOpenAI() { return false, openAIQuotaAutoPauseDecision{} } + // Per-account explicit-disable flags must take precedence over the global default. + // Without these, leaving the account threshold blank means "use global default", + // so an admin has no way to exempt a single account from auto-pause once a global + // default exists. The disable flag is per-window so an account can opt out of + // only 5h or only 7d auto-pause. + disabled5h := resolveAccountExtraBool(account.Extra, "auto_pause_5h_disabled") + disabled7d := resolveAccountExtraBool(account.Extra, "auto_pause_7d_disabled") threshold5h, threshold7d := resolveOpenAIQuotaAutoPauseThresholds(ctx, account) now := time.Now() - if threshold5h > 0 { + if !disabled5h && threshold5h > 0 { if utilization, ok := resolveOpenAIQuotaUtilization(account.Extra, "5h", now); ok && utilization >= threshold5h { return true, openAIQuotaAutoPauseDecision{window: "5h", threshold: threshold5h, utilization: utilization} } } - if threshold7d > 0 { + if !disabled7d && threshold7d > 0 { if utilization, ok := resolveOpenAIQuotaUtilization(account.Extra, "7d", now); ok && utilization >= threshold7d { return true, openAIQuotaAutoPauseDecision{window: "7d", threshold: threshold7d, utilization: utilization} } @@ -1375,6 +1382,39 @@ func shouldAutoPauseOpenAIAccountByQuota(ctx context.Context, account *Account) return false, openAIQuotaAutoPauseDecision{} } +// resolveAccountExtraBool reads a bool-like value from account extra, tolerating +// the few shapes JSON unmarshalling may produce (real bool, "true"/"false" +// strings, 0/1 numbers). +func resolveAccountExtraBool(extra map[string]any, key string) bool { + if len(extra) == 0 { + return false + } + value, ok := extra[key] + if !ok || value == nil { + return false + } + switch v := value.(type) { + case bool: + return v + case string: + parsed, err := strconv.ParseBool(strings.TrimSpace(v)) + return err == nil && parsed + case float64: + return v != 0 + case float32: + return v != 0 + case int: + return v != 0 + case int64: + return v != 0 + case json.Number: + if i, err := v.Int64(); err == nil { + return i != 0 + } + } + return false +} + func resolveOpenAIQuotaAutoPauseThresholds(ctx context.Context, account *Account) (float64, float64) { threshold5h, _ := resolveAccountExtraNumber(account.Extra, "auto_pause_5h_threshold") threshold7d, _ := resolveAccountExtraNumber(account.Extra, "auto_pause_7d_threshold") diff --git a/backend/internal/service/ops_service.go b/backend/internal/service/ops_service.go index 1cea72fa..2d7c5bd4 100644 --- a/backend/internal/service/ops_service.go +++ b/backend/internal/service/ops_service.go @@ -41,6 +41,11 @@ type OpsService struct { // cleanupReloader 由 wire 在 OpsCleanupService 构造完成后通过 SetCleanupReloader 注入。 // 解耦避免 OpsService -> OpsCleanupService 的硬依赖(cleanup 也读 settings,会循环)。 cleanupReloader CleanupReloader + + // quotaAutoPauseSink 由 wire 注入(通常是 SettingService.SetOpenAIQuotaAutoPauseSettings)。 + // UpdateOpsAdvancedSettings 写入新配置后调用,把最新的 quota auto-pause 全局默认阈值 + // 立即同步到调度热路径读取的内存缓存,避免下次请求才能感知新值。 + quotaAutoPauseSink func(OpsOpenAIAccountQuotaAutoPauseSettings) } // CleanupReloader 由 OpsCleanupService 实现。 @@ -57,6 +62,16 @@ func (s *OpsService) SetCleanupReloader(r CleanupReloader) { s.cleanupReloader = r } +// SetOpenAIQuotaAutoPauseSettingsSink 由 wire 注入,把最新的 quota auto-pause 全局默认 +// 阈值 push 到调度热路径读取的内存缓存。同 SetCleanupReloader 的解耦目的:避免 OpsService +// 持有 *SettingService 引入循环依赖。 +func (s *OpsService) SetOpenAIQuotaAutoPauseSettingsSink(sink func(OpsOpenAIAccountQuotaAutoPauseSettings)) { + if s == nil { + return + } + s.quotaAutoPauseSink = sink +} + func NewOpsService( opsRepo OpsRepository, settingRepo SettingRepository, diff --git a/backend/internal/service/ops_settings.go b/backend/internal/service/ops_settings.go index 23e92e5a..472f4e32 100644 --- a/backend/internal/service/ops_settings.go +++ b/backend/internal/service/ops_settings.go @@ -490,12 +490,12 @@ func (s *OpsService) UpdateOpsAdvancedSettings(ctx context.Context, cfg *OpsAdva if err := s.settingRepo.Set(ctx, SettingKeyOpsAdvancedSettings, string(raw)); err != nil { return nil, err } - cacheKey := openAIQuotaAutoPauseSettingsCacheKey(s.settingRepo) - openAIQuotaAutoPauseSettingsSF.Forget(cacheKey) - storeOpenAIQuotaAutoPauseSettingsCache(s.settingRepo, &cachedOpenAIQuotaAutoPauseSettings{ - settings: cfg.OpenAIAccountQuotaAutoPause, - expiresAt: time.Now().Add(openAIQuotaAutoPauseSettingsCacheTTL).UnixNano(), - }) + // Push the new quota auto-pause settings straight into the in-memory cache that + // the OpenAI scheduling hot path reads, so the next request observes the new value + // without waiting for the background refresher's TTL. + if s.quotaAutoPauseSink != nil { + s.quotaAutoPauseSink(cfg.OpenAIAccountQuotaAutoPause) + } // notify cleanup service to reload schedule/enabled. if s.cleanupReloader != nil { diff --git a/backend/internal/service/ops_settings_advanced_test.go b/backend/internal/service/ops_settings_advanced_test.go index d8598fe0..62803f94 100644 --- a/backend/internal/service/ops_settings_advanced_test.go +++ b/backend/internal/service/ops_settings_advanced_test.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "testing" + "time" "github.com/Wei-Shaw/sub2api/internal/config" ) @@ -103,11 +104,58 @@ func TestGetOpenAIQuotaAutoPauseSettings_ReadsDefaultsFromOpsAdvancedSettings(t repo.values[SettingKeyOpsAdvancedSettings] = `{"openai_account_quota_auto_pause":{"default_threshold_5h":0.95,"default_threshold_7d":0.9}}` svc := NewSettingService(repo, &config.Config{}) - settings := svc.GetOpenAIQuotaAutoPauseSettings(context.Background()) + // Warm the in-memory cache synchronously so the assertion below is deterministic. + // GetOpenAIQuotaAutoPauseSettings is non-blocking on the hot path (returns the + // cached value, refreshes asynchronously); for tests and startup, Warm is the + // synchronous entry point that guarantees a populated cache. + settings := svc.WarmOpenAIQuotaAutoPauseSettings(context.Background()) if settings.DefaultThreshold5h != 0.95 { t.Fatalf("DefaultThreshold5h = %v, want 0.95", settings.DefaultThreshold5h) } if settings.DefaultThreshold7d != 0.9 { t.Fatalf("DefaultThreshold7d = %v, want 0.9", settings.DefaultThreshold7d) } + + // Subsequent Get must hit the warm cache and return the same value without any DB + // access — that's the hot-path invariant. + cached := svc.GetOpenAIQuotaAutoPauseSettings(context.Background()) + if cached.DefaultThreshold5h != 0.95 || cached.DefaultThreshold7d != 0.9 { + t.Fatalf("cached read = %+v, want {0.95, 0.9}", cached) + } +} + +// Hot-path invariant: a Get with cold cache must return immediately (zero defaults) +// rather than blocking on the DB. The async refresher will populate the cache for +// subsequent calls. +func TestGetOpenAIQuotaAutoPauseSettings_ColdCacheNonBlocking(t *testing.T) { + repo := newRuntimeSettingRepoStub() + repo.values[SettingKeyOpsAdvancedSettings] = `{"openai_account_quota_auto_pause":{"default_threshold_5h":0.7}}` + svc := NewSettingService(repo, &config.Config{}) + + start := time.Now() + settings := svc.GetOpenAIQuotaAutoPauseSettings(context.Background()) + elapsed := time.Since(start) + if elapsed > 50*time.Millisecond { + t.Fatalf("cold-cache Get must be non-blocking, took %v", elapsed) + } + // Cold cache means we get zero defaults (the async refresh hasn't completed yet). + if settings.DefaultThreshold5h != 0 || settings.DefaultThreshold7d != 0 { + t.Fatalf("cold-cache Get = %+v, want zeroes", settings) + } +} + +// Explicit cache write (e.g. from UpdateOpsAdvancedSettings) must be visible on the +// very next read without any DB roundtrip. +func TestSetOpenAIQuotaAutoPauseSettings_VisibleImmediately(t *testing.T) { + svc := NewSettingService(newRuntimeSettingRepoStub(), &config.Config{}) + + svc.SetOpenAIQuotaAutoPauseSettings(OpsOpenAIAccountQuotaAutoPauseSettings{ + DefaultThreshold5h: 0.88, + DefaultThreshold7d: 0.77, + }) + + got := svc.GetOpenAIQuotaAutoPauseSettings(context.Background()) + if got.DefaultThreshold5h != 0.88 || got.DefaultThreshold7d != 0.77 { + t.Fatalf("after Set, Get = %+v, want {0.88, 0.77}", got) + } } diff --git a/backend/internal/service/setting_service.go b/backend/internal/service/setting_service.go index 40e98b88..98acdb80 100644 --- a/backend/internal/service/setting_service.go +++ b/backend/internal/service/setting_service.go @@ -14,7 +14,6 @@ import ( "sort" "strconv" "strings" - "sync" "sync/atomic" "time" @@ -162,28 +161,7 @@ const openAIQuotaAutoPauseSettingsCacheTTL = 60 * time.Second const openAIQuotaAutoPauseSettingsErrorTTL = 5 * time.Second const openAIQuotaAutoPauseSettingsDBTimeout = 5 * time.Second -var openAIQuotaAutoPauseSettingsCache sync.Map // map[string]*cachedOpenAIQuotaAutoPauseSettings -var openAIQuotaAutoPauseSettingsSF singleflight.Group - -func openAIQuotaAutoPauseSettingsCacheKey(repo SettingRepository) string { - if repo == nil { - return "nil" - } - return fmt.Sprintf("%T:%p", repo, repo) -} - -func loadOpenAIQuotaAutoPauseSettingsCache(repo SettingRepository) (*cachedOpenAIQuotaAutoPauseSettings, bool) { - value, ok := openAIQuotaAutoPauseSettingsCache.Load(openAIQuotaAutoPauseSettingsCacheKey(repo)) - if !ok || value == nil { - return nil, false - } - cached, ok := value.(*cachedOpenAIQuotaAutoPauseSettings) - return cached, ok && cached != nil -} - -func storeOpenAIQuotaAutoPauseSettingsCache(repo SettingRepository, cached *cachedOpenAIQuotaAutoPauseSettings) { - openAIQuotaAutoPauseSettingsCache.Store(openAIQuotaAutoPauseSettingsCacheKey(repo), cached) -} +const openAIQuotaAutoPauseSettingsRefreshKey = "openai_quota_auto_pause_settings" // DefaultSubscriptionGroupReader validates group references used by default subscriptions. type DefaultSubscriptionGroupReader interface { @@ -209,6 +187,15 @@ type SettingService struct { openAICodexUASF singleflight.Group openAIAllowCodexPluginCache atomic.Value // *cachedOpenAIAllowCodexPlugin openAIAllowCodexPluginSF singleflight.Group + + // openAIQuotaAutoPauseSettingsCache holds the most recently observed quota auto-pause + // settings. GetOpenAIQuotaAutoPauseSettings reads this atomic.Value on the request hot + // path without ever blocking on the DB; when the cached entry expires, a background + // goroutine refreshes it via openAIQuotaAutoPauseSettingsSF (stale-while-revalidate). + // This per-service field also gives tests natural isolation — each SettingService + // instance owns its own cache, no shared package-level state. + openAIQuotaAutoPauseSettingsCache atomic.Value // *cachedOpenAIQuotaAutoPauseSettings + openAIQuotaAutoPauseSettingsSF singleflight.Group } // DefaultPlatformQuotaSetting 单 platform 三档限额(nil = 沿用上层;0 = 显式禁用;>0 = 上限) @@ -2060,9 +2047,17 @@ func (s *SettingService) refreshCachedSettings(settings *SystemSettings) { enabled: settings.OpenAIAdvancedSchedulerEnabled, expiresAt: time.Now().Add(openAIAdvancedSchedulerSettingCacheTTL).UnixNano(), }) - cacheKey := openAIQuotaAutoPauseSettingsCacheKey(s.settingRepo) - openAIQuotaAutoPauseSettingsSF.Forget(cacheKey) - openAIQuotaAutoPauseSettingsCache.Delete(cacheKey) + // Invalidate the quota auto-pause cache and let the next read trigger a fresh load. + // We can't know from here whether ops_advanced_settings was also touched, so be + // defensive: store an expired entry — GetOpenAIQuotaAutoPauseSettings will serve + // stale and kick off an async refresh, never blocking the request that follows. + s.openAIQuotaAutoPauseSettingsSF.Forget(openAIQuotaAutoPauseSettingsRefreshKey) + if cached, _ := s.openAIQuotaAutoPauseSettingsCache.Load().(*cachedOpenAIQuotaAutoPauseSettings); cached != nil { + s.openAIQuotaAutoPauseSettingsCache.Store(&cachedOpenAIQuotaAutoPauseSettings{ + settings: cached.settings, + expiresAt: 0, + }) + } if s.cfg != nil { s.cfg.SetTrustForwardedIPForAPIKeyACL(settings.APIKeyACLTrustForwardedIP) } @@ -4484,49 +4479,104 @@ func (s *SettingService) GetClaudeCodeVersionBounds(ctx context.Context) (min, m return b.min, b.max } +// GetOpenAIQuotaAutoPauseSettings returns the current global default quota auto-pause +// settings. It is invoked on the OpenAI scheduling hot path (once per request) and is +// therefore designed to never block on the DB: +// +// - Fresh cached value → returned immediately. +// - Stale or empty cache → the last known value is returned, and a background +// goroutine refreshes the cache via singleflight (stale-while-revalidate). +// - First call with no cache yet → zero defaults are returned and the same async +// refresh is kicked off; the next call gets the freshly populated value. +// +// Callers that need the freshly persisted value synchronously (tests, post-update +// confirmation, optional startup warm-up) should call WarmOpenAIQuotaAutoPauseSettings. func (s *SettingService) GetOpenAIQuotaAutoPauseSettings(ctx context.Context) OpsOpenAIAccountQuotaAutoPauseSettings { - if cached, ok := loadOpenAIQuotaAutoPauseSettingsCache(s.settingRepo); ok { - if time.Now().UnixNano() < cached.expiresAt { - return cached.settings + if s == nil { + return OpsOpenAIAccountQuotaAutoPauseSettings{} + } + cached, _ := s.openAIQuotaAutoPauseSettingsCache.Load().(*cachedOpenAIQuotaAutoPauseSettings) + now := time.Now().UnixNano() + if cached != nil && now < cached.expiresAt { + return cached.settings + } + // Stale or unset: trigger background refresh without blocking this request. + // singleflight.DoChan dedupes concurrent refreshes; we deliberately ignore the + // returned channel — the result is observable via the atomic cache. + s.openAIQuotaAutoPauseSettingsSF.DoChan(openAIQuotaAutoPauseSettingsRefreshKey, func() (any, error) { + s.refreshOpenAIQuotaAutoPauseSettings(context.Background()) + return nil, nil + }) + if cached != nil { + return cached.settings // serve stale value while revalidating + } + return OpsOpenAIAccountQuotaAutoPauseSettings{} +} + +// WarmOpenAIQuotaAutoPauseSettings synchronously loads the quota auto-pause settings +// into the in-memory cache. Useful for application startup (so the first request hits +// a warm cache) and for tests that need deterministic reads immediately after +// constructing the service. +func (s *SettingService) WarmOpenAIQuotaAutoPauseSettings(ctx context.Context) OpsOpenAIAccountQuotaAutoPauseSettings { + if s == nil { + return OpsOpenAIAccountQuotaAutoPauseSettings{} + } + s.refreshOpenAIQuotaAutoPauseSettings(ctx) + cached, _ := s.openAIQuotaAutoPauseSettingsCache.Load().(*cachedOpenAIQuotaAutoPauseSettings) + if cached == nil { + return OpsOpenAIAccountQuotaAutoPauseSettings{} + } + return cached.settings +} + +// refreshOpenAIQuotaAutoPauseSettings reads the latest settings from the DB and stores +// them into the in-memory cache. On error it stores the prior value (or zero defaults +// if nothing is cached yet) with the shorter error TTL so the next refresh comes +// sooner. Always uses its own timeout-bounded context to keep refresh latency +// predictable regardless of the caller. +func (s *SettingService) refreshOpenAIQuotaAutoPauseSettings(ctx context.Context) { + if s == nil || s.settingRepo == nil { + return + } + dbCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), openAIQuotaAutoPauseSettingsDBTimeout) + defer cancel() + + settings := OpsOpenAIAccountQuotaAutoPauseSettings{} + ttl := openAIQuotaAutoPauseSettingsCacheTTL + raw, err := s.settingRepo.GetValue(dbCtx, SettingKeyOpsAdvancedSettings) + if err == nil { + cfg := defaultOpsAdvancedSettings() + if strings.TrimSpace(raw) != "" { + if jsonErr := json.Unmarshal([]byte(raw), cfg); jsonErr == nil { + normalizeOpsAdvancedSettings(cfg) + } } + settings = cfg.OpenAIAccountQuotaAutoPause + } else if !errors.Is(err, ErrSettingNotFound) { + // Real error: keep serving prior value but refresh sooner. + if prior, _ := s.openAIQuotaAutoPauseSettingsCache.Load().(*cachedOpenAIQuotaAutoPauseSettings); prior != nil { + settings = prior.settings + } + ttl = openAIQuotaAutoPauseSettingsErrorTTL } - cacheKey := openAIQuotaAutoPauseSettingsCacheKey(s.settingRepo) - result, _, _ := openAIQuotaAutoPauseSettingsSF.Do(cacheKey, func() (any, error) { - if cached, ok := loadOpenAIQuotaAutoPauseSettingsCache(s.settingRepo); ok { - if time.Now().UnixNano() < cached.expiresAt { - return cached.settings, nil - } - } - - settings := OpsOpenAIAccountQuotaAutoPauseSettings{} - ttl := openAIQuotaAutoPauseSettingsCacheTTL - if s != nil && s.settingRepo != nil { - dbCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), openAIQuotaAutoPauseSettingsDBTimeout) - defer cancel() - raw, err := s.settingRepo.GetValue(dbCtx, SettingKeyOpsAdvancedSettings) - if err == nil { - cfg := defaultOpsAdvancedSettings() - if strings.TrimSpace(raw) != "" { - if jsonErr := json.Unmarshal([]byte(raw), cfg); jsonErr == nil { - normalizeOpsAdvancedSettings(cfg) - } - } - settings = cfg.OpenAIAccountQuotaAutoPause - } else { - ttl = openAIQuotaAutoPauseSettingsErrorTTL - } - } - - storeOpenAIQuotaAutoPauseSettingsCache(s.settingRepo, &cachedOpenAIQuotaAutoPauseSettings{ - settings: settings, - expiresAt: time.Now().Add(ttl).UnixNano(), - }) - return settings, nil + s.openAIQuotaAutoPauseSettingsCache.Store(&cachedOpenAIQuotaAutoPauseSettings{ + settings: settings, + expiresAt: time.Now().Add(ttl).UnixNano(), }) +} - settings, _ := result.(OpsOpenAIAccountQuotaAutoPauseSettings) - return settings +// SetOpenAIQuotaAutoPauseSettings writes the given settings directly into the in-memory +// cache. Called from settings-write code paths so that the next read reflects the new +// value immediately, without waiting for the background refresh. +func (s *SettingService) SetOpenAIQuotaAutoPauseSettings(settings OpsOpenAIAccountQuotaAutoPauseSettings) { + if s == nil { + return + } + s.openAIQuotaAutoPauseSettingsCache.Store(&cachedOpenAIQuotaAutoPauseSettings{ + settings: settings, + expiresAt: time.Now().Add(openAIQuotaAutoPauseSettingsCacheTTL).UnixNano(), + }) } // GetRectifierSettings 获取请求整流器配置 diff --git a/backend/internal/service/wire.go b/backend/internal/service/wire.go index b22e10ae..d3e4ce51 100644 --- a/backend/internal/service/wire.go +++ b/backend/internal/service/wire.go @@ -396,6 +396,46 @@ func ProvideBackupService( return svc } +// ProvideOpsService constructs OpsService and wires the SettingService-backed quota +// auto-pause cache sink. Mirrors the SetCleanupReloader pattern: OpsService doesn't +// hold a *SettingService reference, but wire injects a tiny callback so writes to +// ops_advanced_settings immediately propagate into the scheduler hot-path cache. +func ProvideOpsService( + opsRepo OpsRepository, + settingRepo SettingRepository, + cfg *config.Config, + accountRepo AccountRepository, + userRepo UserRepository, + concurrencyService *ConcurrencyService, + gatewayService *GatewayService, + openAIGatewayService *OpenAIGatewayService, + geminiCompatService *GeminiMessagesCompatService, + antigravityGatewayService *AntigravityGatewayService, + systemLogSink *OpsSystemLogSink, + settingService *SettingService, +) *OpsService { + svc := NewOpsService( + opsRepo, + settingRepo, + cfg, + accountRepo, + userRepo, + concurrencyService, + gatewayService, + openAIGatewayService, + geminiCompatService, + antigravityGatewayService, + systemLogSink, + ) + if settingService != nil { + svc.SetOpenAIQuotaAutoPauseSettingsSink(settingService.SetOpenAIQuotaAutoPauseSettings) + // Optional warm-up so the first scheduled request after process start observes + // a populated cache rather than zero defaults. Best-effort, sync-bounded. + settingService.WarmOpenAIQuotaAutoPauseSettings(context.Background()) + } + return svc +} + // ProvideSettingService wires SettingService with group reader and proxy repo. func ProvideSettingService(settingRepo SettingRepository, groupRepo GroupRepository, proxyRepo ProxyRepository, cfg *config.Config) *SettingService { svc := NewSettingService(settingRepo, cfg) @@ -481,7 +521,7 @@ var ProviderSet = wire.NewSet( NewDataManagementService, ProvideBackupService, ProvideOpsSystemLogSink, - NewOpsService, + ProvideOpsService, ProvideOpsMetricsCollector, ProvideOpsAggregationService, ProvideOpsAlertEvaluatorService, diff --git a/frontend/src/components/account/EditAccountModal.vue b/frontend/src/components/account/EditAccountModal.vue index 470c0bfb..8dc85d0e 100644 --- a/frontend/src/components/account/EditAccountModal.vue +++ b/frontend/src/components/account/EditAccountModal.vue @@ -1791,6 +1791,28 @@ v-if="account?.platform === 'openai'" class="border-t border-gray-200 pt-4 dark:border-dark-600 space-y-4" > +
+
+ + +
+

{{ t('admin.accounts.autoPauseDisabledHint') }}

+

{{ t('admin.accounts.autoPauseThresholdHint') }}

+
+
+ + +
+

{{ t('admin.accounts.autoPauseDisabledHint') }}

+

{{ t('admin.accounts.autoPauseThresholdHint') }}

@@ -2481,6 +2527,8 @@ const interceptWarmupRequests = ref(false) const autoPauseOnExpired = ref(false) const autoPause5hThreshold = ref(null) const autoPause7dThreshold = ref(null) +const autoPause5hDisabled = ref(false) +const autoPause7dDisabled = ref(false) const mixedScheduling = ref(false) // For antigravity accounts: enable mixed scheduling const allowOverages = ref(false) // For antigravity accounts: enable AI Credits overages const antigravityModelRestrictionMode = ref<'whitelist' | 'mapping'>('whitelist') @@ -2901,6 +2949,8 @@ const syncFormFromAccount = (newAccount: Account | null) => { allowOverages.value = extra?.allow_overages === true autoPause5hThreshold.value = typeof extra?.auto_pause_5h_threshold === 'number' ? extra.auto_pause_5h_threshold * 100 : null autoPause7dThreshold.value = typeof extra?.auto_pause_7d_threshold === 'number' ? extra.auto_pause_7d_threshold * 100 : null + autoPause5hDisabled.value = extra?.auto_pause_5h_disabled === true + autoPause7dDisabled.value = extra?.auto_pause_7d_disabled === true // Load OpenAI passthrough toggle (OpenAI OAuth/API Key) openaiPassthroughEnabled.value = false @@ -4064,6 +4114,16 @@ const handleSubmit = async () => { } else { delete newExtra.auto_pause_7d_threshold } + if (autoPause5hDisabled.value) { + newExtra.auto_pause_5h_disabled = true + } else { + delete newExtra.auto_pause_5h_disabled + } + if (autoPause7dDisabled.value) { + newExtra.auto_pause_7d_disabled = true + } else { + delete newExtra.auto_pause_7d_disabled + } delete newExtra.codex_image_generation_bridge_enabled if (codexImageGenerationBridgeMode.value === 'inherit') { diff --git a/frontend/src/components/account/__tests__/EditAccountModal.spec.ts b/frontend/src/components/account/__tests__/EditAccountModal.spec.ts index 6db63831..f4865de9 100644 --- a/frontend/src/components/account/__tests__/EditAccountModal.spec.ts +++ b/frontend/src/components/account/__tests__/EditAccountModal.spec.ts @@ -352,6 +352,27 @@ describe('EditAccountModal', () => { expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.auto_pause_7d_threshold).toBe(0.96) }) + it('submits OpenAI quota auto-pause disable flag in extra', async () => { + // Toggling the per-account disable flag must persist as auto_pause_5h_disabled + // so an admin can exempt one account from auto-pause even when a global default + // threshold is configured (otherwise leaving the threshold blank would silently + // fall back to the global default). + const account = buildAccount() + updateAccountMock.mockReset() + checkMixedChannelRiskMock.mockReset() + checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false }) + updateAccountMock.mockResolvedValue(account) + + const wrapper = mountModal(account) + + await wrapper.get('[data-testid="auto-pause-5h-disabled"]').trigger('click') + await wrapper.get('form#edit-account-form').trigger('submit.prevent') + + expect(updateAccountMock).toHaveBeenCalledTimes(1) + expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.auto_pause_5h_disabled).toBe(true) + expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.auto_pause_7d_disabled).toBeUndefined() + }) + it('keeps at least one OpenAI APIKey endpoint capability selected', async () => { const account = buildAccount() updateAccountMock.mockReset() diff --git a/frontend/src/i18n/locales/en.ts b/frontend/src/i18n/locales/en.ts index 8ab90961..6735029c 100644 --- a/frontend/src/i18n/locales/en.ts +++ b/frontend/src/i18n/locales/en.ts @@ -3477,7 +3477,10 @@ export default { autoPauseOnExpiredDesc: 'When enabled, the account will auto pause scheduling after it expires', autoPause5hThreshold: '5h Usage Threshold (%)', autoPause7dThreshold: '7d Usage Threshold (%)', - autoPauseThresholdHint: 'Leave empty or set 0 to disable. Reaching the threshold only skips the account during scheduling and does not modify schedulable.', + autoPauseThresholdHint: 'Leave empty or set 0 to use the global default threshold (configured in Ops settings); set a value to override the global default. Reaching the threshold only skips the account during scheduling and does not modify schedulable.', + autoPause5hDisabled: 'Disable 5h auto-pause', + autoPause7dDisabled: 'Disable 7d auto-pause', + autoPauseDisabledHint: 'When enabled, this account is never auto-paused (even if a global default threshold is configured).', // Quota control (Anthropic OAuth/SetupToken only) quotaControl: { title: 'Quota Control', diff --git a/frontend/src/i18n/locales/zh.ts b/frontend/src/i18n/locales/zh.ts index 4f1d1f13..abb8dff7 100644 --- a/frontend/src/i18n/locales/zh.ts +++ b/frontend/src/i18n/locales/zh.ts @@ -3615,7 +3615,10 @@ export default { autoPauseOnExpiredDesc: '启用后,账号过期将自动暂停调度', autoPause5hThreshold: '5h 用量阈值(%)', autoPause7dThreshold: '7d 用量阈值(%)', - autoPauseThresholdHint: '填 0 或留空表示不启用;达到阈值后仅在调度时跳过账号,不修改 schedulable。', + autoPauseThresholdHint: '留空或填 0 表示使用全局默认阈值(在运维设置中配置);填具体值则覆盖全局默认。达到阈值后仅在调度时跳过账号,不修改 schedulable。', + autoPause5hDisabled: '禁用 5h 自动暂停', + autoPause7dDisabled: '禁用 7d 自动暂停', + autoPauseDisabledHint: '开启后该账号永不进入自动暂停(即使全局默认阈值已配置)。', // Quota control (Anthropic OAuth/SetupToken only) quotaControl: { title: '配额控制', From 0a521f09fbc481ccdd2a0d6bccfaa8f80df220fc Mon Sep 17 00:00:00 2001 From: Pluviobyte Date: Fri, 29 May 2026 06:46:49 +0000 Subject: [PATCH 38/79] fix(gemini): close tool_use block before text in messages streaming When the Gemini->Anthropic streaming bridge for the /v1/messages endpoint receives a functionCall part followed by a text part, the text branch in handleStreamingResponse opened a new text content block without closing the already-open tool_use block. The tool block's content_block_stop was only emitted at end-of-stream, after the text block's content_block_start, so the Anthropic SSE stream contained overlapping/unterminated content blocks. Clients that assemble messages by block index (e.g. Claude Code) can drop the tool input or mis-parse the response. The functionCall branch already closes an open text block before opening a tool block, and the chat-completions sibling closes the tool block in its text branch via closeOpenTool(). This applies the same symmetric handling to the messages variant: close any open tool_use block (resetting openToolIndex/openToolName/ seenToolJSON) before starting text. Adds a regression test that replays a tool->text Gemini stream and asserts the Anthropic content-block lifecycle never overlaps. --- .../service/gemini_messages_compat_service.go | 16 +++ .../gemini_messages_compat_service_test.go | 105 ++++++++++++++++++ 2 files changed, 121 insertions(+) diff --git a/backend/internal/service/gemini_messages_compat_service.go b/backend/internal/service/gemini_messages_compat_service.go index 516556ca..64f19b2e 100644 --- a/backend/internal/service/gemini_messages_compat_service.go +++ b/backend/internal/service/gemini_messages_compat_service.go @@ -2031,6 +2031,22 @@ func (s *GeminiMessagesCompatService) handleStreamingResponse(c *gin.Context, re parts := extractGeminiParts(geminiResp) for _, part := range parts { if text, ok := part["text"].(string); ok && text != "" { + // Close an open tool_use block before starting text, mirroring + // the functionCall branch (which closes open text blocks) and + // the chat-completions sibling's closeOpenTool(). Otherwise a + // tool→text sequence keeps the tool_use block open while the + // text block starts, emitting overlapping Anthropic content + // blocks that violate the SSE contract. + if openToolIndex >= 0 { + writeSSE(c.Writer, "content_block_stop", map[string]any{ + "type": "content_block_stop", + "index": openToolIndex, + }) + openToolIndex = -1 + openToolName = "" + seenToolJSON = "" + } + delta, newSeen := computeGeminiTextDelta(seenText, text) seenText = newSeen if delta == "" { diff --git a/backend/internal/service/gemini_messages_compat_service_test.go b/backend/internal/service/gemini_messages_compat_service_test.go index d0560344..79db633a 100644 --- a/backend/internal/service/gemini_messages_compat_service_test.go +++ b/backend/internal/service/gemini_messages_compat_service_test.go @@ -832,3 +832,108 @@ func TestParseGeminiRateLimitResetTime(t *testing.T) { }) } } + +// TestGeminiMessagesHandleStreamingResponse_ClosesToolBlockBeforeText guards the +// tool→text ordering in the Gemini→Anthropic (messages) streaming bridge. When +// Gemini emits a functionCall part followed by a text part, the tool_use content +// block must be closed before the text block opens; otherwise the Anthropic SSE +// stream contains overlapping content blocks. The chat-completions sibling +// already enforces this via closeOpenTool(). +func TestGeminiMessagesHandleStreamingResponse_ClosesToolBlockBeforeText(t *testing.T) { + gin.SetMode(gin.TestMode) + + upstreamBody := `data: {"candidates":[{"content":{"parts":[{"functionCall":{"name":"get_weather","args":{"city":"SF"}}}]}}]}` + "\n\n" + + `data: {"candidates":[{"content":{"parts":[{"text":"All done."}]},"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":5,"candidatesTokenCount":3}}` + "\n\n" + + "data: [DONE]\n\n" + + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(upstreamBody)), + } + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + + svc := &GeminiMessagesCompatService{} + result, err := svc.handleStreamingResponse(c, resp, time.Now(), "claude-3-5-sonnet") + require.NoError(t, err) + require.NotNil(t, result) + + events := parseAnthropicContentBlockEvents(t, rec.Body.String()) + + // Anthropic allows at most one content block open at a time: every + // content_block_start must be matched by a content_block_stop before the + // next start. Replay the lifecycle and assert there is no overlap. + open := -1 + blockTypes := map[int]string{} + textStarted := false + toolClosed := false + toolClosedBeforeText := false + for _, ev := range events { + switch ev.event { + case "content_block_start": + require.Equalf(t, -1, open, + "content block %d opened while block %d was still open (overlapping blocks)", ev.index, open) + open = ev.index + blockTypes[ev.index] = ev.blockType + if ev.blockType == "text" { + textStarted = true + if toolClosed { + toolClosedBeforeText = true + } + } + case "content_block_stop": + require.Equalf(t, open, ev.index, + "content_block_stop index %d does not match the open block %d", ev.index, open) + if blockTypes[ev.index] == "tool_use" { + toolClosed = true + } + open = -1 + } + } + + require.True(t, textStarted, "expected a text content block to be emitted after the tool call") + require.True(t, toolClosedBeforeText, "tool_use block must be closed before the text block starts") + require.Equal(t, -1, open, "stream ended with a content block still open") +} + +type anthropicContentBlockEvent struct { + event string + index int + blockType string +} + +// parseAnthropicContentBlockEvents extracts content_block_start/stop events (with +// their index and, for starts, the content block type) from an Anthropic SSE body. +func parseAnthropicContentBlockEvents(t *testing.T, raw string) []anthropicContentBlockEvent { + t.Helper() + var events []anthropicContentBlockEvent + for _, chunk := range strings.Split(raw, "\n\n") { + var eventName, dataLine string + for _, line := range strings.Split(chunk, "\n") { + switch { + case strings.HasPrefix(line, "event:"): + eventName = strings.TrimSpace(strings.TrimPrefix(line, "event:")) + case strings.HasPrefix(line, "data:"): + dataLine = strings.TrimSpace(strings.TrimPrefix(line, "data:")) + } + } + if eventName != "content_block_start" && eventName != "content_block_stop" { + continue + } + var payload struct { + Index int `json:"index"` + ContentBlock struct { + Type string `json:"type"` + } `json:"content_block"` + } + require.NoError(t, json.Unmarshal([]byte(dataLine), &payload)) + events = append(events, anthropicContentBlockEvent{ + event: eventName, + index: payload.Index, + blockType: payload.ContentBlock.Type, + }) + } + return events +} From 68901cbfff783af794d96028be8dad3e532c0fe7 Mon Sep 17 00:00:00 2001 From: shaw Date: Fri, 29 May 2026 16:29:29 +0800 Subject: [PATCH 39/79] chore(pricing): update model pricing metadata --- .../model_prices_and_context_window.json | 3307 ++++++++--------- 1 file changed, 1578 insertions(+), 1729 deletions(-) diff --git a/backend/resources/model-pricing/model_prices_and_context_window.json b/backend/resources/model-pricing/model_prices_and_context_window.json index 3cae8c8b..e88ed2da 100644 --- a/backend/resources/model-pricing/model_prices_and_context_window.json +++ b/backend/resources/model-pricing/model_prices_and_context_window.json @@ -1,135 +1,4 @@ { - "claude-3-5-haiku-20241022": { - "cache_creation_input_token_cost": 1e-06, - "cache_creation_input_token_cost_above_1hr": 6e-06, - "cache_read_input_token_cost": 8e-08, - "deprecation_date": "2025-10-01", - "input_cost_per_token": 8e-07, - "litellm_provider": "anthropic", - "max_input_tokens": 200000, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_token": 4e-06, - "search_context_cost_per_query": { - "search_context_size_high": 0.01, - "search_context_size_low": 0.01, - "search_context_size_medium": 0.01 - }, - "supports_assistant_prefill": true, - "supports_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_response_schema": true, - "supports_tool_choice": true, - "supports_vision": true, - "supports_web_search": true, - "tool_use_system_prompt_tokens": 264 - }, - "claude-3-5-haiku-latest": { - "cache_creation_input_token_cost": 1.25e-06, - "cache_creation_input_token_cost_above_1hr": 6e-06, - "cache_read_input_token_cost": 1e-07, - "deprecation_date": "2025-10-01", - "input_cost_per_token": 1e-06, - "litellm_provider": "anthropic", - "max_input_tokens": 200000, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_token": 5e-06, - "search_context_cost_per_query": { - "search_context_size_high": 0.01, - "search_context_size_low": 0.01, - "search_context_size_medium": 0.01 - }, - "supports_assistant_prefill": true, - "supports_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_response_schema": true, - "supports_tool_choice": true, - "supports_vision": true, - "supports_web_search": true, - "tool_use_system_prompt_tokens": 264 - }, - "claude-3-5-sonnet-20240620": { - "cache_creation_input_token_cost": 3.75e-06, - "cache_creation_input_token_cost_above_1hr": 6e-06, - "cache_read_input_token_cost": 3e-07, - "deprecation_date": "2025-06-01", - "input_cost_per_token": 3e-06, - "litellm_provider": "anthropic", - "max_input_tokens": 200000, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_token": 1.5e-05, - "supports_assistant_prefill": true, - "supports_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_response_schema": true, - "supports_tool_choice": true, - "supports_vision": true, - "tool_use_system_prompt_tokens": 159 - }, - "claude-3-5-sonnet-20241022": { - "cache_creation_input_token_cost": 3.75e-06, - "cache_creation_input_token_cost_above_1hr": 6e-06, - "cache_read_input_token_cost": 3e-07, - "deprecation_date": "2025-10-01", - "input_cost_per_token": 3e-06, - "litellm_provider": "anthropic", - "max_input_tokens": 200000, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_token": 1.5e-05, - "search_context_cost_per_query": { - "search_context_size_high": 0.01, - "search_context_size_low": 0.01, - "search_context_size_medium": 0.01 - }, - "supports_assistant_prefill": true, - "supports_computer_use": true, - "supports_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_response_schema": true, - "supports_tool_choice": true, - "supports_vision": true, - "supports_web_search": true, - "tool_use_system_prompt_tokens": 159 - }, - "claude-3-5-sonnet-latest": { - "cache_creation_input_token_cost": 3.75e-06, - "cache_creation_input_token_cost_above_1hr": 6e-06, - "cache_read_input_token_cost": 3e-07, - "deprecation_date": "2025-06-01", - "input_cost_per_token": 3e-06, - "litellm_provider": "anthropic", - "max_input_tokens": 200000, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_token": 1.5e-05, - "search_context_cost_per_query": { - "search_context_size_high": 0.01, - "search_context_size_low": 0.01, - "search_context_size_medium": 0.01 - }, - "supports_assistant_prefill": true, - "supports_computer_use": true, - "supports_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_response_schema": true, - "supports_tool_choice": true, - "supports_vision": true, - "supports_web_search": true, - "tool_use_system_prompt_tokens": 159 - }, "claude-3-7-sonnet-20250219": { "cache_creation_input_token_cost": 3.75e-06, "cache_creation_input_token_cost_above_1hr": 6e-06, @@ -159,34 +28,6 @@ "supports_web_search": true, "tool_use_system_prompt_tokens": 159 }, - "claude-3-7-sonnet-latest": { - "cache_creation_input_token_cost": 3.75e-06, - "cache_creation_input_token_cost_above_1hr": 6e-06, - "cache_read_input_token_cost": 3e-07, - "deprecation_date": "2025-06-01", - "input_cost_per_token": 3e-06, - "litellm_provider": "anthropic", - "max_input_tokens": 200000, - "max_output_tokens": 64000, - "max_tokens": 64000, - "mode": "chat", - "output_cost_per_token": 1.5e-05, - "search_context_cost_per_query": { - "search_context_size_high": 0.01, - "search_context_size_low": 0.01, - "search_context_size_medium": 0.01 - }, - "supports_assistant_prefill": true, - "supports_computer_use": true, - "supports_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_response_schema": true, - "supports_tool_choice": true, - "supports_vision": true, - "tool_use_system_prompt_tokens": 159 - }, "claude-3-haiku-20240307": { "cache_creation_input_token_cost": 3e-07, "cache_creation_input_token_cost_above_1hr": 6e-06, @@ -226,28 +67,9 @@ "supports_vision": true, "tool_use_system_prompt_tokens": 395 }, - "claude-3-opus-latest": { - "cache_creation_input_token_cost": 1.875e-05, - "cache_creation_input_token_cost_above_1hr": 6e-06, - "cache_read_input_token_cost": 1.5e-06, - "deprecation_date": "2025-03-01", - "input_cost_per_token": 1.5e-05, - "litellm_provider": "anthropic", - "max_input_tokens": 200000, - "max_output_tokens": 4096, - "max_tokens": 4096, - "mode": "chat", - "output_cost_per_token": 7.5e-05, - "supports_assistant_prefill": true, - "supports_function_calling": true, - "supports_prompt_caching": true, - "supports_response_schema": true, - "supports_tool_choice": true, - "supports_vision": true, - "tool_use_system_prompt_tokens": 395 - }, "claude-4-opus-20250514": { "cache_creation_input_token_cost": 1.875e-05, + "cache_creation_input_token_cost_above_1hr": 3e-05, "cache_read_input_token_cost": 1.5e-06, "input_cost_per_token": 1.5e-05, "litellm_provider": "anthropic", @@ -274,6 +96,7 @@ }, "claude-4-sonnet-20250514": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06, "cache_read_input_token_cost": 3e-07, "cache_read_input_token_cost_above_200k_tokens": 6e-07, @@ -447,6 +270,7 @@ "supports_assistant_prefill": true, "supports_computer_use": true, "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_pdf_input": true, "supports_prompt_caching": true, "supports_reasoning": true, @@ -474,6 +298,7 @@ "supports_assistant_prefill": true, "supports_computer_use": true, "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_pdf_input": true, "supports_prompt_caching": true, "supports_reasoning": true, @@ -485,18 +310,14 @@ "claude-opus-4-6": { "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, - "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05, "cache_read_input_token_cost": 5e-07, - "cache_read_input_token_cost_above_200k_tokens": 1e-06, "input_cost_per_token": 5e-06, - "input_cost_per_token_above_200k_tokens": 1e-05, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, - "output_cost_per_token_above_200k_tokens": 3.75e-05, "provider_specific_entry": { "fast": 6.0, "us": 1.1 @@ -506,9 +327,13 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "supports_adaptive_thinking": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": true, + "supports_output_config": true, "supports_pdf_input": true, "supports_prompt_caching": true, "supports_reasoning": true, @@ -520,18 +345,14 @@ "claude-opus-4-6-20260205": { "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, - "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05, "cache_read_input_token_cost": 5e-07, - "cache_read_input_token_cost_above_200k_tokens": 1e-06, "input_cost_per_token": 5e-06, - "input_cost_per_token_above_200k_tokens": 1e-05, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, - "output_cost_per_token_above_200k_tokens": 3.75e-05, "provider_specific_entry": { "fast": 6.0, "us": 1.1 @@ -541,9 +362,13 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "supports_adaptive_thinking": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": true, + "supports_output_config": true, "supports_pdf_input": true, "supports_prompt_caching": true, "supports_reasoning": true, @@ -555,18 +380,14 @@ "claude-opus-4-6-thinking": { "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, - "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05, "cache_read_input_token_cost": 5e-07, - "cache_read_input_token_cost_above_200k_tokens": 1e-06, "input_cost_per_token": 5e-06, - "input_cost_per_token_above_200k_tokens": 1e-05, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, - "output_cost_per_token_above_200k_tokens": 3.75e-05, "provider_specific_entry": { "fast": 6.0, "us": 1.1 @@ -576,9 +397,13 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "supports_adaptive_thinking": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": true, + "supports_output_config": true, "supports_pdf_input": true, "supports_prompt_caching": true, "supports_reasoning": true, @@ -587,6 +412,114 @@ "supports_vision": true, "tool_use_system_prompt_tokens": 346 }, + "claude-opus-4-7": { + "cache_creation_input_token_cost": 6.25e-06, + "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_read_input_token_cost": 5e-07, + "input_cost_per_token": 5e-06, + "litellm_provider": "anthropic", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 2.5e-05, + "provider_specific_entry": { + "fast": 6.0, + "us": 1.1 + }, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": true, + "supports_output_config": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "tool_use_system_prompt_tokens": 346 + }, + "claude-opus-4-7-20260416": { + "cache_creation_input_token_cost": 6.25e-06, + "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_read_input_token_cost": 5e-07, + "input_cost_per_token": 5e-06, + "litellm_provider": "anthropic", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 2.5e-05, + "provider_specific_entry": { + "fast": 6.0, + "us": 1.1 + }, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": true, + "supports_output_config": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "tool_use_system_prompt_tokens": 346 + }, + "claude-opus-4-8": { + "cache_creation_input_token_cost": 6.25e-06, + "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_read_input_token_cost": 5e-07, + "input_cost_per_token": 5e-06, + "litellm_provider": "anthropic", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 2.5e-05, + "provider_specific_entry": { + "fast": 6.0, + "us": 1.1 + }, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_adaptive_thinking": true, + "supports_assistant_prefill": false, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": true, + "supports_output_config": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true, + "tool_use_system_prompt_tokens": 346 + }, "claude-sonnet-4-20250514": { "cache_creation_input_token_cost": 3.75e-06, "cache_creation_input_token_cost_above_1hr": 6e-06, @@ -621,6 +554,7 @@ }, "claude-sonnet-4-5": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06, "cache_read_input_token_cost": 3e-07, "cache_read_input_token_cost_above_200k_tokens": 6e-07, @@ -651,6 +585,7 @@ }, "claude-sonnet-4-5-20250929": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06, "cache_read_input_token_cost": 3e-07, "cache_read_input_token_cost_above_200k_tokens": 6e-07, @@ -682,6 +617,7 @@ }, "claude-sonnet-4-5-20250929-v1:0": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06, "cache_read_input_token_cost": 3e-07, "cache_read_input_token_cost_above_200k_tokens": 6e-07, @@ -707,26 +643,27 @@ }, "claude-sonnet-4-6": { "cache_creation_input_token_cost": 3.75e-06, - "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_read_input_token_cost": 3e-07, - "cache_read_input_token_cost_above_200k_tokens": 6e-07, "input_cost_per_token": 3e-06, - "input_cost_per_token_above_200k_tokens": 6e-06, "litellm_provider": "anthropic", - "max_input_tokens": 200000, + "max_input_tokens": 1000000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 1.5e-05, - "output_cost_per_token_above_200k_tokens": 2.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "supports_adaptive_thinking": true, "supports_assistant_prefill": true, "supports_computer_use": true, "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": true, + "supports_output_config": true, "supports_pdf_input": true, "supports_prompt_caching": true, "supports_reasoning": true, @@ -735,6 +672,54 @@ "supports_vision": true, "tool_use_system_prompt_tokens": 346 }, + "codex-auto-review": { + "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_above_272k_tokens": 1e-06, + "cache_read_input_token_cost_flex": 2.5e-07, + "cache_read_input_token_cost_priority": 1e-06, + "input_cost_per_token": 5e-06, + "input_cost_per_token_above_272k_tokens": 1e-05, + "input_cost_per_token_batches": 2.5e-06, + "input_cost_per_token_flex": 2.5e-06, + "input_cost_per_token_priority": 1e-05, + "litellm_provider": "openai", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 3e-05, + "output_cost_per_token_above_272k_tokens": 4.5e-05, + "output_cost_per_token_batches": 1.5e-05, + "output_cost_per_token_flex": 1.5e-05, + "output_cost_per_token_priority": 6e-05, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, "deepseek-chat": { "cache_read_input_token_cost": 2.8e-08, "input_cost_per_token": 2.8e-07, @@ -792,478 +777,9 @@ "supports_reasoning": true, "supports_tool_choice": true }, - "gemini-1.0-pro": { - "input_cost_per_character": 1.25e-07, - "input_cost_per_image": 0.0025, - "input_cost_per_token": 5e-07, - "input_cost_per_video_per_second": 0.002, - "litellm_provider": "vertex_ai-language-models", - "max_input_tokens": 32760, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_character": 3.75e-07, - "output_cost_per_token": 1.5e-06, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#google_models", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_tool_choice": true - }, - "gemini-1.0-pro-001": { - "deprecation_date": "2025-04-09", - "input_cost_per_character": 1.25e-07, - "input_cost_per_image": 0.0025, - "input_cost_per_token": 5e-07, - "input_cost_per_video_per_second": 0.002, - "litellm_provider": "vertex_ai-language-models", - "max_input_tokens": 32760, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_character": 3.75e-07, - "output_cost_per_token": 1.5e-06, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_tool_choice": true - }, - "gemini-1.0-pro-002": { - "deprecation_date": "2025-04-09", - "input_cost_per_character": 1.25e-07, - "input_cost_per_image": 0.0025, - "input_cost_per_token": 5e-07, - "input_cost_per_video_per_second": 0.002, - "litellm_provider": "vertex_ai-language-models", - "max_input_tokens": 32760, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_character": 3.75e-07, - "output_cost_per_token": 1.5e-06, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_tool_choice": true - }, - "gemini-1.0-pro-vision": { - "input_cost_per_image": 0.0025, - "input_cost_per_token": 5e-07, - "litellm_provider": "vertex_ai-vision-models", - "max_images_per_prompt": 16, - "max_input_tokens": 16384, - "max_output_tokens": 2048, - "max_tokens": 2048, - "max_video_length": 2, - "max_videos_per_prompt": 1, - "mode": "chat", - "output_cost_per_token": 1.5e-06, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_tool_choice": true, - "supports_vision": true - }, - "gemini-1.0-pro-vision-001": { - "deprecation_date": "2025-04-09", - "input_cost_per_image": 0.0025, - "input_cost_per_token": 5e-07, - "litellm_provider": "vertex_ai-vision-models", - "max_images_per_prompt": 16, - "max_input_tokens": 16384, - "max_output_tokens": 2048, - "max_tokens": 2048, - "max_video_length": 2, - "max_videos_per_prompt": 1, - "mode": "chat", - "output_cost_per_token": 1.5e-06, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_tool_choice": true, - "supports_vision": true - }, - "gemini-1.0-ultra": { - "input_cost_per_character": 1.25e-07, - "input_cost_per_image": 0.0025, - "input_cost_per_token": 5e-07, - "input_cost_per_video_per_second": 0.002, - "litellm_provider": "vertex_ai-language-models", - "max_input_tokens": 8192, - "max_output_tokens": 2048, - "max_tokens": 2048, - "mode": "chat", - "output_cost_per_character": 3.75e-07, - "output_cost_per_token": 1.5e-06, - "source": "As of Jun, 2024. There is no available doc on vertex ai pricing gemini-1.0-ultra-001. Using gemini-1.0-pro pricing. Got max_tokens info here: https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_tool_choice": true - }, - "gemini-1.0-ultra-001": { - "input_cost_per_character": 1.25e-07, - "input_cost_per_image": 0.0025, - "input_cost_per_token": 5e-07, - "input_cost_per_video_per_second": 0.002, - "litellm_provider": "vertex_ai-language-models", - "max_input_tokens": 8192, - "max_output_tokens": 2048, - "max_tokens": 2048, - "mode": "chat", - "output_cost_per_character": 3.75e-07, - "output_cost_per_token": 1.5e-06, - "source": "As of Jun, 2024. There is no available doc on vertex ai pricing gemini-1.0-ultra-001. Using gemini-1.0-pro pricing. Got max_tokens info here: https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_tool_choice": true - }, - "gemini-1.5-flash": { - "deprecation_date": "2025-09-29", - "input_cost_per_audio_per_second": 2e-06, - "input_cost_per_audio_per_second_above_128k_tokens": 4e-06, - "input_cost_per_character": 1.875e-08, - "input_cost_per_character_above_128k_tokens": 2.5e-07, - "input_cost_per_image": 2e-05, - "input_cost_per_image_above_128k_tokens": 4e-05, - "input_cost_per_token": 7.5e-08, - "input_cost_per_token_above_128k_tokens": 1e-06, - "input_cost_per_video_per_second": 2e-05, - "input_cost_per_video_per_second_above_128k_tokens": 4e-05, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1000000, - "max_output_tokens": 8192, - "max_pdf_size_mb": 30, - "max_tokens": 8192, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_character": 7.5e-08, - "output_cost_per_character_above_128k_tokens": 1.5e-07, - "output_cost_per_token": 3e-07, - "output_cost_per_token_above_128k_tokens": 6e-07, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true - }, - "gemini-1.5-flash-001": { - "deprecation_date": "2025-05-24", - "input_cost_per_audio_per_second": 2e-06, - "input_cost_per_audio_per_second_above_128k_tokens": 4e-06, - "input_cost_per_character": 1.875e-08, - "input_cost_per_character_above_128k_tokens": 2.5e-07, - "input_cost_per_image": 2e-05, - "input_cost_per_image_above_128k_tokens": 4e-05, - "input_cost_per_token": 7.5e-08, - "input_cost_per_token_above_128k_tokens": 1e-06, - "input_cost_per_video_per_second": 2e-05, - "input_cost_per_video_per_second_above_128k_tokens": 4e-05, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1000000, - "max_output_tokens": 8192, - "max_pdf_size_mb": 30, - "max_tokens": 8192, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_character": 7.5e-08, - "output_cost_per_character_above_128k_tokens": 1.5e-07, - "output_cost_per_token": 3e-07, - "output_cost_per_token_above_128k_tokens": 6e-07, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true - }, - "gemini-1.5-flash-002": { - "deprecation_date": "2025-09-24", - "input_cost_per_audio_per_second": 2e-06, - "input_cost_per_audio_per_second_above_128k_tokens": 4e-06, - "input_cost_per_character": 1.875e-08, - "input_cost_per_character_above_128k_tokens": 2.5e-07, - "input_cost_per_image": 2e-05, - "input_cost_per_image_above_128k_tokens": 4e-05, - "input_cost_per_token": 7.5e-08, - "input_cost_per_token_above_128k_tokens": 1e-06, - "input_cost_per_video_per_second": 2e-05, - "input_cost_per_video_per_second_above_128k_tokens": 4e-05, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1048576, - "max_output_tokens": 8192, - "max_pdf_size_mb": 30, - "max_tokens": 8192, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_character": 7.5e-08, - "output_cost_per_character_above_128k_tokens": 1.5e-07, - "output_cost_per_token": 3e-07, - "output_cost_per_token_above_128k_tokens": 6e-07, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#gemini-1.5-flash", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true - }, - "gemini-1.5-flash-exp-0827": { - "deprecation_date": "2025-09-29", - "input_cost_per_audio_per_second": 2e-06, - "input_cost_per_audio_per_second_above_128k_tokens": 4e-06, - "input_cost_per_character": 1.875e-08, - "input_cost_per_character_above_128k_tokens": 2.5e-07, - "input_cost_per_image": 2e-05, - "input_cost_per_image_above_128k_tokens": 4e-05, - "input_cost_per_token": 4.688e-09, - "input_cost_per_token_above_128k_tokens": 1e-06, - "input_cost_per_video_per_second": 2e-05, - "input_cost_per_video_per_second_above_128k_tokens": 4e-05, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1000000, - "max_output_tokens": 8192, - "max_pdf_size_mb": 30, - "max_tokens": 8192, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_character": 1.875e-08, - "output_cost_per_character_above_128k_tokens": 3.75e-08, - "output_cost_per_token": 4.6875e-09, - "output_cost_per_token_above_128k_tokens": 9.375e-09, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true - }, - "gemini-1.5-flash-preview-0514": { - "deprecation_date": "2025-09-29", - "input_cost_per_audio_per_second": 2e-06, - "input_cost_per_audio_per_second_above_128k_tokens": 4e-06, - "input_cost_per_character": 1.875e-08, - "input_cost_per_character_above_128k_tokens": 2.5e-07, - "input_cost_per_image": 2e-05, - "input_cost_per_image_above_128k_tokens": 4e-05, - "input_cost_per_token": 7.5e-08, - "input_cost_per_token_above_128k_tokens": 1e-06, - "input_cost_per_video_per_second": 2e-05, - "input_cost_per_video_per_second_above_128k_tokens": 4e-05, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1000000, - "max_output_tokens": 8192, - "max_pdf_size_mb": 30, - "max_tokens": 8192, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_character": 1.875e-08, - "output_cost_per_character_above_128k_tokens": 3.75e-08, - "output_cost_per_token": 4.6875e-09, - "output_cost_per_token_above_128k_tokens": 9.375e-09, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true - }, - "gemini-1.5-pro": { - "deprecation_date": "2025-09-29", - "input_cost_per_audio_per_second": 3.125e-05, - "input_cost_per_audio_per_second_above_128k_tokens": 6.25e-05, - "input_cost_per_character": 3.125e-07, - "input_cost_per_character_above_128k_tokens": 6.25e-07, - "input_cost_per_image": 0.00032875, - "input_cost_per_image_above_128k_tokens": 0.0006575, - "input_cost_per_token": 1.25e-06, - "input_cost_per_token_above_128k_tokens": 2.5e-06, - "input_cost_per_video_per_second": 0.00032875, - "input_cost_per_video_per_second_above_128k_tokens": 0.0006575, - "litellm_provider": "vertex_ai-language-models", - "max_input_tokens": 2097152, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_character": 1.25e-06, - "output_cost_per_character_above_128k_tokens": 2.5e-06, - "output_cost_per_token": 5e-06, - "output_cost_per_token_above_128k_tokens": 1e-05, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_pdf_input": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true - }, - "gemini-1.5-pro-001": { - "deprecation_date": "2025-05-24", - "input_cost_per_audio_per_second": 3.125e-05, - "input_cost_per_audio_per_second_above_128k_tokens": 6.25e-05, - "input_cost_per_character": 3.125e-07, - "input_cost_per_character_above_128k_tokens": 6.25e-07, - "input_cost_per_image": 0.00032875, - "input_cost_per_image_above_128k_tokens": 0.0006575, - "input_cost_per_token": 1.25e-06, - "input_cost_per_token_above_128k_tokens": 2.5e-06, - "input_cost_per_video_per_second": 0.00032875, - "input_cost_per_video_per_second_above_128k_tokens": 0.0006575, - "litellm_provider": "vertex_ai-language-models", - "max_input_tokens": 1000000, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_character": 1.25e-06, - "output_cost_per_character_above_128k_tokens": 2.5e-06, - "output_cost_per_token": 5e-06, - "output_cost_per_token_above_128k_tokens": 1e-05, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true - }, - "gemini-1.5-pro-002": { - "deprecation_date": "2025-09-24", - "input_cost_per_audio_per_second": 3.125e-05, - "input_cost_per_audio_per_second_above_128k_tokens": 6.25e-05, - "input_cost_per_character": 3.125e-07, - "input_cost_per_character_above_128k_tokens": 6.25e-07, - "input_cost_per_image": 0.00032875, - "input_cost_per_image_above_128k_tokens": 0.0006575, - "input_cost_per_token": 1.25e-06, - "input_cost_per_token_above_128k_tokens": 2.5e-06, - "input_cost_per_video_per_second": 0.00032875, - "input_cost_per_video_per_second_above_128k_tokens": 0.0006575, - "litellm_provider": "vertex_ai-language-models", - "max_input_tokens": 2097152, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_character": 1.25e-06, - "output_cost_per_character_above_128k_tokens": 2.5e-06, - "output_cost_per_token": 5e-06, - "output_cost_per_token_above_128k_tokens": 1e-05, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#gemini-1.5-pro", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true - }, - "gemini-1.5-pro-preview-0215": { - "deprecation_date": "2025-09-29", - "input_cost_per_audio_per_second": 3.125e-05, - "input_cost_per_audio_per_second_above_128k_tokens": 6.25e-05, - "input_cost_per_character": 3.125e-07, - "input_cost_per_character_above_128k_tokens": 6.25e-07, - "input_cost_per_image": 0.00032875, - "input_cost_per_image_above_128k_tokens": 0.0006575, - "input_cost_per_token": 7.8125e-08, - "input_cost_per_token_above_128k_tokens": 1.5625e-07, - "input_cost_per_video_per_second": 0.00032875, - "input_cost_per_video_per_second_above_128k_tokens": 0.0006575, - "litellm_provider": "vertex_ai-language-models", - "max_input_tokens": 1000000, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_character": 1.25e-06, - "output_cost_per_character_above_128k_tokens": 2.5e-06, - "output_cost_per_token": 3.125e-07, - "output_cost_per_token_above_128k_tokens": 6.25e-07, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true - }, - "gemini-1.5-pro-preview-0409": { - "deprecation_date": "2025-09-29", - "input_cost_per_audio_per_second": 3.125e-05, - "input_cost_per_audio_per_second_above_128k_tokens": 6.25e-05, - "input_cost_per_character": 3.125e-07, - "input_cost_per_character_above_128k_tokens": 6.25e-07, - "input_cost_per_image": 0.00032875, - "input_cost_per_image_above_128k_tokens": 0.0006575, - "input_cost_per_token": 7.8125e-08, - "input_cost_per_token_above_128k_tokens": 1.5625e-07, - "input_cost_per_video_per_second": 0.00032875, - "input_cost_per_video_per_second_above_128k_tokens": 0.0006575, - "litellm_provider": "vertex_ai-language-models", - "max_input_tokens": 1000000, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_character": 1.25e-06, - "output_cost_per_character_above_128k_tokens": 2.5e-06, - "output_cost_per_token": 3.125e-07, - "output_cost_per_token_above_128k_tokens": 6.25e-07, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_response_schema": true, - "supports_tool_choice": true - }, - "gemini-1.5-pro-preview-0514": { - "deprecation_date": "2025-09-29", - "input_cost_per_audio_per_second": 3.125e-05, - "input_cost_per_audio_per_second_above_128k_tokens": 6.25e-05, - "input_cost_per_character": 3.125e-07, - "input_cost_per_character_above_128k_tokens": 6.25e-07, - "input_cost_per_image": 0.00032875, - "input_cost_per_image_above_128k_tokens": 0.0006575, - "input_cost_per_token": 7.8125e-08, - "input_cost_per_token_above_128k_tokens": 1.5625e-07, - "input_cost_per_video_per_second": 0.00032875, - "input_cost_per_video_per_second_above_128k_tokens": 0.0006575, - "litellm_provider": "vertex_ai-language-models", - "max_input_tokens": 1000000, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_character": 1.25e-06, - "output_cost_per_character_above_128k_tokens": 2.5e-06, - "output_cost_per_token": 3.125e-07, - "output_cost_per_token_above_128k_tokens": 6.25e-07, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true - }, "gemini-2.0-flash": { "cache_read_input_token_cost": 2.5e-08, - "deprecation_date": "2026-03-31", + "deprecation_date": "2026-06-01", "input_cost_per_audio_token": 7e-07, "input_cost_per_token": 1e-07, "litellm_provider": "vertex_ai-language-models", @@ -1278,6 +794,11 @@ "max_videos_per_prompt": 10, "mode": "chat", "output_cost_per_token": 4e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.035, + "search_context_size_low": 0.035, + "search_context_size_medium": 0.035 + }, "source": "https://ai.google.dev/pricing#2_0flash", "supported_modalities": [ "text", @@ -1303,7 +824,7 @@ }, "gemini-2.0-flash-001": { "cache_read_input_token_cost": 3.75e-08, - "deprecation_date": "2026-03-31", + "deprecation_date": "2026-06-01", "input_cost_per_audio_token": 1e-06, "input_cost_per_token": 1.5e-07, "litellm_provider": "vertex_ai-language-models", @@ -1318,54 +839,11 @@ "max_videos_per_prompt": 10, "mode": "chat", "output_cost_per_token": 6e-07, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", - "supported_modalities": [ - "text", - "image", - "audio", - "video" - ], - "supported_output_modalities": [ - "text", - "image" - ], - "supports_audio_output": true, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_prompt_caching": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true, - "supports_web_search": true - }, - "gemini-2.0-flash-exp": { - "cache_read_input_token_cost": 3.75e-08, - "input_cost_per_audio_per_second": 0, - "input_cost_per_audio_per_second_above_128k_tokens": 0, - "input_cost_per_character": 0, - "input_cost_per_character_above_128k_tokens": 0, - "input_cost_per_image": 0, - "input_cost_per_image_above_128k_tokens": 0, - "input_cost_per_token": 1.5e-07, - "input_cost_per_token_above_128k_tokens": 0, - "input_cost_per_video_per_second": 0, - "input_cost_per_video_per_second_above_128k_tokens": 0, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1048576, - "max_output_tokens": 8192, - "max_pdf_size_mb": 30, - "max_tokens": 8192, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_character": 0, - "output_cost_per_character_above_128k_tokens": 0, - "output_cost_per_token": 6e-07, - "output_cost_per_token_above_128k_tokens": 0, + "search_context_cost_per_query": { + "search_context_size_high": 0.035, + "search_context_size_low": 0.035, + "search_context_size_medium": 0.035 + }, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", "supported_modalities": [ "text", @@ -1410,7 +888,7 @@ }, "gemini-2.0-flash-lite": { "cache_read_input_token_cost": 1.875e-08, - "deprecation_date": "2026-03-31", + "deprecation_date": "2026-06-01", "input_cost_per_audio_token": 7.5e-08, "input_cost_per_token": 7.5e-08, "litellm_provider": "vertex_ai-language-models", @@ -1424,6 +902,11 @@ "max_videos_per_prompt": 10, "mode": "chat", "output_cost_per_token": 3e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.035, + "search_context_size_low": 0.035, + "search_context_size_medium": 0.035 + }, "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#gemini-2.0-flash", "supported_modalities": [ "text", @@ -1446,7 +929,7 @@ }, "gemini-2.0-flash-lite-001": { "cache_read_input_token_cost": 1.875e-08, - "deprecation_date": "2026-03-31", + "deprecation_date": "2026-06-01", "input_cost_per_audio_token": 7.5e-08, "input_cost_per_token": 7.5e-08, "litellm_provider": "vertex_ai-language-models", @@ -1460,6 +943,11 @@ "max_videos_per_prompt": 10, "mode": "chat", "output_cost_per_token": 3e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.035, + "search_context_size_low": 0.035, + "search_context_size_medium": 0.035 + }, "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#gemini-2.0-flash", "supported_modalities": [ "text", @@ -1480,235 +968,6 @@ "supports_vision": true, "supports_web_search": true }, - "gemini-2.0-flash-live-preview-04-09": { - "cache_read_input_token_cost": 7.5e-08, - "input_cost_per_audio_token": 3e-06, - "input_cost_per_image": 3e-06, - "input_cost_per_token": 5e-07, - "input_cost_per_video_per_second": 3e-06, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1048576, - "max_output_tokens": 65535, - "max_pdf_size_mb": 30, - "max_tokens": 65535, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_audio_token": 1.2e-05, - "output_cost_per_token": 2e-06, - "rpm": 10, - "source": "https://cloud.google.com/vertex-ai/docs/generative-ai/model-reference/gemini#gemini-2-0-flash-live-preview-04-09", - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/completions" - ], - "supported_modalities": [ - "text", - "image", - "audio", - "video" - ], - "supported_output_modalities": [ - "text", - "audio" - ], - "supports_audio_output": true, - "supports_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_url_context": true, - "supports_vision": true, - "supports_web_search": true, - "tpm": 250000 - }, - "gemini-2.0-flash-preview-image-generation": { - "cache_read_input_token_cost": 2.5e-08, - "deprecation_date": "2025-11-14", - "input_cost_per_audio_token": 7e-07, - "input_cost_per_token": 1e-07, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1048576, - "max_output_tokens": 8192, - "max_pdf_size_mb": 30, - "max_tokens": 8192, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_token": 4e-07, - "source": "https://ai.google.dev/pricing#2_0flash", - "supported_modalities": [ - "text", - "image", - "audio", - "video" - ], - "supported_output_modalities": [ - "text", - "image" - ], - "supports_audio_input": true, - "supports_audio_output": true, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_prompt_caching": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true, - "supports_web_search": true - }, - "gemini-2.0-flash-thinking-exp": { - "cache_read_input_token_cost": 0.0, - "deprecation_date": "2025-12-02", - "input_cost_per_audio_per_second": 0, - "input_cost_per_audio_per_second_above_128k_tokens": 0, - "input_cost_per_character": 0, - "input_cost_per_character_above_128k_tokens": 0, - "input_cost_per_image": 0, - "input_cost_per_image_above_128k_tokens": 0, - "input_cost_per_token": 0, - "input_cost_per_token_above_128k_tokens": 0, - "input_cost_per_video_per_second": 0, - "input_cost_per_video_per_second_above_128k_tokens": 0, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1048576, - "max_output_tokens": 8192, - "max_pdf_size_mb": 30, - "max_tokens": 8192, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_character": 0, - "output_cost_per_character_above_128k_tokens": 0, - "output_cost_per_token": 0, - "output_cost_per_token_above_128k_tokens": 0, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#gemini-2.0-flash", - "supported_modalities": [ - "text", - "image", - "audio", - "video" - ], - "supported_output_modalities": [ - "text", - "image" - ], - "supports_audio_output": true, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_prompt_caching": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true, - "supports_web_search": true - }, - "gemini-2.0-flash-thinking-exp-01-21": { - "cache_read_input_token_cost": 0.0, - "deprecation_date": "2025-12-02", - "input_cost_per_audio_per_second": 0, - "input_cost_per_audio_per_second_above_128k_tokens": 0, - "input_cost_per_character": 0, - "input_cost_per_character_above_128k_tokens": 0, - "input_cost_per_image": 0, - "input_cost_per_image_above_128k_tokens": 0, - "input_cost_per_token": 0, - "input_cost_per_token_above_128k_tokens": 0, - "input_cost_per_video_per_second": 0, - "input_cost_per_video_per_second_above_128k_tokens": 0, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1048576, - "max_output_tokens": 65536, - "max_pdf_size_mb": 30, - "max_tokens": 65536, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_character": 0, - "output_cost_per_character_above_128k_tokens": 0, - "output_cost_per_token": 0, - "output_cost_per_token_above_128k_tokens": 0, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#gemini-2.0-flash", - "supported_modalities": [ - "text", - "image", - "audio", - "video" - ], - "supported_output_modalities": [ - "text", - "image" - ], - "supports_audio_output": false, - "supports_function_calling": false, - "supports_parallel_function_calling": true, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_response_schema": false, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true, - "supports_web_search": true - }, - "gemini-2.0-pro-exp-02-05": { - "cache_read_input_token_cost": 3.125e-07, - "input_cost_per_token": 1.25e-06, - "input_cost_per_token_above_200k_tokens": 2.5e-06, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 2097152, - "max_output_tokens": 8192, - "max_pdf_size_mb": 30, - "max_tokens": 8192, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_token": 1e-05, - "output_cost_per_token_above_200k_tokens": 1.5e-05, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/completions" - ], - "supported_modalities": [ - "text", - "image", - "audio", - "video" - ], - "supported_output_modalities": [ - "text" - ], - "supports_audio_input": true, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_video_input": true, - "supports_vision": true, - "supports_web_search": true - }, "gemini-2.5-computer-use-preview-10-2025": { "input_cost_per_token": 1.25e-06, "input_cost_per_token_above_200k_tokens": 2.5e-06, @@ -1751,6 +1010,11 @@ "mode": "chat", "output_cost_per_reasoning_token": 2.5e-06, "output_cost_per_token": 2.5e-06, + "search_context_cost_per_query": { + "search_context_size_high": 0.035, + "search_context_size_low": 0.035, + "search_context_size_medium": 0.035 + }, "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview", "supported_endpoints": [ "/v1/chat/completions", @@ -1773,6 +1037,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_service_tier": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_url_context": true, @@ -1821,6 +1086,7 @@ "supports_pdf_input": true, "supports_prompt_caching": true, "supports_response_schema": true, + "supports_service_tier": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_url_context": true, @@ -1828,57 +1094,6 @@ "supports_web_search": false, "tpm": 8000000 }, - "gemini-2.5-flash-image-preview": { - "cache_read_input_token_cost": 7.5e-08, - "deprecation_date": "2026-01-15", - "input_cost_per_audio_token": 1e-06, - "input_cost_per_image_token": 3e-07, - "input_cost_per_token": 3e-07, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1048576, - "max_output_tokens": 65535, - "max_pdf_size_mb": 30, - "max_tokens": 65535, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "image_generation", - "output_cost_per_image": 0.039, - "output_cost_per_image_token": 3e-05, - "output_cost_per_reasoning_token": 3e-05, - "output_cost_per_token": 3e-05, - "rpm": 100000, - "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview", - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/completions", - "/v1/batch" - ], - "supported_modalities": [ - "text", - "image", - "audio", - "video" - ], - "supported_output_modalities": [ - "text", - "image" - ], - "supports_audio_output": false, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_url_context": true, - "supports_vision": true, - "supports_web_search": true, - "tpm": 8000000 - }, "gemini-2.5-flash-lite": { "cache_read_input_token_cost": 1e-08, "input_cost_per_audio_token": 3e-07, @@ -1896,6 +1111,11 @@ "mode": "chat", "output_cost_per_reasoning_token": 4e-07, "output_cost_per_token": 4e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.035, + "search_context_size_low": 0.035, + "search_context_size_medium": 0.035 + }, "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview", "supported_endpoints": [ "/v1/chat/completions", @@ -1918,6 +1138,7 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_service_tier": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_url_context": true, @@ -1942,6 +1163,11 @@ "mode": "chat", "output_cost_per_reasoning_token": 4e-07, "output_cost_per_token": 4e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.035, + "search_context_size_low": 0.035, + "search_context_size_medium": 0.035 + }, "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview", "supported_endpoints": [ "/v1/chat/completions", @@ -1987,6 +1213,11 @@ "mode": "chat", "output_cost_per_reasoning_token": 4e-07, "output_cost_per_token": 4e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.035, + "search_context_size_low": 0.035, + "search_context_size_medium": 0.035 + }, "source": "https://developers.googleblog.com/en/continuing-to-bring-you-our-latest-models-with-an-improved-gemini-2-5-flash-and-flash-lite-release/", "supported_endpoints": [ "/v1/chat/completions", @@ -2087,96 +1318,6 @@ "supports_audio_input": true, "supports_audio_output": true }, - "gemini-2.5-flash-preview-04-17": { - "cache_read_input_token_cost": 3.75e-08, - "input_cost_per_audio_token": 1e-06, - "input_cost_per_token": 1.5e-07, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1048576, - "max_output_tokens": 65535, - "max_pdf_size_mb": 30, - "max_tokens": 65535, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_reasoning_token": 3.5e-06, - "output_cost_per_token": 6e-07, - "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview", - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/completions", - "/v1/batch" - ], - "supported_modalities": [ - "text", - "image", - "audio", - "video" - ], - "supported_output_modalities": [ - "text" - ], - "supports_audio_output": false, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true, - "supports_web_search": true - }, - "gemini-2.5-flash-preview-05-20": { - "cache_read_input_token_cost": 7.5e-08, - "deprecation_date": "2025-11-18", - "input_cost_per_audio_token": 1e-06, - "input_cost_per_token": 3e-07, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1048576, - "max_output_tokens": 65535, - "max_pdf_size_mb": 30, - "max_tokens": 65535, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_reasoning_token": 2.5e-06, - "output_cost_per_token": 2.5e-06, - "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview", - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/completions", - "/v1/batch" - ], - "supported_modalities": [ - "text", - "image", - "audio", - "video" - ], - "supported_output_modalities": [ - "text" - ], - "supports_audio_output": false, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_url_context": true, - "supports_vision": true, - "supports_web_search": true - }, "gemini-2.5-flash-preview-09-2025": { "cache_read_input_token_cost": 7.5e-08, "input_cost_per_audio_token": 1e-06, @@ -2194,6 +1335,11 @@ "mode": "chat", "output_cost_per_reasoning_token": 2.5e-06, "output_cost_per_token": 2.5e-06, + "search_context_cost_per_query": { + "search_context_size_high": 0.035, + "search_context_size_low": 0.035, + "search_context_size_medium": 0.035 + }, "source": "https://developers.googleblog.com/en/continuing-to-bring-you-our-latest-models-with-an-improved-gemini-2-5-flash-and-flash-lite-release/", "supported_endpoints": [ "/v1/chat/completions", @@ -2251,6 +1397,11 @@ "mode": "chat", "output_cost_per_token": 1e-05, "output_cost_per_token_above_200k_tokens": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.035, + "search_context_size_low": 0.035, + "search_context_size_medium": 0.035 + }, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", @@ -2271,199 +1422,13 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, + "supports_service_tier": true, "supports_system_messages": true, "supports_tool_choice": true, "supports_video_input": true, "supports_vision": true, "supports_web_search": true }, - "gemini-2.5-pro-exp-03-25": { - "cache_read_input_token_cost": 1.25e-07, - "cache_read_input_token_cost_above_200k_tokens": 2.5e-07, - "input_cost_per_token": 1.25e-06, - "input_cost_per_token_above_200k_tokens": 2.5e-06, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1048576, - "max_output_tokens": 65535, - "max_pdf_size_mb": 30, - "max_tokens": 65535, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_token": 1e-05, - "output_cost_per_token_above_200k_tokens": 1.5e-05, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/completions" - ], - "supported_modalities": [ - "text", - "image", - "audio", - "video" - ], - "supported_output_modalities": [ - "text" - ], - "supports_audio_input": true, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_video_input": true, - "supports_vision": true, - "supports_web_search": true - }, - "gemini-2.5-pro-preview-03-25": { - "cache_read_input_token_cost": 1.25e-07, - "cache_read_input_token_cost_above_200k_tokens": 2.5e-07, - "deprecation_date": "2025-12-02", - "input_cost_per_audio_token": 1.25e-06, - "input_cost_per_token": 1.25e-06, - "input_cost_per_token_above_200k_tokens": 2.5e-06, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1048576, - "max_output_tokens": 65535, - "max_pdf_size_mb": 30, - "max_tokens": 65535, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_token": 1e-05, - "output_cost_per_token_above_200k_tokens": 1.5e-05, - "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview", - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/completions", - "/v1/batch" - ], - "supported_modalities": [ - "text", - "image", - "audio", - "video" - ], - "supported_output_modalities": [ - "text" - ], - "supports_audio_output": false, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true, - "supports_web_search": true - }, - "gemini-2.5-pro-preview-05-06": { - "cache_read_input_token_cost": 1.25e-07, - "cache_read_input_token_cost_above_200k_tokens": 2.5e-07, - "deprecation_date": "2025-12-02", - "input_cost_per_audio_token": 1.25e-06, - "input_cost_per_token": 1.25e-06, - "input_cost_per_token_above_200k_tokens": 2.5e-06, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1048576, - "max_output_tokens": 65535, - "max_pdf_size_mb": 30, - "max_tokens": 65535, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_token": 1e-05, - "output_cost_per_token_above_200k_tokens": 1.5e-05, - "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview", - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/completions", - "/v1/batch" - ], - "supported_modalities": [ - "text", - "image", - "audio", - "video" - ], - "supported_output_modalities": [ - "text" - ], - "supported_regions": [ - "global" - ], - "supports_audio_output": false, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true, - "supports_web_search": true - }, - "gemini-2.5-pro-preview-06-05": { - "cache_read_input_token_cost": 1.25e-07, - "cache_read_input_token_cost_above_200k_tokens": 2.5e-07, - "input_cost_per_audio_token": 1.25e-06, - "input_cost_per_token": 1.25e-06, - "input_cost_per_token_above_200k_tokens": 2.5e-06, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1048576, - "max_output_tokens": 65535, - "max_pdf_size_mb": 30, - "max_tokens": 65535, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_token": 1e-05, - "output_cost_per_token_above_200k_tokens": 1.5e-05, - "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview", - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/completions", - "/v1/batch" - ], - "supported_modalities": [ - "text", - "image", - "audio", - "video" - ], - "supported_output_modalities": [ - "text" - ], - "supports_audio_output": false, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true, - "supports_web_search": true - }, "gemini-2.5-pro-preview-tts": { "cache_read_input_token_cost": 1.25e-07, "cache_read_input_token_cost_above_200k_tokens": 2.5e-07, @@ -2483,6 +1448,11 @@ "mode": "chat", "output_cost_per_token": 1e-05, "output_cost_per_token_above_200k_tokens": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.035, + "search_context_size_low": 0.035, + "search_context_size_medium": 0.035 + }, "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-2.5-pro-preview", "supported_modalities": [ "text" @@ -2500,7 +1470,7 @@ "supports_vision": true, "supports_web_search": true }, - "gemini-3-flash-preview": { + "gemini-3-flash": { "cache_read_input_token_cost": 5e-08, "cache_read_input_token_cost_priority": 9e-08, "input_cost_per_audio_token": 1e-06, @@ -2521,6 +1491,11 @@ "output_cost_per_reasoning_token": 3e-06, "output_cost_per_token": 3e-06, "output_cost_per_token_priority": 5.4e-06, + "search_context_cost_per_query": { + "search_context_size_high": 0.014, + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014 + }, "source": "https://ai.google.dev/pricing/gemini-3", "supported_endpoints": [ "/v1/chat/completions", @@ -2549,7 +1524,65 @@ "supports_tool_choice": true, "supports_url_context": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "web_search_billing_unit": "per_query" + }, + "gemini-3-flash-preview": { + "cache_read_input_token_cost": 5e-08, + "cache_read_input_token_cost_priority": 9e-08, + "input_cost_per_audio_token": 1e-06, + "input_cost_per_audio_token_priority": 1.8e-06, + "input_cost_per_token": 5e-07, + "input_cost_per_token_priority": 9e-07, + "litellm_provider": "vertex_ai-language-models", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 1048576, + "max_output_tokens": 65535, + "max_pdf_size_mb": 30, + "max_tokens": 65535, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "chat", + "output_cost_per_reasoning_token": 3e-06, + "output_cost_per_token": 3e-06, + "output_cost_per_token_priority": 5.4e-06, + "search_context_cost_per_query": { + "search_context_size_high": 0.014, + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014 + }, + "source": "https://ai.google.dev/pricing/gemini-3", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_output": false, + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_vision": true, + "supports_web_search": true, + "web_search_billing_unit": "per_query" }, "gemini-3-pro-image-preview": { "input_cost_per_image": 0.0011, @@ -2564,6 +1597,11 @@ "output_cost_per_image_token": 0.00012, "output_cost_per_token": 1.2e-05, "output_cost_per_token_batches": 6e-06, + "search_context_cost_per_query": { + "search_context_size_high": 0.014, + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014 + }, "source": "https://ai.google.dev/gemini-api/docs/pricing", "supported_endpoints": [ "/v1/chat/completions", @@ -2581,9 +1619,11 @@ "supports_function_calling": false, "supports_prompt_caching": true, "supports_response_schema": true, + "supports_service_tier": true, "supports_system_messages": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "web_search_billing_unit": "per_query" }, "gemini-3-pro-preview": { "cache_creation_input_token_cost_above_200k_tokens": 2.5e-07, @@ -2591,6 +1631,7 @@ "cache_read_input_token_cost_above_200k_tokens": 4e-07, "cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07, "cache_read_input_token_cost_priority": 3.6e-07, + "deprecation_date": "2026-03-26", "input_cost_per_token": 2e-06, "input_cost_per_token_above_200k_tokens": 4e-06, "input_cost_per_token_above_200k_tokens_priority": 7.2e-06, @@ -2612,6 +1653,11 @@ "output_cost_per_token_above_200k_tokens_priority": 3.24e-05, "output_cost_per_token_batches": 6e-06, "output_cost_per_token_priority": 2.16e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.014, + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014 + }, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", @@ -2639,6 +1685,240 @@ "supports_tool_choice": true, "supports_video_input": true, "supports_vision": true, + "supports_web_search": true, + "web_search_billing_unit": "per_query" + }, + "gemini-3.1-flash-image": { + "input_cost_per_image": 0.00056, + "input_cost_per_token": 5e-07, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "image_generation", + "output_cost_per_image": 0.0672, + "output_cost_per_image_token": 6e-05, + "output_cost_per_token": 3e-06, + "search_context_cost_per_query": { + "search_context_size_high": 0.014, + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014 + }, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text", + "image" + ], + "supports_function_calling": false, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_web_search": true, + "web_search_billing_unit": "per_query" + }, + "gemini-3.1-flash-image-preview": { + "input_cost_per_image": 0.00056, + "input_cost_per_token": 5e-07, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "image_generation", + "output_cost_per_image": 0.0672, + "output_cost_per_image_token": 6e-05, + "output_cost_per_token": 3e-06, + "search_context_cost_per_query": { + "search_context_size_high": 0.014, + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014 + }, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text", + "image" + ], + "supports_function_calling": false, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_web_search": true, + "web_search_billing_unit": "per_query" + }, + "gemini-3.1-flash-lite": { + "cache_read_input_token_cost": 2.5e-08, + "cache_read_input_token_cost_batches": 1.25e-08, + "cache_read_input_token_cost_flex": 1.25e-08, + "cache_read_input_token_cost_per_audio_token": 5e-08, + "cache_read_input_token_cost_priority": 4.5e-08, + "input_cost_per_audio_token": 5e-07, + "input_cost_per_token": 2.5e-07, + "input_cost_per_token_batches": 1.25e-07, + "input_cost_per_token_flex": 1.25e-07, + "input_cost_per_token_priority": 4.5e-07, + "litellm_provider": "vertex_ai-language-models", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_pdf_size_mb": 30, + "max_tokens": 65536, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "chat", + "output_cost_per_reasoning_token": 1.5e-06, + "output_cost_per_token": 1.5e-06, + "output_cost_per_token_batches": 7.5e-07, + "output_cost_per_token_flex": 7.5e-07, + "output_cost_per_token_priority": 2.7e-06, + "search_context_cost_per_query": { + "search_context_size_high": 0.014, + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014 + }, + "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3.1-flash-lite", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_audio_output": false, + "supports_code_execution": true, + "supports_file_search": true, + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "web_search_billing_unit": "per_query" + }, + "gemini-3.1-flash-lite-preview": { + "cache_read_input_token_cost": 2.5e-08, + "cache_read_input_token_cost_per_audio_token": 5e-08, + "input_cost_per_audio_token": 5e-07, + "input_cost_per_token": 2.5e-07, + "litellm_provider": "vertex_ai-language-models", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_pdf_size_mb": 30, + "max_tokens": 65536, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "chat", + "output_cost_per_reasoning_token": 1.5e-06, + "output_cost_per_token": 1.5e-06, + "search_context_cost_per_query": { + "search_context_size_high": 0.014, + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014 + }, + "source": "https://ai.google.dev/gemini-api/docs/models", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_audio_output": false, + "supports_code_execution": true, + "supports_file_search": true, + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "web_search_billing_unit": "per_query" + }, + "gemini-3.1-flash-live-preview": { + "input_cost_per_audio_token": 3e-06, + "input_cost_per_image_token": 1e-06, + "input_cost_per_token": 7.5e-07, + "input_cost_per_video_per_second": 3.3333333333333335e-05, + "litellm_provider": "gemini", + "max_input_tokens": 131072, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_audio_token": 1.2e-05, + "output_cost_per_token": 4.5e-06, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/realtime" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text", + "audio" + ], + "supports_audio_input": true, + "supports_audio_output": true, + "supports_function_calling": true, + "supports_vision": true, "supports_web_search": true }, "gemini-3.1-pro-high": { @@ -2669,6 +1949,11 @@ "output_cost_per_token_above_200k_tokens_priority": 3.24e-05, "output_cost_per_token_batches": 6e-06, "output_cost_per_token_priority": 2.16e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.014, + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014 + }, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", "supported_endpoints": [ "/v1/chat/completions", @@ -2697,7 +1982,8 @@ "supports_url_context": true, "supports_video_input": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "web_search_billing_unit": "per_query" }, "gemini-3.1-pro-low": { "cache_creation_input_token_cost_above_200k_tokens": 2.5e-07, @@ -2727,6 +2013,11 @@ "output_cost_per_token_above_200k_tokens_priority": 3.24e-05, "output_cost_per_token_batches": 6e-06, "output_cost_per_token_priority": 2.16e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.014, + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014 + }, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", "supported_endpoints": [ "/v1/chat/completions", @@ -2755,7 +2046,8 @@ "supports_url_context": true, "supports_video_input": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "web_search_billing_unit": "per_query" }, "gemini-3.1-pro-preview": { "cache_creation_input_token_cost_above_200k_tokens": 2.5e-07, @@ -2785,6 +2077,11 @@ "output_cost_per_token_above_200k_tokens_priority": 3.24e-05, "output_cost_per_token_batches": 6e-06, "output_cost_per_token_priority": 2.16e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.014, + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014 + }, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", "supported_endpoints": [ "/v1/chat/completions", @@ -2813,7 +2110,8 @@ "supports_url_context": true, "supports_video_input": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "web_search_billing_unit": "per_query" }, "gemini-3.1-pro-preview-customtools": { "cache_creation_input_token_cost_above_200k_tokens": 2.5e-07, @@ -2837,6 +2135,11 @@ "output_cost_per_token": 1.2e-05, "output_cost_per_token_above_200k_tokens": 1.8e-05, "output_cost_per_token_batches": 6e-06, + "search_context_cost_per_query": { + "search_context_size_high": 0.014, + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014 + }, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", "supported_endpoints": [ "/v1/chat/completions", @@ -2864,7 +2167,67 @@ "supports_url_context": true, "supports_video_input": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "web_search_billing_unit": "per_query" + }, + "gemini-3.5-flash": { + "cache_read_input_token_cost": 1.5e-07, + "cache_read_input_token_cost_priority": 2.7e-07, + "input_cost_per_audio_token": 1e-06, + "input_cost_per_audio_token_priority": 1.8e-06, + "input_cost_per_token": 1.5e-06, + "input_cost_per_token_priority": 2.7e-06, + "litellm_provider": "vertex_ai-language-models", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 1048576, + "max_output_tokens": 65535, + "max_pdf_size_mb": 30, + "max_tokens": 65535, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "chat", + "output_cost_per_reasoning_token": 9e-06, + "output_cost_per_token": 9e-06, + "output_cost_per_token_priority": 1.62e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.014, + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014 + }, + "source": "https://ai.google.dev/pricing/gemini-3", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_audio_output": false, + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "web_search_billing_unit": "per_query" }, "gemini-embedding-001": { "input_cost_per_token": 1.5e-07, @@ -2876,6 +2239,35 @@ "output_vector_size": 3072, "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models" }, + "gemini-embedding-2": { + "input_cost_per_audio_per_second": 0.00016, + "input_cost_per_image": 0.00012, + "input_cost_per_token": 2e-07, + "input_cost_per_video_per_second": 0.00079, + "litellm_provider": "vertex_ai-embedding-models", + "max_input_tokens": 8192, + "max_tokens": 8192, + "mode": "embedding", + "output_cost_per_token": 0, + "output_vector_size": 3072, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "supports_multimodal": true, + "uses_embed_content": true + }, + "gemini-embedding-2-preview": { + "input_cost_per_audio_per_second": 0.00016, + "input_cost_per_image": 0.00012, + "input_cost_per_token": 2e-07, + "input_cost_per_video_per_second": 0.00079, + "litellm_provider": "vertex_ai-embedding-models", + "max_input_tokens": 8192, + "max_tokens": 8192, + "mode": "embedding", + "output_cost_per_token": 0, + "output_vector_size": 3072, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "uses_embed_content": true + }, "gemini-exp-1206": { "cache_read_input_token_cost": 3e-08, "input_cost_per_audio_token": 1e-06, @@ -2894,6 +2286,11 @@ "output_cost_per_reasoning_token": 2.5e-06, "output_cost_per_token": 2.5e-06, "rpm": 100000, + "search_context_cost_per_query": { + "search_context_size_high": 0.035, + "search_context_size_low": 0.035, + "search_context_size_medium": 0.035 + }, "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview", "supported_endpoints": [ "/v1/chat/completions", @@ -2930,13 +2327,11 @@ "max_input_tokens": 1000000, "max_output_tokens": 8192, "max_tokens": 8192, - "mode": "chat", - "output_cost_per_character": 0, + "mode": "embedding", "output_cost_per_token": 0, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/multimodal/gemini-experimental", - "supports_function_calling": false, - "supports_parallel_function_calling": true, - "supports_tool_choice": true + "output_vector_size": 3072, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "uses_embed_content": true }, "gemini-flash-latest": { "cache_read_input_token_cost": 3e-08, @@ -2956,6 +2351,11 @@ "output_cost_per_reasoning_token": 2.5e-06, "output_cost_per_token": 2.5e-06, "rpm": 100000, + "search_context_cost_per_query": { + "search_context_size_high": 0.035, + "search_context_size_low": 0.035, + "search_context_size_medium": 0.035 + }, "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview", "supported_endpoints": [ "/v1/chat/completions", @@ -3003,6 +2403,11 @@ "output_cost_per_reasoning_token": 4e-07, "output_cost_per_token": 4e-07, "rpm": 15, + "search_context_cost_per_query": { + "search_context_size_high": 0.035, + "search_context_size_low": 0.035, + "search_context_size_medium": 0.035 + }, "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-lite", "supported_endpoints": [ "/v1/chat/completions", @@ -3046,13 +2451,17 @@ "max_tokens": 65535, "max_video_length": 1, "max_videos_per_prompt": 10, - "mode": "chat", + "mode": "realtime", "output_cost_per_audio_token": 1.2e-05, "output_cost_per_token": 2e-06, + "search_context_cost_per_query": { + "search_context_size_high": 0.035, + "search_context_size_low": 0.035, + "search_context_size_medium": 0.035 + }, "source": "https://ai.google.dev/gemini-api/docs/pricing", "supported_endpoints": [ - "/v1/chat/completions", - "/v1/completions" + "/vertex_ai/live" ], "supported_modalities": [ "text", @@ -3077,38 +2486,6 @@ "supports_vision": true, "supports_web_search": true }, - "gemini-pro": { - "input_cost_per_character": 1.25e-07, - "input_cost_per_image": 0.0025, - "input_cost_per_token": 5e-07, - "input_cost_per_video_per_second": 0.002, - "litellm_provider": "vertex_ai-language-models", - "max_input_tokens": 32760, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_character": 3.75e-07, - "output_cost_per_token": 1.5e-06, - "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_tool_choice": true - }, - "gemini-pro-experimental": { - "input_cost_per_character": 0, - "input_cost_per_token": 0, - "litellm_provider": "vertex_ai-language-models", - "max_input_tokens": 1000000, - "max_output_tokens": 8192, - "max_tokens": 8192, - "mode": "chat", - "output_cost_per_character": 0, - "output_cost_per_token": 0, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/multimodal/gemini-experimental", - "supports_function_calling": false, - "supports_parallel_function_calling": true, - "supports_tool_choice": true - }, "gemini-pro-latest": { "cache_read_input_token_cost": 1.25e-07, "cache_read_input_token_cost_above_200k_tokens": 2.5e-07, @@ -3128,6 +2505,11 @@ "output_cost_per_token": 1e-05, "output_cost_per_token_above_200k_tokens": 1.5e-05, "rpm": 2000, + "search_context_cost_per_query": { + "search_context_size_high": 0.035, + "search_context_size_low": 0.035, + "search_context_size_medium": 0.035 + }, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", "supported_endpoints": [ "/v1/chat/completions", @@ -3155,24 +2537,6 @@ "supports_web_search": true, "tpm": 800000 }, - "gemini-pro-vision": { - "input_cost_per_image": 0.0025, - "input_cost_per_token": 5e-07, - "litellm_provider": "vertex_ai-vision-models", - "max_images_per_prompt": 16, - "max_input_tokens": 16384, - "max_output_tokens": 2048, - "max_tokens": 2048, - "max_video_length": 2, - "max_videos_per_prompt": 1, - "mode": "chat", - "output_cost_per_token": 1.5e-06, - "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models", - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_tool_choice": true, - "supports_vision": true - }, "gemini-robotics-er-1.5-preview": { "cache_read_input_token_cost": 0, "input_cost_per_audio_token": 1e-06, @@ -3236,31 +2600,6 @@ "supports_system_messages": true, "supports_tool_choice": true }, - "gpt-3.5-turbo-0301": { - "input_cost_per_token": 1.5e-06, - "litellm_provider": "openai", - "max_input_tokens": 4097, - "max_output_tokens": 4096, - "max_tokens": 4096, - "mode": "chat", - "output_cost_per_token": 2e-06, - "supports_prompt_caching": true, - "supports_system_messages": true, - "supports_tool_choice": true - }, - "gpt-3.5-turbo-0613": { - "input_cost_per_token": 1.5e-06, - "litellm_provider": "openai", - "max_input_tokens": 4097, - "max_output_tokens": 4096, - "max_tokens": 4096, - "mode": "chat", - "output_cost_per_token": 2e-06, - "supports_function_calling": true, - "supports_prompt_caching": true, - "supports_system_messages": true, - "supports_tool_choice": true - }, "gpt-3.5-turbo-1106": { "deprecation_date": "2026-09-28", "input_cost_per_token": 1e-06, @@ -3288,18 +2627,6 @@ "supports_system_messages": true, "supports_tool_choice": true }, - "gpt-3.5-turbo-16k-0613": { - "input_cost_per_token": 3e-06, - "litellm_provider": "openai", - "max_input_tokens": 16385, - "max_output_tokens": 4096, - "max_tokens": 4096, - "mode": "chat", - "output_cost_per_token": 4e-06, - "supports_prompt_caching": true, - "supports_system_messages": true, - "supports_tool_choice": true - }, "gpt-3.5-turbo-instruct": { "input_cost_per_token": 1.5e-06, "litellm_provider": "text-completion-openai", @@ -3347,6 +2674,7 @@ "supports_tool_choice": true }, "gpt-4-0314": { + "deprecation_date": "2026-03-26", "input_cost_per_token": 3e-05, "litellm_provider": "openai", "max_input_tokens": 8192, @@ -3354,7 +2682,6 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-05, - "supports_prompt_caching": true, "supports_system_messages": true, "supports_tool_choice": true }, @@ -3387,57 +2714,6 @@ "supports_system_messages": true, "supports_tool_choice": true }, - "gpt-4-1106-vision-preview": { - "deprecation_date": "2024-12-06", - "input_cost_per_token": 1e-05, - "litellm_provider": "openai", - "max_input_tokens": 128000, - "max_output_tokens": 4096, - "max_tokens": 4096, - "mode": "chat", - "output_cost_per_token": 3e-05, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true - }, - "gpt-4-32k": { - "input_cost_per_token": 6e-05, - "litellm_provider": "openai", - "max_input_tokens": 32768, - "max_output_tokens": 4096, - "max_tokens": 4096, - "mode": "chat", - "output_cost_per_token": 0.00012, - "supports_prompt_caching": true, - "supports_system_messages": true, - "supports_tool_choice": true - }, - "gpt-4-32k-0314": { - "input_cost_per_token": 6e-05, - "litellm_provider": "openai", - "max_input_tokens": 32768, - "max_output_tokens": 4096, - "max_tokens": 4096, - "mode": "chat", - "output_cost_per_token": 0.00012, - "supports_prompt_caching": true, - "supports_system_messages": true, - "supports_tool_choice": true - }, - "gpt-4-32k-0613": { - "input_cost_per_token": 6e-05, - "litellm_provider": "openai", - "max_input_tokens": 32768, - "max_output_tokens": 4096, - "max_tokens": 4096, - "mode": "chat", - "output_cost_per_token": 0.00012, - "supports_prompt_caching": true, - "supports_system_messages": true, - "supports_tool_choice": true - }, "gpt-4-turbo": { "input_cost_per_token": 1e-05, "litellm_provider": "openai", @@ -3485,21 +2761,6 @@ "supports_system_messages": true, "supports_tool_choice": true }, - "gpt-4-vision-preview": { - "deprecation_date": "2024-12-06", - "input_cost_per_token": 1e-05, - "litellm_provider": "openai", - "max_input_tokens": 128000, - "max_output_tokens": 4096, - "max_tokens": 4096, - "mode": "chat", - "output_cost_per_token": 3e-05, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true - }, "gpt-4.1": { "cache_read_input_token_cost": 5e-07, "cache_read_input_token_cost_priority": 8.75e-07, @@ -3535,7 +2796,8 @@ "supports_service_tier": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true }, "gpt-4.1-2025-04-14": { "cache_read_input_token_cost": 5e-07, @@ -3569,7 +2831,8 @@ "supports_service_tier": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true }, "gpt-4.1-mini": { "cache_read_input_token_cost": 1e-07, @@ -3606,7 +2869,8 @@ "supports_service_tier": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true }, "gpt-4.1-mini-2025-04-14": { "cache_read_input_token_cost": 1e-07, @@ -3640,7 +2904,8 @@ "supports_service_tier": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true }, "gpt-4.1-nano": { "cache_read_input_token_cost": 2.5e-08, @@ -3713,47 +2978,6 @@ "supports_tool_choice": true, "supports_vision": true }, - "gpt-4.5-preview": { - "cache_read_input_token_cost": 3.75e-05, - "input_cost_per_token": 7.5e-05, - "input_cost_per_token_batches": 3.75e-05, - "litellm_provider": "openai", - "max_input_tokens": 128000, - "max_output_tokens": 16384, - "max_tokens": 16384, - "mode": "chat", - "output_cost_per_token": 0.00015, - "output_cost_per_token_batches": 7.5e-05, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true - }, - "gpt-4.5-preview-2025-02-27": { - "cache_read_input_token_cost": 3.75e-05, - "deprecation_date": "2025-07-14", - "input_cost_per_token": 7.5e-05, - "input_cost_per_token_batches": 3.75e-05, - "litellm_provider": "openai", - "max_input_tokens": 128000, - "max_output_tokens": 16384, - "max_tokens": 16384, - "mode": "chat", - "output_cost_per_token": 0.00015, - "output_cost_per_token_batches": 7.5e-05, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true - }, "gpt-4o": { "cache_read_input_token_cost": 1.25e-06, "cache_read_input_token_cost_priority": 2.125e-06, @@ -3857,23 +3081,6 @@ "supports_system_messages": true, "supports_tool_choice": true }, - "gpt-4o-audio-preview-2024-10-01": { - "input_cost_per_audio_token": 4e-05, - "input_cost_per_token": 2.5e-06, - "litellm_provider": "openai", - "max_input_tokens": 128000, - "max_output_tokens": 16384, - "max_tokens": 16384, - "mode": "chat", - "output_cost_per_audio_token": 8e-05, - "output_cost_per_token": 1e-05, - "supports_audio_input": true, - "supports_audio_output": true, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_system_messages": true, - "supports_tool_choice": true - }, "gpt-4o-audio-preview-2024-12-17": { "input_cost_per_audio_token": 4e-05, "input_cost_per_token": 2.5e-06, @@ -4077,7 +3284,7 @@ "supports_vision": true }, "gpt-4o-mini-transcribe": { - "input_cost_per_audio_token": 3e-06, + "input_cost_per_audio_token": 1.25e-06, "input_cost_per_token": 1.25e-06, "litellm_provider": "openai", "max_input_tokens": 16000, @@ -4089,7 +3296,7 @@ ] }, "gpt-4o-mini-transcribe-2025-03-20": { - "input_cost_per_audio_token": 3e-06, + "input_cost_per_audio_token": 1.25e-06, "input_cost_per_token": 1.25e-06, "litellm_provider": "openai", "max_input_tokens": 16000, @@ -4101,7 +3308,7 @@ ] }, "gpt-4o-mini-transcribe-2025-12-15": { - "input_cost_per_audio_token": 3e-06, + "input_cost_per_audio_token": 1.25e-06, "input_cost_per_token": 1.25e-06, "litellm_provider": "openai", "max_input_tokens": 16000, @@ -4184,25 +3391,6 @@ "supports_system_messages": true, "supports_tool_choice": true }, - "gpt-4o-realtime-preview-2024-10-01": { - "cache_creation_input_audio_token_cost": 2e-05, - "cache_read_input_token_cost": 2.5e-06, - "input_cost_per_audio_token": 0.0001, - "input_cost_per_token": 5e-06, - "litellm_provider": "openai", - "max_input_tokens": 128000, - "max_output_tokens": 4096, - "max_tokens": 4096, - "mode": "chat", - "output_cost_per_audio_token": 0.0002, - "output_cost_per_token": 2e-05, - "supports_audio_input": true, - "supports_audio_output": true, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_system_messages": true, - "supports_tool_choice": true - }, "gpt-4o-realtime-preview-2024-12-17": { "cache_read_input_token_cost": 2.5e-06, "input_cost_per_audio_token": 4e-05, @@ -4286,7 +3474,7 @@ "supports_vision": true }, "gpt-4o-transcribe": { - "input_cost_per_audio_token": 6e-06, + "input_cost_per_audio_token": 2.5e-06, "input_cost_per_token": 2.5e-06, "litellm_provider": "openai", "max_input_tokens": 16000, @@ -4298,7 +3486,7 @@ ] }, "gpt-4o-transcribe-diarize": { - "input_cost_per_audio_token": 6e-06, + "input_cost_per_audio_token": 2.5e-06, "input_cost_per_token": 2.5e-06, "litellm_provider": "openai", "max_input_tokens": 16000, @@ -4337,7 +3525,9 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4346,7 +3536,9 @@ "supports_service_tier": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5-2025-08-07": { "cache_read_input_token_cost": 1.25e-07, @@ -4376,7 +3568,9 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4385,7 +3579,9 @@ "supports_service_tier": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5-chat": { "cache_read_input_token_cost": 1.25e-07, @@ -4409,7 +3605,9 @@ "text" ], "supports_function_calling": false, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": false, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4417,7 +3615,8 @@ "supports_response_schema": true, "supports_system_messages": true, "supports_tool_choice": false, - "supports_vision": true + "supports_vision": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5-chat-latest": { "cache_read_input_token_cost": 1.25e-07, @@ -4441,7 +3640,9 @@ "text" ], "supports_function_calling": false, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": false, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4449,7 +3650,8 @@ "supports_response_schema": true, "supports_system_messages": true, "supports_tool_choice": false, - "supports_vision": true + "supports_vision": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5-codex": { "cache_read_input_token_cost": 1.25e-07, @@ -4471,7 +3673,9 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4479,7 +3683,9 @@ "supports_response_schema": true, "supports_system_messages": false, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5-mini": { "cache_read_input_token_cost": 2.5e-08, @@ -4509,7 +3715,9 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4518,7 +3726,9 @@ "supports_service_tier": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5-mini-2025-08-07": { "cache_read_input_token_cost": 2.5e-08, @@ -4548,7 +3758,9 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4557,7 +3769,9 @@ "supports_service_tier": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5-nano": { "cache_read_input_token_cost": 5e-09, @@ -4585,7 +3799,9 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4593,7 +3809,9 @@ "supports_response_schema": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5-nano-2025-08-07": { "cache_read_input_token_cost": 5e-09, @@ -4620,7 +3838,9 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4628,7 +3848,9 @@ "supports_response_schema": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5-pro": { "input_cost_per_token": 1.5e-05, @@ -4652,7 +3874,9 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": false, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4661,7 +3885,8 @@ "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5-pro-2025-10-06": { "input_cost_per_token": 1.5e-05, @@ -4685,7 +3910,9 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": false, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4694,7 +3921,8 @@ "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5-search-api": { "cache_read_input_token_cost": 1.25e-07, @@ -4706,6 +3934,8 @@ "mode": "chat", "output_cost_per_token": 1e-05, "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4713,7 +3943,8 @@ "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5-search-api-2025-10-14": { "cache_read_input_token_cost": 1.25e-07, @@ -4725,6 +3956,7 @@ "mode": "chat", "output_cost_per_token": 1e-05, "supports_function_calling": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4732,7 +3964,8 @@ "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5.1": { "cache_read_input_token_cost": 1.25e-07, @@ -4759,7 +3992,9 @@ "image" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4768,7 +4003,9 @@ "supports_service_tier": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5.1-2025-11-13": { "cache_read_input_token_cost": 1.25e-07, @@ -4795,7 +4032,9 @@ "image" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4804,7 +4043,9 @@ "supports_service_tier": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5.1-chat-latest": { "cache_read_input_token_cost": 1.25e-07, @@ -4831,7 +4072,9 @@ "image" ], "supports_function_calling": false, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": true, "supports_parallel_function_calling": false, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4839,7 +4082,9 @@ "supports_response_schema": true, "supports_system_messages": true, "supports_tool_choice": false, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5.1-codex": { "cache_read_input_token_cost": 1.25e-07, @@ -4864,7 +4109,9 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4872,7 +4119,9 @@ "supports_response_schema": true, "supports_system_messages": false, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5.1-codex-max": { "cache_read_input_token_cost": 1.25e-07, @@ -4894,7 +4143,9 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4902,7 +4153,9 @@ "supports_response_schema": true, "supports_system_messages": false, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true }, "gpt-5.1-codex-mini": { "cache_read_input_token_cost": 2.5e-08, @@ -4927,7 +4180,9 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4935,7 +4190,9 @@ "supports_response_schema": true, "supports_system_messages": false, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5.2": { "cache_read_input_token_cost": 1.75e-07, @@ -4963,7 +4220,9 @@ "image" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -4972,7 +4231,9 @@ "supports_service_tier": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true }, "gpt-5.2-2025-12-11": { "cache_read_input_token_cost": 1.75e-07, @@ -5000,7 +4261,9 @@ "image" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": true, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -5009,7 +4272,9 @@ "supports_service_tier": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true }, "gpt-5.2-chat-latest": { "cache_read_input_token_cost": 1.75e-07, @@ -5035,7 +4300,9 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -5043,7 +4310,9 @@ "supports_response_schema": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5.2-codex": { "cache_read_input_token_cost": 1.75e-07, @@ -5068,7 +4337,9 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -5076,7 +4347,9 @@ "supports_response_schema": true, "supports_system_messages": false, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true }, "gpt-5.2-pro": { "input_cost_per_token": 2.1e-05, @@ -5098,7 +4371,9 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -5107,7 +4382,8 @@ "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true }, "gpt-5.2-pro-2025-12-11": { "input_cost_per_token": 2.1e-05, @@ -5129,7 +4405,9 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -5138,17 +4416,21 @@ "supports_system_messages": true, "supports_tool_choice": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true }, - "gpt-5.4": { - "cache_read_input_token_cost": 2.5e-07, - "input_cost_per_token": 2.5e-06, + "gpt-5.3-chat-latest": { + "cache_read_input_token_cost": 1.75e-07, + "cache_read_input_token_cost_priority": 3.5e-07, + "input_cost_per_token": 1.75e-06, + "input_cost_per_token_priority": 3.5e-06, "litellm_provider": "openai", - "max_input_tokens": 1050000, - "max_output_tokens": 128000, - "max_tokens": 128000, + "max_input_tokens": 128000, + "max_output_tokens": 16384, + "max_tokens": 16384, "mode": "chat", - "output_cost_per_token": 1.5e-05, + "output_cost_per_token": 1.4e-05, + "output_cost_per_token_priority": 2.8e-05, "supported_endpoints": [ "/v1/chat/completions", "/v1/responses" @@ -5157,111 +4439,13 @@ "text", "image" ], - "supported_output_modalities": [ - "text", - "image" - ], - "supports_function_calling": true, - "supports_native_streaming": true, - "supports_parallel_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_response_schema": true, - "supports_service_tier": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true - }, - "codex-auto-review": { - "cache_read_input_token_cost": 2.5e-07, - "input_cost_per_token": 2.5e-06, - "litellm_provider": "openai", - "max_input_tokens": 1050000, - "max_output_tokens": 128000, - "max_tokens": 128000, - "mode": "chat", - "output_cost_per_token": 1.5e-05, - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/responses" - ], - "supported_modalities": [ - "text", - "image" - ], - "supported_output_modalities": [ - "text", - "image" - ], - "supports_function_calling": true, - "supports_native_streaming": true, - "supports_parallel_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_response_schema": true, - "supports_service_tier": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true - }, - "gpt-5.4-mini": { - "cache_read_input_token_cost": 7.5e-08, - "input_cost_per_token": 7.5e-07, - "litellm_provider": "openai", - "max_input_tokens": 400000, - "max_output_tokens": 128000, - "max_tokens": 128000, - "mode": "chat", - "output_cost_per_token": 4.5e-06, - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/batch", - "/v1/responses" - ], - "supported_modalities": [ - "text", - "image" - ], - "supported_output_modalities": [ - "text" - ], - "supports_function_calling": true, - "supports_native_streaming": true, - "supports_parallel_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_response_schema": true, - "supports_service_tier": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_vision": true - }, - "gpt-5.4-nano": { - "cache_read_input_token_cost": 2e-08, - "input_cost_per_token": 2e-07, - "litellm_provider": "openai", - "max_input_tokens": 400000, - "max_output_tokens": 128000, - "max_tokens": 128000, - "mode": "chat", - "output_cost_per_token": 1.25e-06, - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/batch", - "/v1/responses" - ], - "supported_modalities": [ - "text", - "image" - ], "supported_output_modalities": [ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -5269,7 +4453,9 @@ "supports_response_schema": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false }, "gpt-5.3-codex": { "cache_read_input_token_cost": 1.75e-07, @@ -5294,7 +4480,9 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, "supports_native_streaming": true, + "supports_none_reasoning_effort": false, "supports_parallel_function_calling": true, "supports_pdf_input": true, "supports_prompt_caching": true, @@ -5302,8 +4490,586 @@ "supports_response_schema": true, "supports_system_messages": false, "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false + }, + "gpt-5.3-codex-spark": { + "cache_read_input_token_cost": 1.75e-07, + "cache_read_input_token_cost_priority": 3.5e-07, + "input_cost_per_token": 1.75e-06, + "input_cost_per_token_priority": 3.5e-06, + "litellm_provider": "openai", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "responses", + "output_cost_per_token": 1.4e-05, + "output_cost_per_token_priority": 2.8e-05, + "supported_endpoints": [ + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, + "supports_native_streaming": true, + "supports_none_reasoning_effort": false, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": false, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": false + }, + "gpt-5.4": { + "cache_read_input_token_cost": 2.5e-07, + "cache_read_input_token_cost_above_272k_tokens": 5e-07, + "cache_read_input_token_cost_flex": 1.3e-07, + "cache_read_input_token_cost_priority": 5e-07, + "input_cost_per_token": 2.5e-06, + "input_cost_per_token_above_272k_tokens": 5e-06, + "input_cost_per_token_batches": 1.25e-06, + "input_cost_per_token_flex": 1.25e-06, + "input_cost_per_token_priority": 5e-06, + "litellm_provider": "openai", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "output_cost_per_token_above_272k_tokens": 2.25e-05, + "output_cost_per_token_batches": 7.5e-06, + "output_cost_per_token_flex": 7.5e-06, + "output_cost_per_token_priority": 3e-05, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_xhigh_reasoning_effort": true + }, + "gpt-5.4-2026-03-05": { + "cache_read_input_token_cost": 2.5e-07, + "cache_read_input_token_cost_above_272k_tokens": 5e-07, + "cache_read_input_token_cost_flex": 1.3e-07, + "cache_read_input_token_cost_priority": 5e-07, + "input_cost_per_token": 2.5e-06, + "input_cost_per_token_above_272k_tokens": 5e-06, + "input_cost_per_token_batches": 1.25e-06, + "input_cost_per_token_flex": 1.25e-06, + "input_cost_per_token_priority": 5e-06, + "litellm_provider": "openai", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "output_cost_per_token_above_272k_tokens": 2.25e-05, + "output_cost_per_token_batches": 7.5e-06, + "output_cost_per_token_flex": 7.5e-06, + "output_cost_per_token_priority": 3e-05, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, "supports_vision": true }, + "gpt-5.4-mini": { + "cache_read_input_token_cost": 7.5e-08, + "cache_read_input_token_cost_batches": 3.75e-08, + "cache_read_input_token_cost_flex": 3.75e-08, + "cache_read_input_token_cost_priority": 1.5e-07, + "input_cost_per_token": 7.5e-07, + "input_cost_per_token_batches": 3.75e-07, + "input_cost_per_token_flex": 3.75e-07, + "input_cost_per_token_priority": 1.5e-06, + "litellm_provider": "openai", + "max_input_tokens": 400000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 4.5e-06, + "output_cost_per_token_batches": 2.25e-06, + "output_cost_per_token_flex": 2.25e-06, + "output_cost_per_token_priority": 9e-06, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "gpt-5.4-mini-2026-03-17": { + "cache_read_input_token_cost": 7.5e-08, + "cache_read_input_token_cost_batches": 3.75e-08, + "cache_read_input_token_cost_flex": 3.75e-08, + "cache_read_input_token_cost_priority": 1.5e-07, + "input_cost_per_token": 7.5e-07, + "input_cost_per_token_batches": 3.75e-07, + "input_cost_per_token_flex": 3.75e-07, + "input_cost_per_token_priority": 1.5e-06, + "litellm_provider": "openai", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 4.5e-06, + "output_cost_per_token_batches": 2.25e-06, + "output_cost_per_token_flex": 2.25e-06, + "output_cost_per_token_priority": 9e-06, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "gpt-5.4-nano": { + "cache_read_input_token_cost": 2e-08, + "cache_read_input_token_cost_batches": 1e-08, + "cache_read_input_token_cost_flex": 1e-08, + "input_cost_per_token": 2e-07, + "input_cost_per_token_batches": 1e-07, + "input_cost_per_token_flex": 1e-07, + "litellm_provider": "openai", + "max_input_tokens": 400000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1.25e-06, + "output_cost_per_token_batches": 6.25e-07, + "output_cost_per_token_flex": 6.25e-07, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "gpt-5.4-nano-2026-03-17": { + "cache_read_input_token_cost": 2e-08, + "cache_read_input_token_cost_batches": 1e-08, + "cache_read_input_token_cost_flex": 1e-08, + "input_cost_per_token": 2e-07, + "input_cost_per_token_batches": 1e-07, + "input_cost_per_token_flex": 1e-07, + "litellm_provider": "openai", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1.25e-06, + "output_cost_per_token_batches": 6.25e-07, + "output_cost_per_token_flex": 6.25e-07, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "gpt-5.4-pro": { + "cache_read_input_token_cost": 3e-06, + "cache_read_input_token_cost_above_272k_tokens": 6e-06, + "input_cost_per_token": 3e-05, + "input_cost_per_token_above_272k_tokens": 6e-05, + "input_cost_per_token_batches": 1.5e-05, + "input_cost_per_token_flex": 1.5e-05, + "litellm_provider": "openai", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "responses", + "output_cost_per_token": 0.00018, + "output_cost_per_token_above_272k_tokens": 0.00027, + "output_cost_per_token_batches": 9e-05, + "output_cost_per_token_flex": 9e-05, + "supported_endpoints": [ + "/v1/responses", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, + "supports_native_streaming": true, + "supports_none_reasoning_effort": false, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": false, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "gpt-5.4-pro-2026-03-05": { + "cache_read_input_token_cost": 3e-06, + "cache_read_input_token_cost_above_272k_tokens": 6e-06, + "input_cost_per_token": 3e-05, + "input_cost_per_token_above_272k_tokens": 6e-05, + "input_cost_per_token_batches": 1.5e-05, + "input_cost_per_token_flex": 1.5e-05, + "litellm_provider": "openai", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "responses", + "output_cost_per_token": 0.00018, + "output_cost_per_token_above_272k_tokens": 0.00027, + "output_cost_per_token_batches": 9e-05, + "output_cost_per_token_flex": 9e-05, + "supported_endpoints": [ + "/v1/responses", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_minimal_reasoning_effort": true, + "supports_native_streaming": true, + "supports_none_reasoning_effort": false, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": false, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "gpt-5.5": { + "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_above_272k_tokens": 1e-06, + "cache_read_input_token_cost_flex": 2.5e-07, + "cache_read_input_token_cost_priority": 1e-06, + "input_cost_per_token": 5e-06, + "input_cost_per_token_above_272k_tokens": 1e-05, + "input_cost_per_token_batches": 2.5e-06, + "input_cost_per_token_flex": 2.5e-06, + "input_cost_per_token_priority": 1e-05, + "litellm_provider": "openai", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 3e-05, + "output_cost_per_token_above_272k_tokens": 4.5e-05, + "output_cost_per_token_batches": 1.5e-05, + "output_cost_per_token_flex": 1.5e-05, + "output_cost_per_token_priority": 6e-05, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "gpt-5.5-2026-04-23": { + "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_above_272k_tokens": 1e-06, + "cache_read_input_token_cost_flex": 2.5e-07, + "cache_read_input_token_cost_priority": 1e-06, + "input_cost_per_token": 5e-06, + "input_cost_per_token_above_272k_tokens": 1e-05, + "input_cost_per_token_batches": 2.5e-06, + "input_cost_per_token_flex": 2.5e-06, + "input_cost_per_token_priority": 1e-05, + "litellm_provider": "openai", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 3e-05, + "output_cost_per_token_above_272k_tokens": 4.5e-05, + "output_cost_per_token_batches": 1.5e-05, + "output_cost_per_token_flex": 1.5e-05, + "output_cost_per_token_priority": 6e-05, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "gpt-5.5-pro": { + "cache_read_input_token_cost": 3e-06, + "cache_read_input_token_cost_above_272k_tokens": 6e-06, + "input_cost_per_token": 3e-05, + "input_cost_per_token_above_272k_tokens": 6e-05, + "input_cost_per_token_batches": 1.5e-05, + "input_cost_per_token_flex": 1.5e-05, + "litellm_provider": "openai", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "responses", + "output_cost_per_token": 0.00018, + "output_cost_per_token_above_272k_tokens": 0.00027, + "output_cost_per_token_batches": 9e-05, + "output_cost_per_token_flex": 9e-05, + "supported_endpoints": [ + "/v1/responses", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_low_reasoning_effort": false, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": false, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": false, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "gpt-5.5-pro-2026-04-23": { + "cache_read_input_token_cost": 3e-06, + "cache_read_input_token_cost_above_272k_tokens": 6e-06, + "input_cost_per_token": 3e-05, + "input_cost_per_token_above_272k_tokens": 6e-05, + "input_cost_per_token_batches": 1.5e-05, + "input_cost_per_token_flex": 1.5e-05, + "litellm_provider": "openai", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "responses", + "output_cost_per_token": 0.00018, + "output_cost_per_token_above_272k_tokens": 0.00027, + "output_cost_per_token_batches": 9e-05, + "output_cost_per_token_flex": 9e-05, + "supported_endpoints": [ + "/v1/responses", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_low_reasoning_effort": false, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": false, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": false, + "supports_service_tier": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, "gpt-audio": { "input_cost_per_audio_token": 3.2e-05, "input_cost_per_token": 2.5e-06, @@ -5340,6 +5106,39 @@ "supports_tool_choice": true, "supports_vision": false }, + "gpt-audio-1.5": { + "input_cost_per_audio_token": 3.2e-05, + "input_cost_per_token": 2.5e-06, + "litellm_provider": "openai", + "max_input_tokens": 128000, + "max_output_tokens": 16384, + "max_tokens": 16384, + "mode": "chat", + "output_cost_per_audio_token": 6.4e-05, + "output_cost_per_token": 1e-05, + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supported_modalities": [ + "text", + "audio" + ], + "supported_output_modalities": [ + "text", + "audio" + ], + "supports_audio_input": true, + "supports_audio_output": true, + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_prompt_caching": false, + "supports_reasoning": false, + "supports_response_schema": false, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": false + }, "gpt-audio-2025-08-28": { "input_cost_per_audio_token": 3.2e-05, "input_cost_per_token": 2.5e-06, @@ -5540,6 +5339,38 @@ "supports_pdf_input": true, "supports_vision": true }, + "gpt-image-2": { + "cache_read_input_image_token_cost": 2e-06, + "cache_read_input_token_cost": 1.25e-06, + "input_cost_per_image_token": 8e-06, + "input_cost_per_token": 5e-06, + "litellm_provider": "openai", + "mode": "image_generation", + "output_cost_per_image_token": 3e-05, + "output_cost_per_token": 1e-05, + "supported_endpoints": [ + "/v1/images/generations", + "/v1/images/edits" + ], + "supports_pdf_input": true, + "supports_vision": true + }, + "gpt-image-2-2026-04-21": { + "cache_read_input_image_token_cost": 2e-06, + "cache_read_input_token_cost": 1.25e-06, + "input_cost_per_image_token": 8e-06, + "input_cost_per_token": 5e-06, + "litellm_provider": "openai", + "mode": "image_generation", + "output_cost_per_image_token": 3e-05, + "output_cost_per_token": 1e-05, + "supported_endpoints": [ + "/v1/images/generations", + "/v1/images/edits" + ], + "supports_pdf_input": true, + "supports_vision": true + }, "gpt-realtime": { "cache_creation_input_audio_token_cost": 4e-07, "cache_read_input_token_cost": 4e-07, @@ -5572,6 +5403,70 @@ "supports_system_messages": true, "supports_tool_choice": true }, + "gpt-realtime-1.5": { + "cache_creation_input_audio_token_cost": 4e-07, + "cache_read_input_token_cost": 4e-07, + "input_cost_per_audio_token": 3.2e-05, + "input_cost_per_image": 5e-06, + "input_cost_per_token": 4e-06, + "litellm_provider": "openai", + "max_input_tokens": 32000, + "max_output_tokens": 4096, + "max_tokens": 4096, + "mode": "chat", + "output_cost_per_audio_token": 6.4e-05, + "output_cost_per_token": 1.6e-05, + "supported_endpoints": [ + "/v1/realtime" + ], + "supported_modalities": [ + "text", + "image", + "audio" + ], + "supported_output_modalities": [ + "text", + "audio" + ], + "supports_audio_input": true, + "supports_audio_output": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_system_messages": true, + "supports_tool_choice": true + }, + "gpt-realtime-2": { + "cache_creation_input_audio_token_cost": 4e-07, + "cache_read_input_token_cost": 4e-07, + "input_cost_per_audio_token": 3.2e-05, + "input_cost_per_image": 5e-06, + "input_cost_per_token": 4e-06, + "litellm_provider": "openai", + "max_input_tokens": 32000, + "max_output_tokens": 4096, + "max_tokens": 4096, + "mode": "chat", + "output_cost_per_audio_token": 6.4e-05, + "output_cost_per_token": 1.6e-05, + "supported_endpoints": [ + "/v1/realtime" + ], + "supported_modalities": [ + "text", + "image", + "audio" + ], + "supported_output_modalities": [ + "text", + "audio" + ], + "supports_audio_input": true, + "supports_audio_output": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_system_messages": true, + "supports_tool_choice": true + }, "gpt-realtime-2025-08-28": { "cache_creation_input_audio_token_cost": 4e-07, "cache_read_input_token_cost": 4e-07, @@ -5720,62 +5615,6 @@ "supports_tool_choice": true, "supports_vision": true }, - "o1-mini": { - "cache_read_input_token_cost": 5.5e-07, - "input_cost_per_token": 1.1e-06, - "litellm_provider": "openai", - "max_input_tokens": 128000, - "max_output_tokens": 65536, - "max_tokens": 65536, - "mode": "chat", - "output_cost_per_token": 4.4e-06, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_vision": true - }, - "o1-mini-2024-09-12": { - "cache_read_input_token_cost": 1.5e-06, - "deprecation_date": "2025-10-27", - "input_cost_per_token": 3e-06, - "litellm_provider": "openai", - "max_input_tokens": 128000, - "max_output_tokens": 65536, - "max_tokens": 65536, - "mode": "chat", - "output_cost_per_token": 1.2e-05, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_vision": true - }, - "o1-preview": { - "cache_read_input_token_cost": 7.5e-06, - "input_cost_per_token": 1.5e-05, - "litellm_provider": "openai", - "max_input_tokens": 128000, - "max_output_tokens": 32768, - "max_tokens": 32768, - "mode": "chat", - "output_cost_per_token": 6e-05, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_vision": true - }, - "o1-preview-2024-09-12": { - "cache_read_input_token_cost": 7.5e-06, - "input_cost_per_token": 1.5e-05, - "litellm_provider": "openai", - "max_input_tokens": 128000, - "max_output_tokens": 32768, - "max_tokens": 32768, - "mode": "chat", - "output_cost_per_token": 6e-05, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_vision": true - }, "o1-pro": { "input_cost_per_token": 0.00015, "input_cost_per_token_batches": 7.5e-05, @@ -5876,7 +5715,8 @@ "supports_response_schema": true, "supports_service_tier": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true }, "o3-2025-04-16": { "cache_read_input_token_cost": 5e-07, @@ -5908,7 +5748,8 @@ "supports_response_schema": true, "supports_service_tier": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true }, "o3-deep-research": { "cache_read_input_token_cost": 2.5e-06, @@ -5941,7 +5782,8 @@ "supports_response_schema": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true }, "o3-deep-research-2025-06-26": { "cache_read_input_token_cost": 2.5e-06, @@ -5974,7 +5816,8 @@ "supports_response_schema": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true }, "o3-mini": { "cache_read_input_token_cost": 5.5e-07, @@ -6038,7 +5881,8 @@ "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true }, "o3-pro-2025-06-10": { "input_cost_per_token": 2e-05, @@ -6068,7 +5912,8 @@ "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true }, "o4-mini": { "cache_read_input_token_cost": 2.75e-07, @@ -6093,7 +5938,8 @@ "supports_response_schema": true, "supports_service_tier": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true }, "o4-mini-2025-04-16": { "cache_read_input_token_cost": 2.75e-07, @@ -6112,7 +5958,8 @@ "supports_response_schema": true, "supports_service_tier": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true }, "o4-mini-deep-research": { "cache_read_input_token_cost": 5e-07, @@ -6145,7 +5992,8 @@ "supports_response_schema": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true }, "o4-mini-deep-research-2025-06-26": { "cache_read_input_token_cost": 5e-07, @@ -6178,6 +6026,7 @@ "supports_response_schema": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_web_search": true } } From 7321e4dea807651dd6fc309eee3d057ba75e7e70 Mon Sep 17 00:00:00 2001 From: "github-actions[bot]" <41898282+github-actions[bot]@users.noreply.github.com> Date: Fri, 29 May 2026 08:50:18 +0000 Subject: [PATCH 40/79] chore: sync VERSION to 0.1.133 [skip ci] --- backend/cmd/server/VERSION | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/backend/cmd/server/VERSION b/backend/cmd/server/VERSION index 7b9dfc4d..56ebc9e5 100644 --- a/backend/cmd/server/VERSION +++ b/backend/cmd/server/VERSION @@ -1 +1 @@ -0.1.132 +0.1.133 From 06fca662735302def3b78bcb47f1650dd09350d8 Mon Sep 17 00:00:00 2001 From: DaydreamCoding <22166516+DaydreamCoding@users.noreply.github.com> Date: Fri, 29 May 2026 09:39:02 +0800 Subject: [PATCH 41/79] =?UTF-8?q?feat(quota):=20sentinel=20=E5=9B=9E?= =?UTF-8?q?=E5=A1=AB=E6=B6=88=E9=99=A4=E6=97=A0=E9=85=8D=E9=A2=9D=E8=A1=8C?= =?UTF-8?q?=E7=94=A8=E6=88=B7=20preflight=20=E6=AF=8F=E8=AF=B7=E6=B1=82?= =?UTF-8?q?=E5=9B=9E=E6=BA=90=20DB?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 无 user×platform 配额行的用户,preflight 每次 cache MISS 后回源 DB 查得"无行" 却不缓存该结论,导致每请求一次 DB 往返。本 PR 回填 sentinel 占位 entry,使后续 请求命中 Redis 后稳定判"无 limit",TTL 内不再查 DB。 - config: 加 UserPlatformQuotaSentinelTTLSeconds(默认 3600s,短于普通 quota cache 的 86400s 以控 Redis 内存) - metrics: 加 userPlatformQuotaSentinelSetCacheErrorTotal,并入 GatewayUserPlatformQuotaIncrStats 暴露 - billing_cache: checkUserPlatformQuotaEligibility 在 cache MISS + DB 无行且 cacheErr==nil 时回填 sentinel(三 limit nil、三 window_start non-nil、SchemaV1); TTL<=0 fallback 1h 防 EXPIRE 立即删 key 击穿;SET 失败 fail-open + 计 metric - billing_cache: HIT 路径对 sentinel(三 limit nil)跳过 windowExpired refresh, 避免短 sentinel TTL 被误升级为 86400s 有配额 limit 的用户 enforcement 行为不变(rec!=nil 不回填、isSentinel=false 不跳过 refresh)。 测试:扩展 fakeFullCache 夹具(setCalls/lastSetTTL/getErr/setErr);新增回填正确性 / Redis-GET-故障不回填 / SET-失败 fail-open / sentinel 跨窗口不 refresh 四个单测。 go build、quota+billing unit、三态 go vet 全绿。 Co-Authored-By: Claude Opus 4.8 (1M context) --- backend/internal/config/config.go | 4 + .../internal/service/billing_cache_service.go | 34 ++++- ..._cache_service_user_platform_quota_test.go | 122 +++++++++++++++++- backend/internal/service/gateway_service.go | 13 +- 4 files changed, 167 insertions(+), 6 deletions(-) diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index f689c2a9..dcbf30b4 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -649,6 +649,9 @@ type BillingConfig struct { // - billing_cache_service.checkUserPlatformQuotaEligibility 首次缓存装载 // 读写两端必须共用同一 TTL,避免缓存生命周期不一致导致 quota 计数漂移。 UserPlatformQuotaCacheTTLSeconds int `mapstructure:"user_platform_quota_cache_ttl_seconds"` + // UserPlatformQuotaSentinelTTLSeconds sentinel(无 limit 占位)entry 的 TTL, + // 显著短于 quota cache 默认 86400s 以控 Redis 内存;默认 3600=1h。 + UserPlatformQuotaSentinelTTLSeconds int `mapstructure:"user_platform_quota_sentinel_ttl_seconds"` } type CircuitBreakerConfig struct { @@ -1571,6 +1574,7 @@ func setDefaults() { viper.SetDefault("billing.circuit_breaker.reset_timeout_seconds", 30) viper.SetDefault("billing.circuit_breaker.half_open_requests", 3) viper.SetDefault("billing.user_platform_quota_cache_ttl_seconds", 86400) + viper.SetDefault("billing.user_platform_quota_sentinel_ttl_seconds", 3600) // Turnstile viper.SetDefault("turnstile.required", false) diff --git a/backend/internal/service/billing_cache_service.go b/backend/internal/service/billing_cache_service.go index 2b7c06ba..8a5172f4 100644 --- a/backend/internal/service/billing_cache_service.go +++ b/backend/internal/service/billing_cache_service.go @@ -1096,7 +1096,12 @@ func (s *BillingCacheService) checkUserPlatformQuotaEligibility( // 超时 50ms:覆盖正常路径与可接受抖动;Redis 异常时 hot path 不阻塞超过此值。 // 用 context.Background()+短超时,避免请求 ctx 取消导致刷新丢失。 // 显式 setCancel()(而非 defer):缩短 context 生命周期,避免 defer 延迟到函数返回。 - if windowExpired && s.cache != nil { + // isSentinel 判定「该 entry 无任何 limit」,涵盖两类,跨窗口命中时都跳过 refresh: + // 1) A3 回填的 sentinel(DB 无行,短 TTL):refresh 会把短 TTL 误升级为 86400s,有害; + // 2) DB 有行但三 limit 全未配置的用户(TTL 86400s):refresh 纯属无意义(TTL 升级本身无害)。 + // 两类的 enforcement(下方 limit!=nil 比较)都因 limit 全 nil 永远放行,跳过 refresh 均正确。 + isSentinel := entry.DailyLimitUSD == nil && entry.WeeklyLimitUSD == nil && entry.MonthlyLimitUSD == nil + if windowExpired && s.cache != nil && !isSentinel { refreshed := &UserPlatformQuotaCacheEntry{ DailyUsageUSD: dailyUsage, WeeklyUsageUSD: weeklyUsage, @@ -1159,6 +1164,33 @@ func (s *BillingCacheService) checkUserPlatformQuotaEligibility( } rec, _ := v.(*UserPlatformQuotaRecord) if rec == nil { + // 仅在 cache 可用且本次 GET 未出错时回填 sentinel:Redis GET 故障(cacheErr!=nil) + // 时不回填,与下方 line ~1201 "Redis 故障时 fail-open:不回填" 保持一致, + // 避免在 Redis 异常期做一次注定失败的 SET。 + if s.cache != nil && cacheErr == nil { + now := time.Now() + startOfDay := timezone.StartOfDay(now) + startOfWeek := timezone.StartOfWeek(now) + sentinel := &UserPlatformQuotaCacheEntry{ + SchemaVersion: UserPlatformQuotaCacheSchemaV1, + DailyWindowStart: &startOfDay, + WeeklyWindowStart: &startOfWeek, + MonthlyWindowStart: &now, + // limits 全 nil, usage 全 0(零值) + } + sentinelTTL := time.Duration(s.cfg.Billing.UserPlatformQuotaSentinelTTLSeconds) * time.Second + if sentinelTTL <= 0 { + // 防御:TTL<=0 时 Redis EXPIRE 会立即删除整个 key(见 billing_cache.go 的 pipe.Expire), + // sentinel 不持久化 → 每请求击穿 DB。配置缺失/误配为 0 时 fallback 到 1h。 + sentinelTTL = time.Hour + } + setCtx, setCancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + if setErr := s.cache.SetUserPlatformQuotaCache(setCtx, userID, platform, sentinel, sentinelTTL); setErr != nil { + userPlatformQuotaSentinelSetCacheErrorTotal.Add(1) + logger.LegacyPrintf("service.billing_cache", "Warning: set sentinel quota cache failed user=%d platform=%s: %v", userID, platform, setErr) + } + setCancel() + } return nil } diff --git a/backend/internal/service/billing_cache_service_user_platform_quota_test.go b/backend/internal/service/billing_cache_service_user_platform_quota_test.go index 57697ddb..674aa9a5 100644 --- a/backend/internal/service/billing_cache_service_user_platform_quota_test.go +++ b/backend/internal/service/billing_cache_service_user_platform_quota_test.go @@ -95,6 +95,10 @@ type fakeFullCache struct { mu sync.Mutex entry *UserPlatformQuotaCacheEntry deleteCalls int + setCalls int // SetUserPlatformQuotaCache 调用次数 + lastSetTTL time.Duration // 最近一次 Set 的 ttl + getErr error // 非 nil 时 Get 先返回 (nil,false,getErr) + setErr error // 非 nil 时 Set 返回该 err(setCalls 仍+1) } // getDeleteCalls 线程安全地读取 deleteCalls。 @@ -111,19 +115,41 @@ func (f *fakeFullCache) getEntry() *UserPlatformQuotaCacheEntry { return f.entry } +// getSetCalls 线程安全地读取 setCalls。 +func (f *fakeFullCache) getSetCalls() int { + f.mu.Lock() + defer f.mu.Unlock() + return f.setCalls +} + +// getLastSetTTL 线程安全地读取 lastSetTTL。 +func (f *fakeFullCache) getLastSetTTL() time.Duration { + f.mu.Lock() + defer f.mu.Unlock() + return f.lastSetTTL +} + func (f *fakeFullCache) GetUserPlatformQuotaCache(_ context.Context, _ int64, _ string) (*UserPlatformQuotaCacheEntry, bool, error) { f.mu.Lock() defer f.mu.Unlock() + if f.getErr != nil { + return nil, false, f.getErr + } if f.entry == nil { return nil, false, nil } return f.entry, true, nil } -func (f *fakeFullCache) SetUserPlatformQuotaCache(_ context.Context, _ int64, _ string, e *UserPlatformQuotaCacheEntry, _ time.Duration) error { +func (f *fakeFullCache) SetUserPlatformQuotaCache(_ context.Context, _ int64, _ string, e *UserPlatformQuotaCacheEntry, ttl time.Duration) error { f.mu.Lock() defer f.mu.Unlock() + f.setCalls++ + if f.setErr != nil { + return f.setErr + } f.entry = e + f.lastSetTTL = ttl return nil } @@ -593,3 +619,97 @@ func TestMonthlyQuotaWindowExpired_BoundaryTable(t *testing.T) { }) } } + +// TestCheckUserPlatformQuotaEligibility_NoRow_WritesSentinel 验证: +// cache MISS + DB 无行时,回填 sentinel entry(三 limit 全 nil,三 window_start 全 non-nil), +// TTL = UserPlatformQuotaSentinelTTLSeconds,函数返回 nil(fail-open)。 +func TestCheckUserPlatformQuotaEligibility_NoRow_WritesSentinel(t *testing.T) { + repo := &fakeQuotaRepo{rec: nil} // DB 无行 + cache := &fakeFullCache{} // entry=nil → Get 返回 MISS + svc := newServiceForPreflight(t, repo, cache) + svc.cfg.Billing.UserPlatformQuotaSentinelTTLSeconds = 3600 + + if err := svc.checkUserPlatformQuotaEligibility(context.Background(), 1, "anthropic"); err != nil { + t.Fatalf("expected nil (fail-open), got %v", err) + } + if cache.getSetCalls() != 1 { + t.Fatalf("expected 1 SetUserPlatformQuotaCache call for sentinel, got %d", cache.getSetCalls()) + } + sentinel := cache.getEntry() + if sentinel == nil { + t.Fatal("expected sentinel entry backfilled") + } + if sentinel.DailyLimitUSD != nil || sentinel.WeeklyLimitUSD != nil || sentinel.MonthlyLimitUSD != nil { + t.Errorf("sentinel must have all-nil limits") + } + if sentinel.DailyWindowStart == nil || sentinel.WeeklyWindowStart == nil || sentinel.MonthlyWindowStart == nil { + t.Errorf("sentinel must have non-nil window_start to avoid refresh churn") + } + if sentinel.SchemaVersion != UserPlatformQuotaCacheSchemaV1 { + t.Errorf("sentinel schema = %d, want V1", sentinel.SchemaVersion) + } + if cache.getLastSetTTL() != 3600*time.Second { + t.Errorf("sentinel ttl = %v, want 3600s", cache.getLastSetTTL()) + } +} + +// TestCheckUserPlatformQuotaEligibility_RedisGetError_NoSentinelBackfill 验证: +// Redis GET 故障(cacheErr!=nil)+ DB 无行时,不应回填 sentinel(与 "Redis 故障时不回填" 一致),且 fail-open。 +func TestCheckUserPlatformQuotaEligibility_RedisGetError_NoSentinelBackfill(t *testing.T) { + repo := &fakeQuotaRepo{rec: nil} + cache := &fakeFullCache{getErr: errors.New("redis get down")} + svc := newServiceForPreflight(t, repo, cache) + svc.cfg.Billing.UserPlatformQuotaSentinelTTLSeconds = 3600 + + if err := svc.checkUserPlatformQuotaEligibility(context.Background(), 1, "anthropic"); err != nil { + t.Fatalf("redis 故障应 fail-open, got %v", err) + } + if cache.getSetCalls() != 0 { + t.Errorf("redis-get-error 时不应回填 sentinel, got %d set calls", cache.getSetCalls()) + } +} + +// TestCheckUserPlatformQuotaEligibility_NoRow_SentinelSetFailsFailOpen 验证: +// sentinel SET 失败时 fail-open(返回 nil)且计 metric。 +func TestCheckUserPlatformQuotaEligibility_NoRow_SentinelSetFailsFailOpen(t *testing.T) { + before := userPlatformQuotaSentinelSetCacheErrorTotal.Load() + repo := &fakeQuotaRepo{rec: nil} + cache := &fakeFullCache{setErr: errors.New("redis set timeout")} + svc := newServiceForPreflight(t, repo, cache) + svc.cfg.Billing.UserPlatformQuotaSentinelTTLSeconds = 3600 + + if err := svc.checkUserPlatformQuotaEligibility(context.Background(), 1, "anthropic"); err != nil { + t.Fatalf("sentinel set 失败应 fail-open, got %v", err) + } + if cache.getSetCalls() != 1 { + t.Errorf("应尝试 set sentinel 恰好一次, got %d", cache.getSetCalls()) + } + if got := userPlatformQuotaSentinelSetCacheErrorTotal.Load() - before; got != 1 { + t.Errorf("set 失败应使 metric +1, got delta %d", got) + } +} + +// TestCheckUserPlatformQuotaEligibility_SentinelCrossDay_NoRefresh 验证: +// 命中 sentinel(三 limit 全 nil)且跨窗口(daily/weekly 过期)时,不应触发 refresh SetCache +// (否则会把短 sentinel TTL 误升级为 quota cache 默认 86400s)。 +func TestCheckUserPlatformQuotaEligibility_SentinelCrossDay_NoRefresh(t *testing.T) { + yesterday := timezone.StartOfDay(time.Now().AddDate(0, 0, -1)) + lastWeek := timezone.StartOfWeek(time.Now().AddDate(0, 0, -7)) + monthAgoOK := time.Now().AddDate(0, 0, -5) // <30d, monthly 不过期 + sentinel := &UserPlatformQuotaCacheEntry{ + SchemaVersion: UserPlatformQuotaCacheSchemaV1, + DailyWindowStart: &yesterday, // 跨日 → daily windowExpired = true + WeeklyWindowStart: &lastWeek, // 跨周 → weekly windowExpired = true + MonthlyWindowStart: &monthAgoOK, + // limits 全 nil → sentinel + } + cache := &fakeFullCache{entry: sentinel} // entry 非 nil → Get HIT + svc := newServiceForPreflight(t, &fakeQuotaRepo{}, cache) + + if err := svc.checkUserPlatformQuotaEligibility(context.Background(), 1, "anthropic"); err != nil { + t.Fatalf("sentinel = no limit, expected nil, got %v", err) + } + if cache.getSetCalls() != 0 { + t.Errorf("sentinel cross-window must NOT trigger refresh SetCache, got %d calls", cache.getSetCalls()) + } +} diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index 94197f37..effa803a 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -105,6 +105,9 @@ var ( // (applyUsageBilling 在 repo==nil 时 fallback)路径下的失败次数; // 与 DB Incr 失败分开计数,便于区分"主路径暂时故障"vs"基础设施长期未配齐"。 userPlatformQuotaDBIncrLegacyErrorTotal atomic.Int64 + // userPlatformQuotaSentinelSetCacheErrorTotal 统计 checkUserPlatformQuotaEligibility + // 在 DB 无行时回填 sentinel cache entry 写 Redis 失败的次数(phase A)。 + userPlatformQuotaSentinelSetCacheErrorTotal atomic.Int64 ) func GatewayWindowCostPrefetchStats() (cacheHit, cacheMiss, batchSQL, fallback, errCount int64) { @@ -127,13 +130,15 @@ func GatewayModelsListCacheStats() (cacheHit, cacheMiss, store int64) { return modelsListCacheHitTotal.Load(), modelsListCacheMissTotal.Load(), modelsListCacheStoreTotal.Load() } -// GatewayUserPlatformQuotaIncrStats 返回 (mainPathErr, legacyPathErr)。 +// GatewayUserPlatformQuotaIncrStats 返回 (mainPathErr, legacyPathErr, sentinelSetErr)。 // mainPathErr:finalizePostUsageBilling 异步 goroutine 写 DB 失败累计次数; -// legacyPathErr:postUsageBilling fallback 路径写 DB 失败累计次数。 +// legacyPathErr:postUsageBilling fallback 路径写 DB 失败累计次数; +// sentinelSetErr:DB 无行时回填 sentinel cache entry 写 Redis 失败累计次数。 // ops 监控面板可以按"持续上升斜率"做告警阈值。 -func GatewayUserPlatformQuotaIncrStats() (mainPathErr, legacyPathErr int64) { +func GatewayUserPlatformQuotaIncrStats() (mainPathErr, legacyPathErr, sentinelSetErr int64) { return userPlatformQuotaDBIncrErrorTotal.Load(), - userPlatformQuotaDBIncrLegacyErrorTotal.Load() + userPlatformQuotaDBIncrLegacyErrorTotal.Load(), + userPlatformQuotaSentinelSetCacheErrorTotal.Load() } func openAIStreamEventIsTerminal(data string) bool { From f7f5e33830bcf4d8998544cfb1d726c4c436af1f Mon Sep 17 00:00:00 2001 From: DaydreamCoding <22166516+DaydreamCoding@users.noreply.github.com> Date: Fri, 29 May 2026 13:11:36 +0800 Subject: [PATCH 42/79] =?UTF-8?q?feat(quota):=20user=C3=97platform=20?= =?UTF-8?q?=E9=85=8D=E9=A2=9D=20DB=20=E5=86=99=E8=81=9A=E5=90=88=20flusher?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Redis 同步权威 + DB 镜像,不在进程内维护 delta: - 写入点 HasUserPlatformQuotaLimit 守卫:无 limit 跳过 Redis 写与持久化 - 累加 usage 的 Lua 在 flusher_enabled 时 SADD 脏集 billing:upq:dirty - UserPlatformQuotaUsageFlusher 定时 SPOP 脏集 → 批量 HGETALL 读当前窗口 usage 快照 → BatchSnapshotUsage 绝对值 UPSERT 覆盖 DB(去 SELECT FOR UPDATE 行锁) → 失败 SADD 回 / FK(23503)整批丢弃 - flusher 单批 clamp 到 ≤6000,保证一次 flush 只生成一条 UPSERT(单事务原子) - flusher_enabled 默认 false(降级=旧异步直写 DB) 效果:DB 写连接从 O(QPS) 收敛到 O(副本)。 循环依赖:service 层独立 Snapshot/FK 类型,repository adapter 转换 + %w 映射 FK error。 admin reset/upsert 后失效 cache(脏残留被 flusher 当 MISS 跳过)。 健壮性与可观测性: - flusher_enabled=false 时 Start 不注册定时器;flush_interval_ms 非法回退 2s - Readd 回填失败单独计 dirty_lost(不再误记 dirty_readd)并 ALERT;脏集 Readd 补兜底 TTL - 单 tick 达 max batches 上限仍有积压时记 log - admin 失效 cache 失败升级为 ALERT(提示 enforcement 可能延迟至 sentinel TTL) - BatchGet 单条命令失败 / usage 字段损坏均记 log,避免静默以 0 覆写 DB 三态 go vet + 单测/集成测全绿。 已知取舍(默认 flusher_enabled=false 不触发): - FK 整批丢弃牵连同批正常 key(活跃 key 靠下次 SADD+绝对值快照自愈;Redis 仍权威) - admin reset/upsert 直写 DB 与 flusher 异步刷存在覆盖竞态:flusher 持旧快照在途时可能覆盖 admin 刚写值(limit 列不受影响;usage 有 preflight windowExpired 兜底;低频)。彻底消除需 version OCC。 Co-Authored-By: Claude Opus 4.8 --- backend/cmd/server/wire.go | 7 + backend/cmd/server/wire_gen.go | 10 +- backend/cmd/server/wire_gen_test.go | 1 + backend/internal/config/config.go | 10 + .../internal/handler/admin/user_handler.go | 4 +- backend/internal/repository/billing_cache.go | 210 +++++-- .../billing_cache_user_platform_quota_test.go | 6 +- .../user_platform_quota_adapter_test.go | 3 + .../repository/user_platform_quota_repo.go | 92 ++++ ...er_platform_quota_repo_integration_test.go | 98 ++++ .../user_platform_quota_service_adapter.go | 46 ++ .../service/admin_service_delete_test.go | 14 +- .../auth_service_platform_quota_test.go | 4 + .../service/auth_service_register_test.go | 4 + .../internal/service/billing_cache_service.go | 20 +- ...billing_cache_service_singleflight_test.go | 14 +- .../service/billing_cache_service_test.go | 14 +- ..._cache_service_user_platform_quota_test.go | 135 ++++- backend/internal/service/billing_service.go | 14 +- backend/internal/service/gateway_service.go | 96 ++-- .../service/user_platform_quota_flusher.go | 267 +++++++++ .../user_platform_quota_flusher_test.go | 511 ++++++++++++++++++ .../service/user_platform_quota_port.go | 19 + backend/internal/service/user_service_test.go | 14 +- backend/internal/service/wire.go | 8 + 25 files changed, 1515 insertions(+), 106 deletions(-) create mode 100644 backend/internal/service/user_platform_quota_flusher.go create mode 100644 backend/internal/service/user_platform_quota_flusher_test.go diff --git a/backend/cmd/server/wire.go b/backend/cmd/server/wire.go index 9bfa2717..b474cfa1 100644 --- a/backend/cmd/server/wire.go +++ b/backend/cmd/server/wire.go @@ -98,6 +98,7 @@ func provideCleanup( backupSvc *service.BackupService, paymentOrderExpiry *service.PaymentOrderExpiryService, channelMonitorRunner *service.ChannelMonitorRunner, + quotaFlusher *service.UserPlatformQuotaUsageFlusher, ) func() { return func() { ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) @@ -246,6 +247,12 @@ func provideCleanup( } return nil }}, + {"UserPlatformQuotaUsageFlusher", func() error { + if quotaFlusher != nil { + quotaFlusher.Stop() + } + return nil + }}, } infraSteps := []cleanupStep{ diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index 6e8be8fc..9b059c5d 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -269,7 +269,8 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { scheduledTestRunnerService := service.ProvideScheduledTestRunnerService(scheduledTestPlanRepository, scheduledTestService, accountTestService, rateLimitService, configConfig) paymentOrderExpiryService := service.ProvidePaymentOrderExpiryService(paymentService) channelMonitorRunner := service.ProvideChannelMonitorRunner(channelMonitorService, settingService) - v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, schedulerSnapshotService, tokenRefreshService, accountExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner) + userPlatformQuotaUsageFlusher := service.ProvideUserPlatformQuotaUsageFlusher(configConfig, billingCache, serviceUserPlatformQuotaRepository, timingWheelService) + v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, schedulerSnapshotService, tokenRefreshService, accountExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, userPlatformQuotaUsageFlusher) application := &Application{ Server: httpServer, Cleanup: v, @@ -324,6 +325,7 @@ func provideCleanup( backupSvc *service.BackupService, paymentOrderExpiry *service.PaymentOrderExpiryService, channelMonitorRunner *service.ChannelMonitorRunner, + quotaFlusher *service.UserPlatformQuotaUsageFlusher, ) func() { return func() { ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) @@ -471,6 +473,12 @@ func provideCleanup( } return nil }}, + {"UserPlatformQuotaUsageFlusher", func() error { + if quotaFlusher != nil { + quotaFlusher.Stop() + } + return nil + }}, } infraSteps := []cleanupStep{ diff --git a/backend/cmd/server/wire_gen_test.go b/backend/cmd/server/wire_gen_test.go index a44b2d5c..7f4e4773 100644 --- a/backend/cmd/server/wire_gen_test.go +++ b/backend/cmd/server/wire_gen_test.go @@ -77,6 +77,7 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) { nil, // backupSvc nil, // paymentOrderExpiry nil, // channelMonitorRunner + nil, // quotaFlusher ) require.NotPanics(t, func() { diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index dcbf30b4..df9dcefc 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -1094,6 +1094,13 @@ type DatabaseConfig struct { ConnMaxLifetimeMinutes int `mapstructure:"conn_max_lifetime_minutes"` // ConnMaxIdleTimeMinutes: 空闲连接最大存活时间,及时释放不活跃连接 ConnMaxIdleTimeMinutes int `mapstructure:"conn_max_idle_time_minutes"` + // UserPlatformQuotaFlusherEnabled: 是否启用 user×platform 配额写聚合 flusher + UserPlatformQuotaFlusherEnabled bool `mapstructure:"user_platform_quota_flusher_enabled"` + // UserPlatformQuotaFlushIntervalMs: flusher 刷写间隔(毫秒) + UserPlatformQuotaFlushIntervalMs int `mapstructure:"user_platform_quota_flush_interval_ms"` + // UserPlatformQuotaFlushBatchSize: flusher 单批最大条数 + // 建议 ≤ 6000(单条 UPSERT 原子上限) + UserPlatformQuotaFlushBatchSize int `mapstructure:"user_platform_quota_flush_batch_size"` } func (d *DatabaseConfig) DSN() string { @@ -1661,6 +1668,9 @@ func setDefaults() { viper.SetDefault("database.max_idle_conns", 128) viper.SetDefault("database.conn_max_lifetime_minutes", 30) viper.SetDefault("database.conn_max_idle_time_minutes", 5) + viper.SetDefault("database.user_platform_quota_flusher_enabled", false) + viper.SetDefault("database.user_platform_quota_flush_interval_ms", 2000) + viper.SetDefault("database.user_platform_quota_flush_batch_size", 1000) // Redis viper.SetDefault("redis.host", "localhost") diff --git a/backend/internal/handler/admin/user_handler.go b/backend/internal/handler/admin/user_handler.go index 32a21692..6c0a02ff 100644 --- a/backend/internal/handler/admin/user_handler.go +++ b/backend/internal/handler/admin/user_handler.go @@ -743,7 +743,7 @@ func (h *UserHandler) UpdateUserPlatformQuotas(c *gin.Context) { if h.billingCache != nil { for _, p := range service.AllowedQuotaPlatforms { if err := h.billingCache.DeleteUserPlatformQuotaCache(ctx, userID, p); err != nil { - slog.Warn("quota cache invalidation failed", "user_id", userID, "platform", p, "err", err) + slog.Error("ALERT: quota cache invalidation failed after UpsertForUser; limit 生效可能延迟至 sentinel TTL(最长 1h),需人工确认或重试失效", "user_id", userID, "platform", p, "err", err) } } } @@ -827,7 +827,7 @@ func (h *UserHandler) ResetUserPlatformQuotaWindow(c *gin.Context) { if h.billingCache != nil { if err := h.billingCache.DeleteUserPlatformQuotaCache(ctx, userID, req.Platform); err != nil { - slog.Warn("quota cache invalidation failed", "user_id", userID, "platform", req.Platform, "err", err) + slog.Error("ALERT: quota cache invalidation failed after ResetExpiredWindow; 窗口重置可能延迟至 sentinel TTL(最长 1h)", "user_id", userID, "platform", req.Platform, "err", err) } } diff --git a/backend/internal/repository/billing_cache.go b/backend/internal/repository/billing_cache.go index 60dae954..de229da9 100644 --- a/backend/internal/repository/billing_cache.go +++ b/backend/internal/repository/billing_cache.go @@ -7,6 +7,7 @@ import ( "log" "math/rand/v2" "strconv" + "strings" "time" "github.com/Wei-Shaw/sub2api/internal/service" @@ -338,38 +339,26 @@ func userPlatformQuotaCacheKey(userID int64, platform string) string { return fmt.Sprintf("billing:user_platform_quota:%d:%s", userID, platform) } -func (c *billingCache) GetUserPlatformQuotaCache(ctx context.Context, userID int64, platform string) (*service.UserPlatformQuotaCacheEntry, bool, error) { - key := userPlatformQuotaCacheKey(userID, platform) - fields := []string{ - "daily_usage", "weekly_usage", "monthly_usage", "version", "schema_version", - "daily_limit", "weekly_limit", "monthly_limit", - "daily_window_start", "weekly_window_start", "monthly_window_start", +// parseUserPlatformQuotaHash 将 Redis HGETALL 返回的 map[string]string 反序列化为 +// *service.UserPlatformQuotaCacheEntry。空 map(key 不存在)返回 nil。 +// GetUserPlatformQuotaCache 和 BatchGetUserPlatformQuotaCache 共用此函数,确保解析逻辑一致。 +func parseUserPlatformQuotaHash(m map[string]string) *service.UserPlatformQuotaCacheEntry { + if len(m) == 0 { + return nil } - vals, err := c.rdb.HMGet(ctx, key, fields...).Result() - if err != nil { - return nil, false, err - } - // 前4个全为nil → key 不存在 - if vals[0] == nil && vals[1] == nil && vals[2] == nil && vals[3] == nil { - return nil, false, nil - } - parseFloat := func(v any) float64 { - if v == nil { + parseFloat := func(s string) float64 { + if s == "" { return 0 } - s, ok := v.(string) - if !ok { + f, err := strconv.ParseFloat(s, 64) + if err != nil { + log.Printf("billing_cache: corrupt quota usage field %q (using 0): %v", s, err) return 0 } - f, _ := strconv.ParseFloat(s, 64) return f } - parseFloatPtr := func(v any) *float64 { - if v == nil { - return nil - } - s, ok := v.(string) - if !ok || s == "" { + parseFloatPtr := func(s string) *float64 { + if s == "" { return nil } f, err := strconv.ParseFloat(s, 64) @@ -378,12 +367,8 @@ func (c *billingCache) GetUserPlatformQuotaCache(ctx context.Context, userID int } return &f } - parseTimePtr := func(v any) *time.Time { - if v == nil { - return nil - } - s, ok := v.(string) - if !ok || s == "" { + parseTimePtr := func(s string) *time.Time { + if s == "" { return nil } n, err := strconv.ParseInt(s, 10, 64) @@ -393,30 +378,37 @@ func (c *billingCache) GetUserPlatformQuotaCache(ctx context.Context, userID int t := time.Unix(n, 0).UTC() return &t } - parseInt64 := func(v any) int64 { - if v == nil { - return 0 - } - s, ok := v.(string) - if !ok { - return 0 - } + parseInt64 := func(s string) int64 { n, _ := strconv.ParseInt(s, 10, 64) return n } return &service.UserPlatformQuotaCacheEntry{ - DailyUsageUSD: parseFloat(vals[0]), - WeeklyUsageUSD: parseFloat(vals[1]), - MonthlyUsageUSD: parseFloat(vals[2]), - Version: parseInt64(vals[3]), - SchemaVersion: parseInt64(vals[4]), - DailyLimitUSD: parseFloatPtr(vals[5]), - WeeklyLimitUSD: parseFloatPtr(vals[6]), - MonthlyLimitUSD: parseFloatPtr(vals[7]), - DailyWindowStart: parseTimePtr(vals[8]), - WeeklyWindowStart: parseTimePtr(vals[9]), - MonthlyWindowStart: parseTimePtr(vals[10]), - }, true, nil + DailyUsageUSD: parseFloat(m["daily_usage"]), + WeeklyUsageUSD: parseFloat(m["weekly_usage"]), + MonthlyUsageUSD: parseFloat(m["monthly_usage"]), + Version: parseInt64(m["version"]), + SchemaVersion: parseInt64(m["schema_version"]), + DailyLimitUSD: parseFloatPtr(m["daily_limit"]), + WeeklyLimitUSD: parseFloatPtr(m["weekly_limit"]), + MonthlyLimitUSD: parseFloatPtr(m["monthly_limit"]), + DailyWindowStart: parseTimePtr(m["daily_window_start"]), + WeeklyWindowStart: parseTimePtr(m["weekly_window_start"]), + MonthlyWindowStart: parseTimePtr(m["monthly_window_start"]), + } +} + +func (c *billingCache) GetUserPlatformQuotaCache(ctx context.Context, userID int64, platform string) (*service.UserPlatformQuotaCacheEntry, bool, error) { + key := userPlatformQuotaCacheKey(userID, platform) + m, err := c.rdb.HGetAll(ctx, key).Result() + if err != nil { + return nil, false, err + } + entry := parseUserPlatformQuotaHash(m) + if entry == nil { + // 空 map → key 不存在 → MISS + return nil, false, nil + } + return entry, true, nil } func (c *billingCache) SetUserPlatformQuotaCache(ctx context.Context, userID int64, platform string, entry *service.UserPlatformQuotaCacheEntry, ttl time.Duration) error { @@ -468,9 +460,12 @@ func (c *billingCache) DeleteUserPlatformQuotaCache(ctx context.Context, userID // SetCache 重建为新版 entry —— 若此处仍累加,上层覆盖时会丢失这部分增量,导致 Redis usage 比真实偏小。 // key 不存在同样跳过(由下次 SetCache 重建)。 // KEYS[1] = hash key +// KEYS[2] = 脏集 key(dirty set) // ARGV[1] = cost (string float) // ARGV[2] = ttl seconds // ARGV[3] = expected schema_version (Go 侧 UserPlatformQuotaCacheSchemaV1) +// ARGV[4] = dirty set member(空串则不 SADD) +// ARGV[5] = 脏集兜底 TTL 秒 const updateUserPlatformQuotaUsageScript = ` if redis.call("EXISTS", KEYS[1]) == 0 then return 0 @@ -484,18 +479,125 @@ redis.call("HINCRBYFLOAT", KEYS[1], "weekly_usage", ARGV[1]) redis.call("HINCRBYFLOAT", KEYS[1], "monthly_usage", ARGV[1]) redis.call("HINCRBY", KEYS[1], "version", 1) redis.call("EXPIRE", KEYS[1], ARGV[2]) +if ARGV[4] ~= "" then + redis.call("SADD", KEYS[2], ARGV[4]) + redis.call("EXPIRE", KEYS[2], ARGV[5]) +end return 1 ` -func (c *billingCache) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration) error { - key := userPlatformQuotaCacheKey(userID, platform) - _, err := c.rdb.Eval(ctx, updateUserPlatformQuotaUsageScript, []string{key}, +// userPlatformQuotaDirtySetKey 返回脏集(dirty set)的 Redis key。 +// 使用与 userPlatformQuotaCacheKey 相同的前缀 "billing:"。 +func userPlatformQuotaDirtySetKey() string { return "billing:" + "upq:dirty" } + +// userPlatformQuotaDirtyTTLSeconds 脏集兜底 TTL(秒):初始 SADD(Lua)与 Readd 共用, +// 确保 flusher 长期停摆时脏集最终过期;正常运行因持续 SADD 不断续期。 +const userPlatformQuotaDirtyTTLSeconds = 86400 + +// userPlatformQuotaDirtyMember 构造脏集成员字符串 "userID:platform"。 +func userPlatformQuotaDirtyMember(userID int64, platform string) string { + return strconv.FormatInt(userID, 10) + ":" + platform +} + +func (c *billingCache) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration, markDirty bool) error { + member := "" + if markDirty { + member = userPlatformQuotaDirtyMember(userID, platform) + } + _, err := c.rdb.Eval(ctx, updateUserPlatformQuotaUsageScript, + []string{userPlatformQuotaCacheKey(userID, platform), userPlatformQuotaDirtySetKey()}, strconv.FormatFloat(cost, 'f', -1, 64), int(ttl.Seconds()), service.UserPlatformQuotaCacheSchemaV1, + member, + userPlatformQuotaDirtyTTLSeconds, ).Result() if err != nil && !errors.Is(err, redis.Nil) { return err } return nil } + +// parseUserPlatformQuotaDirtyMember 将脏集成员字符串 "userID:platform" 解析为 +// service.UserPlatformQuotaKey。解析失败返回 ok=false。 +func parseUserPlatformQuotaDirtyMember(m string) (service.UserPlatformQuotaKey, bool) { + parts := strings.SplitN(m, ":", 2) + if len(parts) != 2 { + return service.UserPlatformQuotaKey{}, false + } + uid, err := strconv.ParseInt(parts[0], 10, 64) + if err != nil { + return service.UserPlatformQuotaKey{}, false + } + return service.UserPlatformQuotaKey{UserID: uid, Platform: parts[1]}, true +} + +// PopDirtyUserPlatformQuotaKeys 从脏集随机弹出最多 n 个 key。 +// 脏集为空时返回 (nil, nil)。 +func (c *billingCache) PopDirtyUserPlatformQuotaKeys(ctx context.Context, n int) ([]service.UserPlatformQuotaKey, error) { + members, err := c.rdb.SPopN(ctx, userPlatformQuotaDirtySetKey(), int64(n)).Result() + if err != nil { + if errors.Is(err, redis.Nil) { + return nil, nil + } + return nil, err + } + keys := make([]service.UserPlatformQuotaKey, 0, len(members)) + for _, m := range members { + k, ok := parseUserPlatformQuotaDirtyMember(m) + if !ok { + log.Printf("billing_cache: skipping invalid dirty member %q", m) + continue + } + keys = append(keys, k) + } + return keys, nil +} + +// ReaddDirtyUserPlatformQuotaKeys 将 keys 重新加入脏集(flush 失败时回填)。 +// 通过 pipeline 同时执行 SAdd + Expire,确保 Readd 后脏集具有兜底 TTL。 +// 空切片时直接返回 nil。 +func (c *billingCache) ReaddDirtyUserPlatformQuotaKeys(ctx context.Context, keys []service.UserPlatformQuotaKey) error { + if len(keys) == 0 { + return nil + } + dirtyKey := userPlatformQuotaDirtySetKey() + members := make([]any, len(keys)) + for i, k := range keys { + members[i] = userPlatformQuotaDirtyMember(k.UserID, k.Platform) + } + pipe := c.rdb.Pipeline() + pipe.SAdd(ctx, dirtyKey, members...) + pipe.Expire(ctx, dirtyKey, userPlatformQuotaDirtyTTLSeconds*time.Second) + _, err := pipe.Exec(ctx) + return err +} + +// BatchGetUserPlatformQuotaCache 通过 Pipeline 批量 HGETALL 获取多个 user×platform 的 +// quota cache。返回切片与 keys 顺序、长度对齐;MISS 或解析失败位置返回 nil。 +func (c *billingCache) BatchGetUserPlatformQuotaCache(ctx context.Context, keys []service.UserPlatformQuotaKey) ([]*service.UserPlatformQuotaCacheEntry, error) { + if len(keys) == 0 { + return nil, nil + } + pipe := c.rdb.Pipeline() + cmds := make([]*redis.MapStringStringCmd, len(keys)) + for i, k := range keys { + cmds[i] = pipe.HGetAll(ctx, userPlatformQuotaCacheKey(k.UserID, k.Platform)) + } + if _, err := pipe.Exec(ctx); err != nil && !errors.Is(err, redis.Nil) { + return nil, err + } + results := make([]*service.UserPlatformQuotaCacheEntry, len(keys)) + for i, cmd := range cmds { + m, err := cmd.Result() + if err != nil { + if !errors.Is(err, redis.Nil) { + log.Printf("billing_cache: BatchGet HGETALL cmd[%d] failed: %v (skip, self-heal)", i, err) + } + // 单个命令失败 → 对应位置 nil,继续 + continue + } + results[i] = parseUserPlatformQuotaHash(m) + } + return results, nil +} diff --git a/backend/internal/repository/billing_cache_user_platform_quota_test.go b/backend/internal/repository/billing_cache_user_platform_quota_test.go index 8d49fd31..15b185e7 100644 --- a/backend/internal/repository/billing_cache_user_platform_quota_test.go +++ b/backend/internal/repository/billing_cache_user_platform_quota_test.go @@ -88,7 +88,7 @@ func TestUserPlatformQuotaCache_NilLimitSetThenGet(t *testing.T) { func TestUserPlatformQuotaCache_IncrMissIsNoop(t *testing.T) { c, _ := newMiniRedisCache(t) - if err := c.IncrUserPlatformQuotaUsageCache(context.Background(), 1, "openai", 0.5, time.Minute); err != nil { + if err := c.IncrUserPlatformQuotaUsageCache(context.Background(), 1, "openai", 0.5, time.Minute, false); err != nil { t.Fatal(err) } _, ok, _ := c.GetUserPlatformQuotaCache(context.Background(), 1, "openai") @@ -105,10 +105,10 @@ func TestUserPlatformQuotaCache_IncrHitAccumulates(t *testing.T) { Version: 1, SchemaVersion: service.UserPlatformQuotaCacheSchemaV1, }, time.Minute) - if err := c.IncrUserPlatformQuotaUsageCache(ctx, 1, "openai", 0.5, time.Minute); err != nil { + if err := c.IncrUserPlatformQuotaUsageCache(ctx, 1, "openai", 0.5, time.Minute, false); err != nil { t.Fatal(err) } - if err := c.IncrUserPlatformQuotaUsageCache(ctx, 1, "openai", 0.25, time.Minute); err != nil { + if err := c.IncrUserPlatformQuotaUsageCache(ctx, 1, "openai", 0.25, time.Minute, false); err != nil { t.Fatal(err) } got, _, _ := c.GetUserPlatformQuotaCache(ctx, 1, "openai") diff --git a/backend/internal/repository/user_platform_quota_adapter_test.go b/backend/internal/repository/user_platform_quota_adapter_test.go index a55d2e9c..f31defe5 100644 --- a/backend/internal/repository/user_platform_quota_adapter_test.go +++ b/backend/internal/repository/user_platform_quota_adapter_test.go @@ -38,6 +38,9 @@ func (f *fakeRepoForAdapter) UpsertForUser(_ context.Context, userID int64, reco f.upsertCalledWith = records return f.upsertErr } +func (f *fakeRepoForAdapter) BatchSnapshotUsage(_ context.Context, _ []UserPlatformQuotaSnapshot, _ time.Time) error { + return nil +} func TestGenericAdapter_UpsertForUser_ForwardsRecords(t *testing.T) { fake := &fakeRepoForAdapter{} diff --git a/backend/internal/repository/user_platform_quota_repo.go b/backend/internal/repository/user_platform_quota_repo.go index 1e2e7f51..ccba2330 100644 --- a/backend/internal/repository/user_platform_quota_repo.go +++ b/backend/internal/repository/user_platform_quota_repo.go @@ -2,6 +2,7 @@ package repository import ( "context" + "errors" "fmt" "strings" "time" @@ -9,6 +10,7 @@ import ( dbent "github.com/Wei-Shaw/sub2api/ent" "github.com/Wei-Shaw/sub2api/ent/userplatformquota" "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" + "github.com/lib/pq" ) // UserPlatformQuotaRecord 是 repository 层的传输结构体, @@ -30,6 +32,22 @@ type UserPlatformQuotaRecord struct { // ErrUserPlatformQuotaNotFound 用于 ResetExpiredWindow 等需要"必须命中已有记录"的方法。 var ErrUserPlatformQuotaNotFound = fmt.Errorf("user platform quota record not found") +// ErrUserPlatformQuotaFKViolation 当批量 UPSERT 中存在 user_id 不在 users 表的记录时返回。 +var ErrUserPlatformQuotaFKViolation = errors.New("user platform quota snapshot FK violation") + +// UserPlatformQuotaSnapshot 是 BatchSnapshotUsage 的输入结构体, +// 表示 Redis 当前窗口快照(用于绝对值覆盖写入 DB)。 +type UserPlatformQuotaSnapshot struct { + UserID int64 + Platform string + DailyUsageUSD float64 + WeeklyUsageUSD float64 + MonthlyUsageUSD float64 + DailyWindowStart time.Time + WeeklyWindowStart time.Time + MonthlyWindowStart time.Time +} + // UserPlatformQuotaRepository 定义用户平台配额的数据访问接口。 type UserPlatformQuotaRepository interface { // BulkInsertInitial 幂等批量插入初始配额记录(ON CONFLICT DO NOTHING)。 @@ -44,6 +62,10 @@ type UserPlatformQuotaRepository interface { ResetExpiredWindow(ctx context.Context, userID int64, platform string, window string, newStart time.Time) error // UpsertForUser 全量替换该用户所有平台限额配置(详见 service.UserPlatformQuotaRepository.UpsertForUser)。 UpsertForUser(ctx context.Context, userID int64, records []UserPlatformQuotaRecord) error + // BatchSnapshotUsage 用一条多行 UPSERT 把整批 usage 以绝对值覆盖写入(非累加)。 + // usage/window_start 直接取 EXCLUDED(Redis 当前窗口快照),无 CASE。整批共用 now 作 created/updated_at。 + // 要求 snapshots 内 (user,platform) 不重复。FK 违反返回 ErrUserPlatformQuotaFKViolation。 + BatchSnapshotUsage(ctx context.Context, snapshots []UserPlatformQuotaSnapshot, now time.Time) error } type userPlatformQuotaRepository struct { @@ -414,3 +436,73 @@ func insertLimitsRow(ctx context.Context, client *dbent.Client, userID int64, re } return nil } + +// batchRows 是 BatchSnapshotUsage 每批最大行数(9 参/行 × 6000 ≈ 54000 参,低于 Postgres 65535 上限)。 +const batchRows = 6000 + +// BatchSnapshotUsage 用一条多行 UPSERT 把整批 usage 以绝对值覆盖写入(非累加)。 +// 每批最多 batchRows 行;$1=now 共用;每行 8 个 per-row 参(user_id, platform, 3×usage, 3×window_start)。 +// FK 违反(user_id 不存在)返回 ErrUserPlatformQuotaFKViolation。 +// +// 注意:snapshots 超过 batchRows 会分多条 SQL 执行且【非单事务】——若某子批 FK 失败, +// 先前子批已写入无法回滚。调用方(flusher)应保证单次 batchSize ≤ batchRows +// (默认 flush_batch_size=1000 < 6000,安全)。 +// 另注:启用 flusher 后,本绝对值覆盖与 admin 直写 DB(ResetExpiredWindow/UpsertForUser)存在覆盖竞态, +// 详见 service/user_platform_quota_flusher.go 中 flushOneBatch 的"已知竞态"注释。 +func (r *userPlatformQuotaRepository) BatchSnapshotUsage(ctx context.Context, snapshots []UserPlatformQuotaSnapshot, now time.Time) error { + if len(snapshots) == 0 { + return nil + } + + client := clientFromContext(ctx, r.client) + + for start := 0; start < len(snapshots); start += batchRows { + end := start + batchRows + if end > len(snapshots) { + end = len(snapshots) + } + batch := snapshots[start:end] + + var sb strings.Builder + _, _ = sb.WriteString( + "INSERT INTO user_platform_quotas" + + " (user_id, platform, daily_usage_usd, weekly_usage_usd, monthly_usage_usd," + + " daily_window_start, weekly_window_start, monthly_window_start, created_at, updated_at)" + + " VALUES ") + + // $1 = now(共用);每行 8 个 per-row 参,从 $2 起连续编号。 + args := []any{now} + for i, s := range batch { + if i > 0 { + _, _ = sb.WriteString(",") + } + b := len(args) // 当前 per-row 第一个参数的 0-based 索引,实际占位符 = b+1 + fmt.Fprintf(&sb, "($%d,$%d,$%d,$%d,$%d,$%d,$%d,$%d,$1,$1)", + b+1, b+2, b+3, b+4, b+5, b+6, b+7, b+8) + args = append(args, + s.UserID, s.Platform, + s.DailyUsageUSD, s.WeeklyUsageUSD, s.MonthlyUsageUSD, + s.DailyWindowStart, s.WeeklyWindowStart, s.MonthlyWindowStart, + ) + } + + _, _ = sb.WriteString( + " ON CONFLICT (user_id, platform) WHERE deleted_at IS NULL DO UPDATE SET" + + " daily_usage_usd = EXCLUDED.daily_usage_usd," + + " weekly_usage_usd = EXCLUDED.weekly_usage_usd," + + " monthly_usage_usd = EXCLUDED.monthly_usage_usd," + + " daily_window_start = EXCLUDED.daily_window_start," + + " weekly_window_start = EXCLUDED.weekly_window_start," + + " monthly_window_start = EXCLUDED.monthly_window_start," + + " updated_at = EXCLUDED.updated_at") + + if _, err := client.ExecContext(ctx, sb.String(), args...); err != nil { + var pqErr *pq.Error + if errors.As(err, &pqErr) && pqErr.Code == "23503" { + return ErrUserPlatformQuotaFKViolation + } + return err + } + } + return nil +} diff --git a/backend/internal/repository/user_platform_quota_repo_integration_test.go b/backend/internal/repository/user_platform_quota_repo_integration_test.go index f02eeaa9..39e2f6e0 100644 --- a/backend/internal/repository/user_platform_quota_repo_integration_test.go +++ b/backend/internal/repository/user_platform_quota_repo_integration_test.go @@ -267,3 +267,101 @@ func TestUserPlatformQuotaRepository_ResetExpiredWindow_NotFoundReturnsSentinel( require.True(t, errors.Is(err, ErrUserPlatformQuotaNotFound), "expected ErrUserPlatformQuotaNotFound, got %v", err) } + +// TestBatchSnapshotUsage_InsertOverwriteMultiKey 验证 BatchSnapshotUsage 的绝对值覆盖语义: +// 1. 首批插入 2 条(不同 user),验证 daily 等于首批值; +// 2. 对同一 key 传不同值,验证 daily 等于新值(绝对覆盖,非累加)。 +func TestBatchSnapshotUsage_InsertOverwriteMultiKey(t *testing.T) { + ctx := context.Background() + // BatchSnapshotUsage 不开事务(直接写),使用独立 client 保证跨调用可见性。 + client := testEntClient(t) + + userID1 := mustCreateUserForQuota(t, client) + userID2 := mustCreateUserForQuota(t, client) + + repo := NewUserPlatformQuotaRepository(client) + + now := time.Date(2026, 5, 29, 12, 0, 0, 0, time.UTC) + dailyStart := time.Date(2026, 5, 29, 0, 0, 0, 0, time.UTC) + weeklyStart := time.Date(2026, 5, 25, 0, 0, 0, 0, time.UTC) // 当周一 + monthlyStart := time.Date(2026, 5, 1, 0, 0, 0, 0, time.UTC) + + // ── 第一批:插入 2 行 ────────────────────────────────────────────────────── + firstBatch := []UserPlatformQuotaSnapshot{ + { + UserID: userID1, + Platform: "anthropic", + DailyUsageUSD: 1.0, + WeeklyUsageUSD: 3.0, + MonthlyUsageUSD: 5.0, + DailyWindowStart: dailyStart, + WeeklyWindowStart: weeklyStart, + MonthlyWindowStart: monthlyStart, + }, + { + UserID: userID2, + Platform: "openai", + DailyUsageUSD: 2.0, + WeeklyUsageUSD: 4.0, + MonthlyUsageUSD: 6.0, + DailyWindowStart: dailyStart, + WeeklyWindowStart: weeklyStart, + MonthlyWindowStart: monthlyStart, + }, + } + require.NoError(t, repo.BatchSnapshotUsage(ctx, firstBatch, now), "first batch upsert") + + // 验证首批 daily 值 + rec1, err := repo.GetByUserPlatform(ctx, userID1, "anthropic") + require.NoError(t, err) + require.NotNil(t, rec1, "user1/anthropic should exist after first batch") + require.InDelta(t, 1.0, rec1.DailyUsageUSD, 1e-9, "user1 daily after first batch") + require.InDelta(t, 3.0, rec1.WeeklyUsageUSD, 1e-9, "user1 weekly after first batch") + require.InDelta(t, 5.0, rec1.MonthlyUsageUSD, 1e-9, "user1 monthly after first batch") + + rec2, err := repo.GetByUserPlatform(ctx, userID2, "openai") + require.NoError(t, err) + require.NotNil(t, rec2, "user2/openai should exist after first batch") + require.InDelta(t, 2.0, rec2.DailyUsageUSD, 1e-9, "user2 daily after first batch") + + // ── 第二批:对同一 key 传不同值,验证绝对覆盖(非累加)────────────────── + now2 := now.Add(5 * time.Minute) + secondBatch := []UserPlatformQuotaSnapshot{ + { + UserID: userID1, + Platform: "anthropic", + DailyUsageUSD: 9.9, // 新值,不是 1.0+9.9=10.9 + WeeklyUsageUSD: 19.9, // 新值,不是 3.0+19.9=22.9 + MonthlyUsageUSD: 29.9, // 新值 + DailyWindowStart: dailyStart, + WeeklyWindowStart: weeklyStart, + MonthlyWindowStart: monthlyStart, + }, + { + UserID: userID2, + Platform: "openai", + DailyUsageUSD: 8.8, + WeeklyUsageUSD: 18.8, + MonthlyUsageUSD: 28.8, + DailyWindowStart: dailyStart, + WeeklyWindowStart: weeklyStart, + MonthlyWindowStart: monthlyStart, + }, + } + require.NoError(t, repo.BatchSnapshotUsage(ctx, secondBatch, now2), "second batch upsert") + + // 验证第二批覆盖:daily 应为新值,不是累加 + rec1After, err := repo.GetByUserPlatform(ctx, userID1, "anthropic") + require.NoError(t, err) + require.NotNil(t, rec1After) + require.InDelta(t, 9.9, rec1After.DailyUsageUSD, 1e-9, "user1 daily must be overwritten to 9.9 (not accumulated)") + require.InDelta(t, 19.9, rec1After.WeeklyUsageUSD, 1e-9, "user1 weekly must be overwritten to 19.9") + require.InDelta(t, 29.9, rec1After.MonthlyUsageUSD, 1e-9, "user1 monthly must be overwritten to 29.9") + + rec2After, err := repo.GetByUserPlatform(ctx, userID2, "openai") + require.NoError(t, err) + require.NotNil(t, rec2After) + require.InDelta(t, 8.8, rec2After.DailyUsageUSD, 1e-9, "user2 daily must be overwritten to 8.8 (not accumulated)") + require.InDelta(t, 18.8, rec2After.WeeklyUsageUSD, 1e-9, "user2 weekly must be overwritten to 18.8") + require.InDelta(t, 28.8, rec2After.MonthlyUsageUSD, 1e-9, "user2 monthly must be overwritten to 28.8") +} diff --git a/backend/internal/repository/user_platform_quota_service_adapter.go b/backend/internal/repository/user_platform_quota_service_adapter.go index 7495cd26..5240bb54 100644 --- a/backend/internal/repository/user_platform_quota_service_adapter.go +++ b/backend/internal/repository/user_platform_quota_service_adapter.go @@ -94,6 +94,29 @@ func (a *userPlatformQuotaServiceAdapter) ResetExpiredWindow(ctx context.Context return err } +// BatchSnapshotUsage 转换 []service.UserPlatformQuotaSnapshot → []UserPlatformQuotaSnapshot, +// 调底层 repo,并将 repository FK sentinel 包装为 service sentinel。 +func (a *userPlatformQuotaServiceAdapter) BatchSnapshotUsage(ctx context.Context, snapshots []service.UserPlatformQuotaSnapshot, now time.Time) error { + repoSnaps := make([]UserPlatformQuotaSnapshot, len(snapshots)) + for i, s := range snapshots { + repoSnaps[i] = UserPlatformQuotaSnapshot{ + UserID: s.UserID, + Platform: s.Platform, + DailyUsageUSD: s.DailyUsageUSD, + WeeklyUsageUSD: s.WeeklyUsageUSD, + MonthlyUsageUSD: s.MonthlyUsageUSD, + DailyWindowStart: s.DailyWindowStart, + WeeklyWindowStart: s.WeeklyWindowStart, + MonthlyWindowStart: s.MonthlyWindowStart, + } + } + err := a.inner.BatchSnapshotUsage(ctx, repoSnaps, now) + if errors.Is(err, ErrUserPlatformQuotaFKViolation) { + return fmt.Errorf("%w: %v", service.ErrUserPlatformQuotaFKViolation, err) + } + return err +} + // genericUserPlatformQuotaAdapter 通过通用接口适配(用于测试 fake 或非标准实现)。 type genericUserPlatformQuotaAdapter struct { inner UserPlatformQuotaRepository @@ -167,6 +190,29 @@ func (a *genericUserPlatformQuotaAdapter) ResetExpiredWindow(ctx context.Context return err } +// BatchSnapshotUsage 转换 []service.UserPlatformQuotaSnapshot → []UserPlatformQuotaSnapshot(通用 adapter), +// 并将 repository FK sentinel 包装为 service sentinel。 +func (a *genericUserPlatformQuotaAdapter) BatchSnapshotUsage(ctx context.Context, snapshots []service.UserPlatformQuotaSnapshot, now time.Time) error { + repoSnaps := make([]UserPlatformQuotaSnapshot, len(snapshots)) + for i, s := range snapshots { + repoSnaps[i] = UserPlatformQuotaSnapshot{ + UserID: s.UserID, + Platform: s.Platform, + DailyUsageUSD: s.DailyUsageUSD, + WeeklyUsageUSD: s.WeeklyUsageUSD, + MonthlyUsageUSD: s.MonthlyUsageUSD, + DailyWindowStart: s.DailyWindowStart, + WeeklyWindowStart: s.WeeklyWindowStart, + MonthlyWindowStart: s.MonthlyWindowStart, + } + } + err := a.inner.BatchSnapshotUsage(ctx, repoSnaps, now) + if errors.Is(err, ErrUserPlatformQuotaFKViolation) { + return fmt.Errorf("%w: %v", service.ErrUserPlatformQuotaFKViolation, err) + } + return err +} + // toServiceRecord 将 repository.UserPlatformQuotaRecord 转换为 service.UserPlatformQuotaRecord。 func toServiceRecord(rec *UserPlatformQuotaRecord) *service.UserPlatformQuotaRecord { return &service.UserPlatformQuotaRecord{ diff --git a/backend/internal/service/admin_service_delete_test.go b/backend/internal/service/admin_service_delete_test.go index d01b11e6..2aae73a9 100644 --- a/backend/internal/service/admin_service_delete_test.go +++ b/backend/internal/service/admin_service_delete_test.go @@ -471,10 +471,22 @@ func (s *billingCacheStub) DeleteUserPlatformQuotaCache(ctx context.Context, use panic("unexpected DeleteUserPlatformQuotaCache call") } -func (s *billingCacheStub) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration) error { +func (s *billingCacheStub) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration, markDirty bool) error { panic("unexpected IncrUserPlatformQuotaUsageCache call") } +func (s *billingCacheStub) PopDirtyUserPlatformQuotaKeys(ctx context.Context, n int) ([]UserPlatformQuotaKey, error) { + panic("unexpected PopDirtyUserPlatformQuotaKeys call") +} + +func (s *billingCacheStub) ReaddDirtyUserPlatformQuotaKeys(ctx context.Context, keys []UserPlatformQuotaKey) error { + panic("unexpected ReaddDirtyUserPlatformQuotaKeys call") +} + +func (s *billingCacheStub) BatchGetUserPlatformQuotaCache(ctx context.Context, keys []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) { + panic("unexpected BatchGetUserPlatformQuotaCache call") +} + func waitForInvalidations(t *testing.T, ch <-chan subscriptionInvalidateCall, expected int) []subscriptionInvalidateCall { t.Helper() calls := make([]subscriptionInvalidateCall, 0, expected) diff --git a/backend/internal/service/auth_service_platform_quota_test.go b/backend/internal/service/auth_service_platform_quota_test.go index f58dc48c..46069814 100644 --- a/backend/internal/service/auth_service_platform_quota_test.go +++ b/backend/internal/service/auth_service_platform_quota_test.go @@ -43,6 +43,10 @@ func (f *fakeInsertRecorder) ResetExpiredWindow(_ context.Context, _ int64, _ st return nil } +func (f *fakeInsertRecorder) BatchSnapshotUsage(_ context.Context, _ []UserPlatformQuotaSnapshot, _ time.Time) error { + return nil +} + func TestSnapshotPlatformQuotaDefaults_PassesToRepoBulkInsert(t *testing.T) { fakeRepo := &fakeInsertRecorder{} s := &AuthService{userPlatformQuotaRepo: fakeRepo} diff --git a/backend/internal/service/auth_service_register_test.go b/backend/internal/service/auth_service_register_test.go index a7c0d260..2ee9f21a 100644 --- a/backend/internal/service/auth_service_register_test.go +++ b/backend/internal/service/auth_service_register_test.go @@ -105,6 +105,10 @@ func (s *userPlatformQuotaRepoStub) ResetExpiredWindow(context.Context, int64, s panic("unexpected ResetExpiredWindow call") } +func (s *userPlatformQuotaRepoStub) BatchSnapshotUsage(_ context.Context, _ []UserPlatformQuotaSnapshot, _ time.Time) error { + return nil +} + func (s *defaultSubscriptionAssignerStub) AssignOrExtendSubscription(_ context.Context, input *AssignSubscriptionInput) (*UserSubscription, bool, error) { if input != nil { s.calls = append(s.calls, *input) diff --git a/backend/internal/service/billing_cache_service.go b/backend/internal/service/billing_cache_service.go index 8a5172f4..b734fab1 100644 --- a/backend/internal/service/billing_cache_service.go +++ b/backend/internal/service/billing_cache_service.go @@ -689,7 +689,8 @@ func (s *BillingCacheService) IncrementUserPlatformQuotaUsage(userID int64, plat ctx, cancel := context.WithTimeout(context.Background(), cacheWriteTimeout) defer cancel() ttl := time.Duration(s.cfg.Billing.UserPlatformQuotaCacheTTLSeconds) * time.Second - if err := s.cache.IncrUserPlatformQuotaUsageCache(ctx, userID, platform, cost, ttl); err != nil { + markDirty := s.cfg.Database.UserPlatformQuotaFlusherEnabled + if err := s.cache.IncrUserPlatformQuotaUsageCache(ctx, userID, platform, cost, ttl, markDirty); err != nil { logger.LegacyPrintf("service.billing_cache", "ALERT: incr user platform quota cache failed user=%d platform=%s cost=%f: %v", userID, platform, cost, err) @@ -1310,3 +1311,20 @@ func monthlyQuotaWindowExpired(start *time.Time, now time.Time) bool { } return now.Sub(*start) >= 30*24*time.Hour } + +// HasUserPlatformQuotaLimit 判断该 user×platform 是否设了任一非 nil limit。 +// 写入点守卫:无 limit 直接跳过 Redis 写 + 脏集标记,消除无谓写入。 +// fail-safe:任何不确定(simple 模式除外)都返回 true 维持写入。 +func (s *BillingCacheService) HasUserPlatformQuotaLimit(ctx context.Context, userID int64, platform string) bool { + if s.cfg.RunMode == config.RunModeSimple { + return false + } + if s.cache == nil { + return true + } + entry, ok, err := s.cache.GetUserPlatformQuotaCache(ctx, userID, platform) + if err != nil || !ok || entry == nil { + return true + } + return entry.DailyLimitUSD != nil || entry.WeeklyLimitUSD != nil || entry.MonthlyLimitUSD != nil +} diff --git a/backend/internal/service/billing_cache_service_singleflight_test.go b/backend/internal/service/billing_cache_service_singleflight_test.go index b443d97e..235b13a6 100644 --- a/backend/internal/service/billing_cache_service_singleflight_test.go +++ b/backend/internal/service/billing_cache_service_singleflight_test.go @@ -79,10 +79,22 @@ func (s *billingCacheMissStub) DeleteUserPlatformQuotaCache(ctx context.Context, return nil } -func (s *billingCacheMissStub) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration) error { +func (s *billingCacheMissStub) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration, markDirty bool) error { return nil } +func (s *billingCacheMissStub) PopDirtyUserPlatformQuotaKeys(ctx context.Context, n int) ([]UserPlatformQuotaKey, error) { + return nil, nil +} + +func (s *billingCacheMissStub) ReaddDirtyUserPlatformQuotaKeys(ctx context.Context, keys []UserPlatformQuotaKey) error { + return nil +} + +func (s *billingCacheMissStub) BatchGetUserPlatformQuotaCache(ctx context.Context, keys []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) { + return nil, nil +} + type balanceLoadUserRepoStub struct { mockUserRepo calls atomic.Int64 diff --git a/backend/internal/service/billing_cache_service_test.go b/backend/internal/service/billing_cache_service_test.go index bcd086fa..c344b417 100644 --- a/backend/internal/service/billing_cache_service_test.go +++ b/backend/internal/service/billing_cache_service_test.go @@ -80,10 +80,22 @@ func (b *billingCacheWorkerStub) DeleteUserPlatformQuotaCache(ctx context.Contex return nil } -func (b *billingCacheWorkerStub) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration) error { +func (b *billingCacheWorkerStub) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration, markDirty bool) error { return nil } +func (b *billingCacheWorkerStub) PopDirtyUserPlatformQuotaKeys(ctx context.Context, n int) ([]UserPlatformQuotaKey, error) { + return nil, nil +} + +func (b *billingCacheWorkerStub) ReaddDirtyUserPlatformQuotaKeys(ctx context.Context, keys []UserPlatformQuotaKey) error { + return nil +} + +func (b *billingCacheWorkerStub) BatchGetUserPlatformQuotaCache(ctx context.Context, keys []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) { + return nil, nil +} + func TestBillingCacheServiceQueueHighLoad(t *testing.T) { cache := &billingCacheWorkerStub{} svc := NewBillingCacheService(cache, nil, nil, nil, nil, nil, &config.Config{}, nil) diff --git a/backend/internal/service/billing_cache_service_user_platform_quota_test.go b/backend/internal/service/billing_cache_service_user_platform_quota_test.go index 674aa9a5..a82c0e00 100644 --- a/backend/internal/service/billing_cache_service_user_platform_quota_test.go +++ b/backend/internal/service/billing_cache_service_user_platform_quota_test.go @@ -20,14 +20,15 @@ type fakeIncrCache struct { } type incrCall struct { - userID int64 - platform string - cost float64 - ttl time.Duration + userID int64 + platform string + cost float64 + ttl time.Duration + markDirty bool } -func (f *fakeIncrCache) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration) error { - f.calls = append(f.calls, incrCall{userID, platform, cost, ttl}) +func (f *fakeIncrCache) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration, markDirty bool) error { + f.calls = append(f.calls, incrCall{userID, platform, cost, ttl, markDirty}) return nil } @@ -49,10 +50,10 @@ func TestIncrementUserPlatformQuotaUsage_SyncCallsCache(t *testing.T) { if len(fake.calls) != 2 { t.Fatalf("expected 2 incr calls, got %d", len(fake.calls)) } - if fake.calls[0] != (incrCall{101, "anthropic", 0.25, 120 * time.Second}) { + if fake.calls[0] != (incrCall{userID: 101, platform: "anthropic", cost: 0.25, ttl: 120 * time.Second, markDirty: false}) { t.Errorf("call[0] = %+v", fake.calls[0]) } - if fake.calls[1] != (incrCall{101, "openai", 0.50, 120 * time.Second}) { + if fake.calls[1] != (incrCall{userID: 101, platform: "openai", cost: 0.50, ttl: 120 * time.Second, markDirty: false}) { t.Errorf("call[1] = %+v", fake.calls[1]) } } @@ -88,7 +89,11 @@ func (f *fakeQuotaRepo) ResetExpiredWindow(_ context.Context, _ int64, _ string, return nil } -// fakeFullCache 同时支持 Get + Set + Incr + Delete。 +func (f *fakeQuotaRepo) BatchSnapshotUsage(_ context.Context, _ []UserPlatformQuotaSnapshot, _ time.Time) error { + return nil +} + +// fakeFullCache 同时支持 Get + Set + Incr + Delete + Pop/Readd/BatchGet(脏集读写)。 // mu 保护 entry 和 deleteCalls,防止异步 goroutine 与主 goroutine 之间的 data race。 type fakeFullCache struct { BillingCache @@ -99,6 +104,8 @@ type fakeFullCache struct { lastSetTTL time.Duration // 最近一次 Set 的 ttl getErr error // 非 nil 时 Get 先返回 (nil,false,getErr) setErr error // 非 nil 时 Set 返回该 err(setCalls 仍+1) + // dirty 模拟脏集,供 flusher 测试使用。 + dirty map[UserPlatformQuotaKey]struct{} } // getDeleteCalls 线程安全地读取 deleteCalls。 @@ -161,6 +168,48 @@ func (f *fakeFullCache) DeleteUserPlatformQuotaCache(_ context.Context, _ int64, return nil } +func (f *fakeFullCache) PopDirtyUserPlatformQuotaKeys(_ context.Context, n int) ([]UserPlatformQuotaKey, error) { + f.mu.Lock() + defer f.mu.Unlock() + if len(f.dirty) == 0 { + return nil, nil + } + keys := make([]UserPlatformQuotaKey, 0, n) + for k := range f.dirty { + if len(keys) >= n { + break + } + keys = append(keys, k) + delete(f.dirty, k) + } + return keys, nil +} + +func (f *fakeFullCache) ReaddDirtyUserPlatformQuotaKeys(_ context.Context, keys []UserPlatformQuotaKey) error { + f.mu.Lock() + defer f.mu.Unlock() + if f.dirty == nil { + f.dirty = make(map[UserPlatformQuotaKey]struct{}) + } + for _, k := range keys { + f.dirty[k] = struct{}{} + } + return nil +} + +// BatchGetUserPlatformQuotaCache 对每个 key 返回 f.entry(MISS → nil), +// 保持与输入 keys 顺序/长度对齐。注意此处所有 key 共享同一个 entry, +// 仅用于测试场景。 +func (f *fakeFullCache) BatchGetUserPlatformQuotaCache(_ context.Context, keys []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) { + f.mu.Lock() + defer f.mu.Unlock() + results := make([]*UserPlatformQuotaCacheEntry, len(keys)) + for i := range keys { + results[i] = f.entry + } + return results, nil +} + func newServiceForPreflight(t *testing.T, repo UserPlatformQuotaRepository, cache BillingCache) *BillingCacheService { t.Helper() cfg := &config.Config{} @@ -713,3 +762,71 @@ func TestCheckUserPlatformQuotaEligibility_SentinelCrossDay_NoRefresh(t *testing t.Errorf("sentinel cross-window must NOT trigger refresh SetCache, got %d calls", cache.getSetCalls()) } } + +// ── TestHasUserPlatformQuotaLimit ──────────────────────────────────────────── + +func TestHasUserPlatformQuotaLimit(t *testing.T) { + daily := 5.0 + + tests := []struct { + name string + setup func() *BillingCacheService + want bool + }{ + { + name: "has_limit", + setup: func() *BillingCacheService { + entry := &UserPlatformQuotaCacheEntry{DailyLimitUSD: &daily} + svc := newServiceForPreflight(t, &fakeQuotaRepo{}, &fakeFullCache{entry: entry}) + return svc + }, + want: true, + }, + { + name: "sentinel_no_limit", + setup: func() *BillingCacheService { + entry := &UserPlatformQuotaCacheEntry{} // 三个 limit 字段全 nil + svc := newServiceForPreflight(t, &fakeQuotaRepo{}, &fakeFullCache{entry: entry}) + return svc + }, + want: false, + }, + { + name: "cache_miss", + setup: func() *BillingCacheService { + // entry==nil → GetUserPlatformQuotaCache 返回 (nil,false,nil) + svc := newServiceForPreflight(t, &fakeQuotaRepo{}, &fakeFullCache{}) + return svc + }, + want: true, // fail-safe + }, + { + name: "redis_err", + setup: func() *BillingCacheService { + svc := newServiceForPreflight(t, &fakeQuotaRepo{}, &fakeFullCache{getErr: errors.New("redis down")}) + return svc + }, + want: true, // fail-safe + }, + { + name: "simple_mode", + setup: func() *BillingCacheService { + entry := &UserPlatformQuotaCacheEntry{DailyLimitUSD: &daily} + svc := newServiceForPreflight(t, &fakeQuotaRepo{}, &fakeFullCache{entry: entry}) + svc.cfg.RunMode = config.RunModeSimple + return svc + }, + want: false, // simple 模式始终跳过 + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + svc := tt.setup() + got := svc.HasUserPlatformQuotaLimit(context.Background(), 1, "anthropic") + if got != tt.want { + t.Errorf("HasUserPlatformQuotaLimit() = %v, want %v", got, tt.want) + } + }) + } +} diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index 940a827d..6b1438e8 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -21,6 +21,12 @@ type APIKeyRateLimitCacheData struct { Window7d int64 `json:"window_7d"` } +// UserPlatformQuotaKey 标识一个 user×platform,用于脏集出入与批量读。 +type UserPlatformQuotaKey struct { + UserID int64 + Platform string +} + // UserPlatformQuotaCacheEntry Redis hash 反序列化结果。 // // SchemaVersion 用于向后兼容: @@ -72,7 +78,13 @@ type BillingCache interface { SetUserPlatformQuotaCache(ctx context.Context, userID int64, platform string, entry *UserPlatformQuotaCacheEntry, ttl time.Duration) error DeleteUserPlatformQuotaCache(ctx context.Context, userID int64, platform string) error // IncrUserPlatformQuotaUsageCache 在缓存命中时累加用量;缓存未命中(key 不存在)静默返回 nil。 - IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration) error + // markDirty=true 时将该 key 的 member 写入 Redis 脏集,供 flusher 批量回写 DB。 + IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration, markDirty bool) error + + // 脏集读写,供 flusher 使用。 + PopDirtyUserPlatformQuotaKeys(ctx context.Context, n int) ([]UserPlatformQuotaKey, error) + ReaddDirtyUserPlatformQuotaKeys(ctx context.Context, keys []UserPlatformQuotaKey) error + BatchGetUserPlatformQuotaCache(ctx context.Context, keys []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) } // ModelPricing 模型价格配置(per-token价格,与LiteLLM格式一致) diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index effa803a..f807f3ec 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -96,11 +96,13 @@ var ( modelsListCacheMissTotal atomic.Int64 modelsListCacheStoreTotal atomic.Int64 + // Deprecated: flusher_enabled=true 后不再增长(仅 flag=false 降级直写路径使用);新主路径见 FlusherMetrics。remove after 2026-09。 // userPlatformQuotaDBIncrErrorTotal 统计 finalizePostUsageBilling 异步 goroutine // 中 IncrementUsageWithReset 失败次数。Redis 已成功累加 + DB 写失败意味着 // Redis cache TTL 过期或被清后该笔 cost 会丢失(与实际消费偏差)。 // oncall 通过 GatewayUserPlatformQuotaIncrStats() 暴露给 ops 面板做阈值告警。 userPlatformQuotaDBIncrErrorTotal atomic.Int64 + // Deprecated: flusher_enabled=true 后不再增长(仅 flag=false 降级直写路径使用);新主路径见 FlusherMetrics。remove after 2026-09。 // userPlatformQuotaDBIncrLegacyErrorTotal 统计 legacy postUsageBilling // (applyUsageBilling 在 repo==nil 时 fallback)路径下的失败次数; // 与 DB Incr 失败分开计数,便于区分"主路径暂时故障"vs"基础设施长期未配齐"。 @@ -141,6 +143,23 @@ func GatewayUserPlatformQuotaIncrStats() (mainPathErr, legacyPathErr, sentinelSe userPlatformQuotaSentinelSetCacheErrorTotal.Load() } +// GatewayUserPlatformQuotaFlusherStats 暴露 flusher 运行指标供 ops/health 面板查询。 +func GatewayUserPlatformQuotaFlusherStats(f *UserPlatformQuotaUsageFlusher) map[string]int64 { + if f == nil || f.metrics == nil { + return nil + } + m := f.metrics + return map[string]int64{ + "flush_success": m.FlushSuccessTotal.Load(), + "flush_error": m.FlushErrorTotal.Load(), + "flush_batch_size": m.FlushBatchSizeTotal.Load(), + "flush_latency_ms_max": m.FlushLatencyMsMax.Load(), + "dirty_readd": m.DirtyReaddTotal.Load(), + "dirty_lost": m.DirtyLostTotal.Load(), + "flush_fk_violation": m.FlushFKViolationTotal.Load(), + } +} + func openAIStreamEventIsTerminal(data string) bool { trimmed := strings.TrimSpace(data) if trimmed == "" { @@ -8234,18 +8253,23 @@ func postUsageBilling(ctx context.Context, p *postUsageBillingParams, deps *bill } } - // Platform quota DB-only 累加(与 finalizePostUsageBilling 行为对齐的兜底): - // - 仅对 standard(余额)模式生效;订阅模式豁免 - // - 直接走 DB,不经 Redis Incr 队列:legacy 路径在 repo==nil(仓库未注入) - // 时被触发,此时整套 billing repo 都不可用,没有"双队列"风险 - // - 失败仅记 ALERT log + counter,不阻断主扣费流程;与正常路径一致 - // - // 历史背景:原 legacy path 完全跳过此累加,导致部署中如果 repo 偶然为 nil - // 时用户消费可绕过 platform quota,存在静默资金风险。 + // Platform quota 累加(legacy 兜底路径):仅对 standard(余额)模式生效;订阅模式豁免;仅对有 limit 的用户写 + // - HasUserPlatformQuotaLimit 守卫:与正常路径对齐,无 limit 公司跳过 + // - 新增 Redis 同步写:enforcement 走 Redis,legacy 路径也必须同步写,否则 preflight 看不到消费 + // - flusher_enabled=false(降级):保留原有同步直写 DB + // - flusher_enabled=true:跳过直写 DB,由 flusher 异步批量刷(markDirty 在 IncrementUserPlatformQuotaUsage 内部完成) + // - 失败仅记 ALERT log + counter,不阻断主扣费流程 if !p.IsSubscriptionBill && p.Platform != "" && cost.ActualCost > 0 && p.User != nil && deps.userPlatformQuotaRepo != nil { - if err := deps.userPlatformQuotaRepo.IncrementUsageWithReset(billingCtx, p.User.ID, p.Platform, cost.ActualCost, time.Now().UTC()); err != nil { - userPlatformQuotaDBIncrLegacyErrorTotal.Add(1) - logger.LegacyPrintf("service.gateway", "ALERT: legacy incr user platform quota DB failed user=%d platform=%s cost=%f: %v", p.User.ID, p.Platform, cost.ActualCost, err) + if deps.billingCacheService.HasUserPlatformQuotaLimit(billingCtx, p.User.ID, p.Platform) { + deps.billingCacheService.IncrementUserPlatformQuotaUsage(p.User.ID, p.Platform, cost.ActualCost) + if deps.cfg == nil || !deps.cfg.Database.UserPlatformQuotaFlusherEnabled { + // 降级路径:flusher 未启用时保留原有同步直写 DB + if err := deps.userPlatformQuotaRepo.IncrementUsageWithReset(billingCtx, p.User.ID, p.Platform, cost.ActualCost, time.Now().UTC()); err != nil { + userPlatformQuotaDBIncrLegacyErrorTotal.Add(1) + logger.LegacyPrintf("service.gateway", "ALERT: legacy incr user platform quota DB failed user=%d platform=%s cost=%f: %v", p.User.ID, p.Platform, cost.ActualCost, err) + } + } + // flusher_enabled=true:不直写 DB,flusher 异步批量刷 } } @@ -8395,30 +8419,38 @@ func finalizePostUsageBilling(ctx context.Context, p *postUsageBillingParams, de deps.deferredService.ScheduleLastUsedUpdate(p.Account.ID) - // Platform quota 累加:仅在 standard(余额)模式生效;订阅模式豁免 - // Redis 同步写 + DB 异步持久化: + // Platform quota 累加:仅在 standard(余额)模式生效;订阅模式豁免;仅对有 limit 的用户写 + // Redis 同步写 + DB 异步持久化(flag=false 降级)或 flusher 异步刷(flag=true): + // - HasUserPlatformQuotaLimit 守卫:无 limit 的公司跳过,避免无效写入 + 浪费 Redis 容量 // - Redis 同步:确保下次 preflight 立即看到最新 usage,把 TOCTOU 超支窗口 // 限制在并发 in-flight 请求数量内(旧实现的异步入队会让超支无限累积直到 worker 处理) - // - DB 异步:在独立 goroutine 中走 detached context,失败用 ALERT log 触发 oncall 对账 + // - DB 异步(flusher_enabled=false):在独立 goroutine 中走 detached context,失败用 ALERT log 触发 oncall 对账 + // - flusher_enabled=true:不直写 DB,由 flusher 异步批量刷(markDirty 已在 IncrementUserPlatformQuotaUsage 内部完成) if !p.IsSubscriptionBill && p.Platform != "" && p.Cost.ActualCost > 0 && p.User != nil && deps.userPlatformQuotaRepo != nil { - deps.billingCacheService.IncrementUserPlatformQuotaUsage(p.User.ID, p.Platform, p.Cost.ActualCost) - dbCtx, dbCancel := detachUpstreamContext(ctx) - userID, platform, cost := p.User.ID, p.Platform, p.Cost.ActualCost - go func() { - defer func() { - if r := recover(); r != nil { - logger.LegacyPrintf("service.gateway", "ALERT: panic in user platform quota incr goroutine user=%d platform=%s: %v", userID, platform, r) - } - }() - defer dbCancel() - if err := deps.userPlatformQuotaRepo.IncrementUsageWithReset(dbCtx, userID, platform, cost, time.Now().UTC()); err != nil { - // 失败计数器:暴露给 GatewayUserPlatformQuotaIncrStats(),由 ops 面板做斜率告警。 - userPlatformQuotaDBIncrErrorTotal.Add(1) - // ALERT 级别:DB 持久化失败意味着 Redis cache 失效后该笔 cost 永久丢失, - // 用户配额视图与实际消费会偏差,oncall 需要据此对账或人工补录。 - logger.LegacyPrintf("service.gateway", "ALERT: incr user platform quota DB failed user=%d platform=%s cost=%f: %v", userID, platform, cost, err) + if deps.billingCacheService.HasUserPlatformQuotaLimit(ctx, p.User.ID, p.Platform) { + deps.billingCacheService.IncrementUserPlatformQuotaUsage(p.User.ID, p.Platform, p.Cost.ActualCost) + if deps.cfg == nil || !deps.cfg.Database.UserPlatformQuotaFlusherEnabled { + // 降级路径:flusher 未启用时保留原有异步直写 DB + dbCtx, dbCancel := detachUpstreamContext(ctx) + userID, platform, cost := p.User.ID, p.Platform, p.Cost.ActualCost + go func() { + defer func() { + if r := recover(); r != nil { + logger.LegacyPrintf("service.gateway", "ALERT: panic in user platform quota incr goroutine user=%d platform=%s: %v", userID, platform, r) + } + }() + defer dbCancel() + if err := deps.userPlatformQuotaRepo.IncrementUsageWithReset(dbCtx, userID, platform, cost, time.Now().UTC()); err != nil { + // 失败计数器:暴露给 GatewayUserPlatformQuotaIncrStats(),由 ops 面板做斜率告警。 + userPlatformQuotaDBIncrErrorTotal.Add(1) + // ALERT 级别:DB 持久化失败意味着 Redis cache 失效后该笔 cost 永久丢失, + // 用户配额视图与实际消费会偏差,oncall 需要据此对账或人工补录。 + logger.LegacyPrintf("service.gateway", "ALERT: incr user platform quota DB failed user=%d platform=%s cost=%f: %v", userID, platform, cost, err) + } + }() } - }() + // flusher_enabled=true:不直写 DB,flusher 异步批量刷 + } } // Notification checks run async — all parameters are already captured, @@ -8533,6 +8565,7 @@ type billingDeps struct { deferredService *DeferredService balanceNotifyService *BalanceNotifyService userPlatformQuotaRepo UserPlatformQuotaRepository + cfg *config.Config } func (s *GatewayService) billingDeps() *billingDeps { @@ -8544,6 +8577,7 @@ func (s *GatewayService) billingDeps() *billingDeps { deferredService: s.deferredService, balanceNotifyService: s.balanceNotifyService, userPlatformQuotaRepo: s.userPlatformQuotaRepo, + cfg: s.cfg, } } diff --git a/backend/internal/service/user_platform_quota_flusher.go b/backend/internal/service/user_platform_quota_flusher.go new file mode 100644 index 00000000..3ee23d2c --- /dev/null +++ b/backend/internal/service/user_platform_quota_flusher.go @@ -0,0 +1,267 @@ +package service + +import ( + "context" + "errors" + "sync/atomic" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" +) + +// quotaDirtyCache 是 flusher 依赖的窄接口(来自 BillingCache)。 +type quotaDirtyCache interface { + PopDirtyUserPlatformQuotaKeys(ctx context.Context, n int) ([]UserPlatformQuotaKey, error) + ReaddDirtyUserPlatformQuotaKeys(ctx context.Context, keys []UserPlatformQuotaKey) error + BatchGetUserPlatformQuotaCache(ctx context.Context, keys []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) +} + +// quotaSnapshotWriter 是 flusher 依赖的 DB 写入窄接口。 +// 使用 service 层的 UserPlatformQuotaSnapshot,避免与 repository 包形成循环依赖; +// 实际实现由 repository adapter 在 B7 注入。 +type quotaSnapshotWriter interface { + BatchSnapshotUsage(ctx context.Context, snapshots []UserPlatformQuotaSnapshot, now time.Time) error +} + +// FlusherMetrics 记录 flusher 运行时指标(原子量,零值可用)。 +type FlusherMetrics struct { + FlushSuccessTotal atomic.Int64 + FlushErrorTotal atomic.Int64 + FlushBatchSizeTotal atomic.Int64 + FlushLatencyMsMax atomic.Int64 + DirtyReaddTotal atomic.Int64 + // DirtyLostTotal:Readd 失败导致脏 key 丢失——已 SPOP+主操作失败+Readd 也失败; + // Redis 仍权威,活跃 key 下次 SADD 自愈。 + DirtyLostTotal atomic.Int64 + FlushFKViolationTotal atomic.Int64 +} + +// flusherMaxBatchesPerTick 单次 tick 最多消费的批数,防止 tick 执行时间过长。 +const flusherMaxBatchesPerTick = 16 + +// maxFlushBatchSize 限制单批行数,必须 ≤ repository.BatchSnapshotUsage 的 batchRows(6000), +// 以保证单次 flush 的 snapshots 仅生成一条 UPSERT(单事务原子)。两处需手动保持一致。 +const maxFlushBatchSize = 6000 + +// defaultFlushBatchSize 是配置 flush_batch_size 非法(≤0)时的回退值。 +const defaultFlushBatchSize = 1000 + +// UserPlatformQuotaUsageFlusher 将 Redis 脏集快照定期批量写入 DB。 +// 不维护任何 delta/in-process 状态;每批读取 Redis 当前绝对值覆盖写入。 +type UserPlatformQuotaUsageFlusher struct { + cache quotaDirtyCache + quotaRepo quotaSnapshotWriter + timingWheel *TimingWheelService + // enabled 对应 flusher_enabled 配置;false 时 Start() 不注册定时器。 + enabled bool + interval time.Duration + batchSize int + flushTimeout time.Duration + metrics *FlusherMetrics + stopped atomic.Bool +} + +// NewUserPlatformQuotaUsageFlusher 创建 UserPlatformQuotaUsageFlusher。 +// cache(BillingCache) 隐式满足 quotaDirtyCache;quotaRepo(UserPlatformQuotaRepository) 隐式满足 quotaSnapshotWriter。 +func NewUserPlatformQuotaUsageFlusher(cfg *config.Config, cache BillingCache, quotaRepo UserPlatformQuotaRepository, tw *TimingWheelService) *UserPlatformQuotaUsageFlusher { + batchSize := cfg.Database.UserPlatformQuotaFlushBatchSize + if batchSize <= 0 { + batchSize = defaultFlushBatchSize + } + if batchSize > maxFlushBatchSize { + logger.LegacyPrintf("quota_flusher", + "[QuotaFlusher] flush_batch_size %d 超过上限 %d,已 clamp(避免 BatchSnapshotUsage 多子批非原子)", + cfg.Database.UserPlatformQuotaFlushBatchSize, maxFlushBatchSize) + batchSize = maxFlushBatchSize + } + interval := time.Duration(cfg.Database.UserPlatformQuotaFlushIntervalMs) * time.Millisecond + if interval <= 0 { + logger.LegacyPrintf("quota_flusher", "[QuotaFlusher] flush_interval_ms %d 非法,回退 2000ms", cfg.Database.UserPlatformQuotaFlushIntervalMs) + interval = 2 * time.Second + } + return &UserPlatformQuotaUsageFlusher{ + cache: cache, + quotaRepo: quotaRepo, + timingWheel: tw, + enabled: cfg.Database.UserPlatformQuotaFlusherEnabled, + interval: interval, + batchSize: batchSize, + flushTimeout: 3 * time.Second, + metrics: &FlusherMetrics{}, + } +} + +// updateLatencyMax 用 CAS 单调更新最大延迟。 +func (s *UserPlatformQuotaUsageFlusher) updateLatencyMax(ms int64) { + for { + old := s.metrics.FlushLatencyMsMax.Load() + if ms <= old { + return + } + if s.metrics.FlushLatencyMsMax.CompareAndSwap(old, ms) { + return + } + } +} + +// readdOrCountLost 尝试把 keys 回填脏集:成功计 DirtyReaddTotal,失败计 DirtyLostTotal 并 ALERT。 +func (s *UserPlatformQuotaUsageFlusher) readdOrCountLost(ctx context.Context, keys []UserPlatformQuotaKey, stage string) { + if err := s.cache.ReaddDirtyUserPlatformQuotaKeys(ctx, keys); err != nil { + s.metrics.DirtyLostTotal.Add(int64(len(keys))) + logger.LegacyPrintf("quota_flusher", "[QuotaFlusher] ALERT: Readd after %s failed, %d keys 丢出脏集(DB 镜像缺这批,Redis 仍权威,活跃 key 下次 SADD 自愈): %v", stage, len(keys), err) + return + } + s.metrics.DirtyReaddTotal.Add(int64(len(keys))) +} + +// flushOneBatch 处理单批:Pop → BatchGet → 组装 snaps → BatchSnapshotUsage。 +// 返回 (shouldContinue bool):false 表示本轮循环应停止(空集/错误/最后一批)。 +// 每次调用独立创建带 timeout 的 ctx 并 defer cancel,不会在循环中累积泄漏。 +func (s *UserPlatformQuotaUsageFlusher) flushOneBatch(parentCtx context.Context) bool { + ctx, cancel := context.WithTimeout(parentCtx, s.flushTimeout) + defer cancel() + + // 1. Pop 脏集 + keys, err := s.cache.PopDirtyUserPlatformQuotaKeys(ctx, s.batchSize) + if err != nil { + s.metrics.FlushErrorTotal.Add(1) + logger.LegacyPrintf("quota_flusher", "[QuotaFlusher] PopDirty error: %v", err) + return false + } + if len(keys) == 0 { + // 脏集已空 + return false + } + + // 2. 批量读 Redis 快照 + entries, err := s.cache.BatchGetUserPlatformQuotaCache(ctx, keys) + if err != nil { + s.metrics.FlushErrorTotal.Add(1) + s.readdOrCountLost(ctx, keys, "BatchGet") + logger.LegacyPrintf("quota_flusher", "[QuotaFlusher] BatchGet error: %v", err) + return false + } + + // 3. 组装 snapshots(MISS 或任一 WindowStart==nil → 跳过) + snaps := make([]UserPlatformQuotaSnapshot, 0, len(keys)) + for i, key := range keys { + e := entries[i] + if e == nil { + continue + } + if e.DailyWindowStart == nil || e.WeeklyWindowStart == nil || e.MonthlyWindowStart == nil { + continue + } + snaps = append(snaps, UserPlatformQuotaSnapshot{ + UserID: key.UserID, + Platform: key.Platform, + DailyUsageUSD: e.DailyUsageUSD, + WeeklyUsageUSD: e.WeeklyUsageUSD, + MonthlyUsageUSD: e.MonthlyUsageUSD, + DailyWindowStart: *e.DailyWindowStart, + WeeklyWindowStart: *e.WeeklyWindowStart, + MonthlyWindowStart: *e.MonthlyWindowStart, + }) + } + + // 4. 全部 MISS/异常跳过时 + if len(snaps) == 0 { + // 若 Pop 数量已不满一批,表示脏集将空,停止 + if len(keys) < s.batchSize { + return false + } + // 否则继续下一批(可能还有更多脏 key) + return true + } + + // 已知竞态(admin 写 × flusher 刷,仅 flusher_enabled=true 时存在): + // admin ResetExpiredWindow/UpsertForUser 是"先写 DB 再 DeleteCache"。若本批已 SPOP + BatchGet + // 读到旧 usage 快照(此刻 member 已离开脏集),而 admin 随后写 DB、本行 UPSERT 又在 admin 写之后落库, + // 则旧快照会覆盖 admin 刚写入的值;DeleteCache 后 Redis MISS,下次 preflight 从 DB 重载被覆盖的旧值。 + // 因 member 已被 SPOP,admin 侧 SREM/清脏标记无法拦截本批(故未做)。影响有限,暂列为已知取舍: + // - UpsertForUser 改 limit,而本 UPSERT 不写 limit 列 → limit 配置不受影响; + // - ResetExpiredWindow 改 usage,但 preflight windowExpired 会在窗口真正过期时自愈重置, + // 仅"强制重置未过期窗口"且与本批精确交错时短暂失效; + // - 低频 admin 操作 + 默认 flusher_enabled=false。彻底消除需 version OCC(DB 加 version 列条件 UPSERT), + // 成本高;启用 flusher 后如需强一致再评估。 + + // 5. 写入 DB + start := time.Now() + writeErr := s.quotaRepo.BatchSnapshotUsage(ctx, snaps, time.Now().UTC()) + s.updateLatencyMax(time.Since(start).Milliseconds()) + + if writeErr != nil { + if errors.Is(writeErr, ErrUserPlatformQuotaFKViolation) { + // 注意:PG FK violation 是整条 INSERT 回滚 → 整批(含同批正常用户)均未写入 DB, + // 且这些 key 已被 SPOP 出脏集、此处不 Readd。活跃 key 会在下次请求重新 SADD, + // flusher 读 Redis 当前累计绝对值刷库即自愈;低活跃 key 这轮 DB usage 偏低 + // (Redis 仍是 enforcement 权威,不受影响;DB 仅展示)。已删用户边角的接受取舍,不做逐行重试。 + // FK 违反:用户已被删除,直接丢弃不 Readd + s.metrics.FlushFKViolationTotal.Add(1) + s.metrics.FlushErrorTotal.Add(1) + logger.LegacyPrintf("quota_flusher", "[QuotaFlusher] FK violation (dropped %d snaps): %v", len(snaps), writeErr) + } else { + // 其他错误:回填脏集,保留下次重试 + s.metrics.FlushErrorTotal.Add(1) + s.readdOrCountLost(ctx, keys, "BatchSnapshotUsage") + logger.LegacyPrintf("quota_flusher", "[QuotaFlusher] BatchSnapshotUsage error: %v", writeErr) + } + return false + } + + // 6. 成功 + s.metrics.FlushSuccessTotal.Add(1) + s.metrics.FlushBatchSizeTotal.Add(int64(len(snaps))) + + // 若 Pop 数量不满一批,脏集已空,停止 + if len(keys) < s.batchSize { + return false + } + return true +} + +// flush 执行一次完整的 flush,循环消费至脏集空或达到 maxBatchesPerTick。 +func (s *UserPlatformQuotaUsageFlusher) flush() { + if s == nil { + return + } + parentCtx := context.Background() + for b := 0; b < flusherMaxBatchesPerTick; b++ { + if !s.flushOneBatch(parentCtx) { + return + } + } + // 连续消费满 flusherMaxBatchesPerTick 批仍未取空脏集:本 tick 主动让出,剩余积压留待下一 tick。 + // 记一条 log 便于 oncall 发现 distinct 活跃 key 远超 maxBatchesPerTick×batchSize(DB 镜像延迟上升); + // 可配合 Redis SCARD billing:upq:dirty 观察脏集存量。 + logger.LegacyPrintf("quota_flusher", + "[QuotaFlusher] 单 tick 达到 max batches 上限(%d × batchSize=%d),脏集仍非空,积压顺延至下一 tick", + flusherMaxBatchesPerTick, s.batchSize) +} + +// tick 是 TimingWheel 回调。若 flusher 已停止则直接返回。 +func (s *UserPlatformQuotaUsageFlusher) tick() { + if s == nil || s.stopped.Load() { + return + } + s.flush() +} + +// Start 注册定时 tick。flusher_enabled=false 时直接返回,不注册定时器。 +func (s *UserPlatformQuotaUsageFlusher) Start() { + if s == nil || !s.enabled { + return + } + s.timingWheel.ScheduleRecurring("deferred:platform_quota", s.interval, s.tick) +} + +// Stop 停止 flusher:标记 stopped → Cancel 定时器 → 执行最后一次 flush。 +func (s *UserPlatformQuotaUsageFlusher) Stop() { + if s == nil { + return + } + s.stopped.Store(true) + s.timingWheel.Cancel("deferred:platform_quota") + s.flush() +} diff --git a/backend/internal/service/user_platform_quota_flusher_test.go b/backend/internal/service/user_platform_quota_flusher_test.go new file mode 100644 index 00000000..4f734481 --- /dev/null +++ b/backend/internal/service/user_platform_quota_flusher_test.go @@ -0,0 +1,511 @@ +package service + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" +) + +// --------------------------------------------------------------------------- +// Mock: quotaDirtyCache +// --------------------------------------------------------------------------- + +type mockQuotaDirtyCache struct { + // popSequence: 第 0 次 Pop 返回 popSequence[0],之后返回 nil(空集) + popSequence [][]UserPlatformQuotaKey + popCallIdx int + + // getEntries: BatchGetUserPlatformQuotaCache 返回的 entries(与 keys 对齐) + getEntries []*UserPlatformQuotaCacheEntry + getErr error + + // readdCalled: 记录 Readd 收到的 keys(累积所有次调用) + readdCalled [][]UserPlatformQuotaKey + readdErr error +} + +func (m *mockQuotaDirtyCache) PopDirtyUserPlatformQuotaKeys(_ context.Context, _ int) ([]UserPlatformQuotaKey, error) { + if m.popCallIdx < len(m.popSequence) { + keys := m.popSequence[m.popCallIdx] + m.popCallIdx++ + return keys, nil + } + // 超出序列 → 空集(模拟脏集已清空) + return nil, nil +} + +func (m *mockQuotaDirtyCache) ReaddDirtyUserPlatformQuotaKeys(_ context.Context, keys []UserPlatformQuotaKey) error { + m.readdCalled = append(m.readdCalled, keys) + return m.readdErr +} + +func (m *mockQuotaDirtyCache) BatchGetUserPlatformQuotaCache(_ context.Context, _ []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) { + if m.getErr != nil { + return nil, m.getErr + } + return m.getEntries, nil +} + +// --------------------------------------------------------------------------- +// Mock: quotaSnapshotWriter +// --------------------------------------------------------------------------- + +type mockQuotaSnapshotWriter struct { + receivedSnaps []UserPlatformQuotaSnapshot + returnErr error +} + +func (m *mockQuotaSnapshotWriter) BatchSnapshotUsage(_ context.Context, snaps []UserPlatformQuotaSnapshot, _ time.Time) error { + m.receivedSnaps = append(m.receivedSnaps, snaps...) + return m.returnErr +} + +// --------------------------------------------------------------------------- +// Helper: 构造窗口起始时间(非 nil) +// --------------------------------------------------------------------------- + +func flusherPtrTime(t time.Time) *time.Time { return &t } + +func makeEntry(daily, weekly, monthly float64) *UserPlatformQuotaCacheEntry { + now := time.Now().UTC() + return &UserPlatformQuotaCacheEntry{ + DailyUsageUSD: daily, + WeeklyUsageUSD: weekly, + MonthlyUsageUSD: monthly, + DailyWindowStart: flusherPtrTime(now), + WeeklyWindowStart: flusherPtrTime(now), + MonthlyWindowStart: flusherPtrTime(now), + } +} + +// --------------------------------------------------------------------------- +// newTestFlusher: 直接构造 struct(跳过构造函数,B7 才注入) +// --------------------------------------------------------------------------- + +func newTestFlusher(cache quotaDirtyCache, writer quotaSnapshotWriter) *UserPlatformQuotaUsageFlusher { + return &UserPlatformQuotaUsageFlusher{ + cache: cache, + quotaRepo: writer, + timingWheel: nil, // 单测不启动 TimingWheel + interval: 5 * time.Second, + batchSize: 100, + flushTimeout: 5 * time.Second, + metrics: &FlusherMetrics{}, + } +} + +// --------------------------------------------------------------------------- +// 场景 1: PopSnapshotUpsert — 2 key + 2 个含 window 的 entry → writer 收 2 行 +// --------------------------------------------------------------------------- + +func TestFlusher_PopSnapshotUpsert(t *testing.T) { + keys := []UserPlatformQuotaKey{ + {UserID: 1, Platform: "anthropic"}, + {UserID: 2, Platform: "openai"}, + } + cache := &mockQuotaDirtyCache{ + popSequence: [][]UserPlatformQuotaKey{keys}, // 第 1 次返回 keys,之后空 + getEntries: []*UserPlatformQuotaCacheEntry{ + makeEntry(1.0, 2.0, 3.0), + makeEntry(4.0, 5.0, 6.0), + }, + } + writer := &mockQuotaSnapshotWriter{} + f := newTestFlusher(cache, writer) + + f.flush() + + if len(writer.receivedSnaps) != 2 { + t.Fatalf("expected 2 snaps, got %d", len(writer.receivedSnaps)) + } + if f.metrics.FlushBatchSizeTotal.Load() != 2 { + t.Errorf("FlushBatchSizeTotal = %d, want 2", f.metrics.FlushBatchSizeTotal.Load()) + } + if f.metrics.FlushSuccessTotal.Load() != 1 { + t.Errorf("FlushSuccessTotal = %d, want 1", f.metrics.FlushSuccessTotal.Load()) + } + if f.metrics.FlushErrorTotal.Load() != 0 { + t.Errorf("FlushErrorTotal = %d, want 0", f.metrics.FlushErrorTotal.Load()) + } +} + +// --------------------------------------------------------------------------- +// 场景 2: MissKeySkipped — 2 key,BatchGet 返回 [entry, nil] → 只刷 1 行,nil 跳过,不 Readd +// --------------------------------------------------------------------------- + +func TestFlusher_MissKeySkipped(t *testing.T) { + keys := []UserPlatformQuotaKey{ + {UserID: 1, Platform: "anthropic"}, + {UserID: 2, Platform: "openai"}, + } + cache := &mockQuotaDirtyCache{ + popSequence: [][]UserPlatformQuotaKey{keys}, + getEntries: []*UserPlatformQuotaCacheEntry{ + makeEntry(1.0, 2.0, 3.0), + nil, // MISS + }, + } + writer := &mockQuotaSnapshotWriter{} + f := newTestFlusher(cache, writer) + + f.flush() + + if len(writer.receivedSnaps) != 1 { + t.Fatalf("expected 1 snap, got %d", len(writer.receivedSnaps)) + } + if writer.receivedSnaps[0].UserID != 1 { + t.Errorf("expected snap for UserID=1, got %d", writer.receivedSnaps[0].UserID) + } + if len(cache.readdCalled) != 0 { + t.Errorf("Readd should NOT be called on MISS, got %d calls", len(cache.readdCalled)) + } + if f.metrics.FlushSuccessTotal.Load() != 1 { + t.Errorf("FlushSuccessTotal = %d, want 1", f.metrics.FlushSuccessTotal.Load()) + } +} + +// --------------------------------------------------------------------------- +// 场景 3: UpsertFailReadds — writer 返普通 error → keys 被 Readd,FlushErrorTotal=1,DirtyReaddTotal=len +// --------------------------------------------------------------------------- + +func TestFlusher_UpsertFailReadds(t *testing.T) { + keys := []UserPlatformQuotaKey{ + {UserID: 1, Platform: "anthropic"}, + {UserID: 2, Platform: "openai"}, + } + cache := &mockQuotaDirtyCache{ + popSequence: [][]UserPlatformQuotaKey{keys}, + getEntries: []*UserPlatformQuotaCacheEntry{ + makeEntry(1.0, 2.0, 3.0), + makeEntry(4.0, 5.0, 6.0), + }, + } + writeErr := errors.New("db connection timeout") + writer := &mockQuotaSnapshotWriter{returnErr: writeErr} + f := newTestFlusher(cache, writer) + + f.flush() + + if f.metrics.FlushErrorTotal.Load() != 1 { + t.Errorf("FlushErrorTotal = %d, want 1", f.metrics.FlushErrorTotal.Load()) + } + if len(cache.readdCalled) == 0 { + t.Fatal("Readd should be called after write error") + } + totalReadd := 0 + for _, rk := range cache.readdCalled { + totalReadd += len(rk) + } + if totalReadd != len(keys) { + t.Errorf("DirtyReaddTotal (from Readd calls) = %d, want %d", totalReadd, len(keys)) + } + if f.metrics.DirtyReaddTotal.Load() != int64(len(keys)) { + t.Errorf("DirtyReaddTotal metric = %d, want %d", f.metrics.DirtyReaddTotal.Load(), len(keys)) + } + if f.metrics.FlushSuccessTotal.Load() != 0 { + t.Errorf("FlushSuccessTotal = %d, want 0", f.metrics.FlushSuccessTotal.Load()) + } +} + +// --------------------------------------------------------------------------- +// 场景 4: FKViolationDropsNoReadd — writer 返 ErrUserPlatformQuotaFKViolation → 不 Readd,FlushFKViolationTotal=1 +// --------------------------------------------------------------------------- + +func TestFlusher_FKViolationDropsNoReadd(t *testing.T) { + keys := []UserPlatformQuotaKey{ + {UserID: 999, Platform: "anthropic"}, + } + cache := &mockQuotaDirtyCache{ + popSequence: [][]UserPlatformQuotaKey{keys}, + getEntries: []*UserPlatformQuotaCacheEntry{ + makeEntry(1.0, 2.0, 3.0), + }, + } + writer := &mockQuotaSnapshotWriter{returnErr: ErrUserPlatformQuotaFKViolation} + f := newTestFlusher(cache, writer) + + f.flush() + + if f.metrics.FlushFKViolationTotal.Load() != 1 { + t.Errorf("FlushFKViolationTotal = %d, want 1", f.metrics.FlushFKViolationTotal.Load()) + } + if f.metrics.FlushErrorTotal.Load() != 1 { + t.Errorf("FlushErrorTotal = %d, want 1", f.metrics.FlushErrorTotal.Load()) + } + if len(cache.readdCalled) != 0 { + t.Errorf("Readd should NOT be called for FK violation (drop), got %d calls", len(cache.readdCalled)) + } + if f.metrics.DirtyReaddTotal.Load() != 0 { + t.Errorf("DirtyReaddTotal = %d, want 0 (FK violation drops)", f.metrics.DirtyReaddTotal.Load()) + } +} + +// --------------------------------------------------------------------------- +// 场景 5: NilSafe — var f *UserPlatformQuotaUsageFlusher; f.flush(); f.Stop() 不 panic +// --------------------------------------------------------------------------- + +func TestFlusher_NilSafe(t *testing.T) { + var f *UserPlatformQuotaUsageFlusher + // 下面两行不应 panic + f.flush() + f.Stop() +} + +// --------------------------------------------------------------------------- +// 场景 6: StopPreventsFlush — stopped=true 后 tick() 不调 flush(writer 没收到 snaps) +// --------------------------------------------------------------------------- + +func TestFlusher_StopPreventsFlush(t *testing.T) { + keys := []UserPlatformQuotaKey{ + {UserID: 1, Platform: "anthropic"}, + } + cache := &mockQuotaDirtyCache{ + popSequence: [][]UserPlatformQuotaKey{keys}, + getEntries: []*UserPlatformQuotaCacheEntry{ + makeEntry(1.0, 2.0, 3.0), + }, + } + writer := &mockQuotaSnapshotWriter{} + f := newTestFlusher(cache, writer) + + // 标记为已停止 + f.stopped.Store(true) + + // tick 应该直接返回,不触发 flush + f.tick() + + if len(writer.receivedSnaps) != 0 { + t.Errorf("expected 0 snaps after stop, got %d", len(writer.receivedSnaps)) + } + if cache.popCallIdx != 0 { + t.Errorf("Pop should not be called after stop, popCallIdx = %d", cache.popCallIdx) + } +} + +// --------------------------------------------------------------------------- +// 场景 B13-1: ZeroPercentCompany — 0% 公司脏集恒空,flusher 空跑无 DB 写 +// +// 模拟几乎没有用户配置 quota limit 的公司:脏集始终为空(popSequence 为空切片), +// Pop 每次返回空集。flush() 应早退,不写 DB、不计成功、不 Readd。 +// --------------------------------------------------------------------------- + +func TestScenario_ZeroPercentCompany(t *testing.T) { + cache := &mockQuotaDirtyCache{ + // popSequence 为空 → Pop 超出序列 → 始终返回 nil(空集) + popSequence: [][]UserPlatformQuotaKey{}, + } + writer := &mockQuotaSnapshotWriter{} + f := newTestFlusher(cache, writer) + + f.flush() + + if len(writer.receivedSnaps) != 0 { + t.Errorf("0%% company: expected 0 snaps, got %d", len(writer.receivedSnaps)) + } + if f.metrics.FlushBatchSizeTotal.Load() != 0 { + t.Errorf("0%% company: FlushBatchSizeTotal = %d, want 0", f.metrics.FlushBatchSizeTotal.Load()) + } + if f.metrics.FlushSuccessTotal.Load() != 0 { + t.Errorf("0%% company: FlushSuccessTotal = %d, want 0 (empty-set early return)", f.metrics.FlushSuccessTotal.Load()) + } + if f.metrics.FlushErrorTotal.Load() != 0 { + t.Errorf("0%% company: FlushErrorTotal = %d, want 0", f.metrics.FlushErrorTotal.Load()) + } + if len(cache.readdCalled) != 0 { + t.Errorf("0%% company: Readd should never be called, got %d calls", len(cache.readdCalled)) + } +} + +// --------------------------------------------------------------------------- +// P1: IntervalFallback — flush_interval_ms ≤0 时回退 2s;正常值保留 +// --------------------------------------------------------------------------- + +func TestNewUserPlatformQuotaUsageFlusher_IntervalFallback(t *testing.T) { + cases := []struct { + name string + inMs int + wantDu time.Duration + }{ + {"零值回退 2s", 0, 2 * time.Second}, + {"负数回退 2s", -100, 2 * time.Second}, + {"正常 2000ms 保留", 2000, 2 * time.Second}, + {"正常 500ms 保留", 500, 500 * time.Millisecond}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + cfg := &config.Config{} + cfg.Database.UserPlatformQuotaFlushIntervalMs = tc.inMs + f := NewUserPlatformQuotaUsageFlusher(cfg, nil, nil, nil) + if f.interval != tc.wantDu { + t.Fatalf("interval = %v, want %v", f.interval, tc.wantDu) + } + }) + } +} + +// --------------------------------------------------------------------------- +// P1: EnabledField — flusher_enabled 配置正确写入 f.enabled +// --------------------------------------------------------------------------- + +func TestNewUserPlatformQuotaUsageFlusher_EnabledField(t *testing.T) { + for _, enabled := range []bool{true, false} { + cfg := &config.Config{} + cfg.Database.UserPlatformQuotaFlusherEnabled = enabled + cfg.Database.UserPlatformQuotaFlushIntervalMs = 500 + f := NewUserPlatformQuotaUsageFlusher(cfg, nil, nil, nil) + if f.enabled != enabled { + t.Errorf("enabled = %v, want %v", f.enabled, enabled) + } + } +} + +// --------------------------------------------------------------------------- +// P2: ReaddFailCounts — BatchGet 失败 + Readd 失败 → DirtyLostTotal 增、DirtyReaddTotal 不变 +// BatchGet 失败 + Readd 成功 → DirtyReaddTotal 增、DirtyLostTotal 不变 +// --------------------------------------------------------------------------- + +func TestFlusher_ReaddFailCounts(t *testing.T) { + keys := []UserPlatformQuotaKey{ + {UserID: 10, Platform: "anthropic"}, + {UserID: 11, Platform: "openai"}, + } + + t.Run("Readd 失败计 DirtyLostTotal", func(t *testing.T) { + cache := &mockQuotaDirtyCache{ + popSequence: [][]UserPlatformQuotaKey{keys}, + getErr: errors.New("redis timeout"), // 触发 BatchGet 失败路径 + readdErr: errors.New("redis connection refused"), // Readd 也失败 + } + f := newTestFlusher(cache, &mockQuotaSnapshotWriter{}) + + f.flush() + + if f.metrics.DirtyLostTotal.Load() != int64(len(keys)) { + t.Errorf("DirtyLostTotal = %d, want %d", f.metrics.DirtyLostTotal.Load(), len(keys)) + } + if f.metrics.DirtyReaddTotal.Load() != 0 { + t.Errorf("DirtyReaddTotal = %d, want 0 (Readd 失败不应计入)", f.metrics.DirtyReaddTotal.Load()) + } + }) + + t.Run("Readd 成功计 DirtyReaddTotal", func(t *testing.T) { + cache := &mockQuotaDirtyCache{ + popSequence: [][]UserPlatformQuotaKey{keys}, + getErr: errors.New("redis timeout"), // 触发 BatchGet 失败路径 + readdErr: nil, // Readd 成功 + } + f := newTestFlusher(cache, &mockQuotaSnapshotWriter{}) + + f.flush() + + if f.metrics.DirtyReaddTotal.Load() != int64(len(keys)) { + t.Errorf("DirtyReaddTotal = %d, want %d", f.metrics.DirtyReaddTotal.Load(), len(keys)) + } + if f.metrics.DirtyLostTotal.Load() != 0 { + t.Errorf("DirtyLostTotal = %d, want 0 (Readd 成功不应计 lost)", f.metrics.DirtyLostTotal.Load()) + } + }) +} + +// --------------------------------------------------------------------------- +// ClampsBatchSize — NewUserPlatformQuotaUsageFlusher 构造时按 +// [defaultFlushBatchSize, maxFlushBatchSize] 区间 clamp batchSize +// --------------------------------------------------------------------------- + +func TestNewUserPlatformQuotaUsageFlusher_ClampsBatchSize(t *testing.T) { + cases := []struct { + name string + in int + want int + }{ + {"超上限被 clamp", 7000, maxFlushBatchSize}, + {"恰好上限保留", maxFlushBatchSize, maxFlushBatchSize}, + {"零回退默认", 0, defaultFlushBatchSize}, + {"负数回退默认", -5, defaultFlushBatchSize}, + {"正常值保留", 500, 500}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + cfg := &config.Config{} + cfg.Database.UserPlatformQuotaFlushBatchSize = tc.in + f := NewUserPlatformQuotaUsageFlusher(cfg, nil, nil, nil) + if f.batchSize != tc.want { + t.Fatalf("batchSize = %d, want %d", f.batchSize, tc.want) + } + }) + } +} + +// --------------------------------------------------------------------------- +// 场景 B13-2: NinetyPercentCompany — 90% 公司大量用户配 limit,一批 5 key 批量刷库 +// +// 模拟大量用户配置了 quota limit 的公司:脏集第一次 Pop 返回 5 个不同用户的 key, +// 之后返回空集(避免 flush 循环)。flush() 应构造 5 条 snapshot 写入 DB, +// 断言绝对值语义(snap 的 DailyUsageUSD 等于 entry 的值)、metrics 正确、不 Readd。 +// --------------------------------------------------------------------------- + +func TestScenario_NinetyPercentCompany(t *testing.T) { + keys := []UserPlatformQuotaKey{ + {UserID: 101, Platform: "anthropic"}, + {UserID: 102, Platform: "anthropic"}, + {UserID: 103, Platform: "openai"}, + {UserID: 104, Platform: "openai"}, + {UserID: 105, Platform: "anthropic"}, + } + entries := []*UserPlatformQuotaCacheEntry{ + makeEntry(1.1, 2.2, 3.3), + makeEntry(4.4, 5.5, 6.6), + makeEntry(7.7, 8.8, 9.9), + makeEntry(0.5, 1.0, 1.5), + makeEntry(10.0, 20.0, 30.0), + } + cache := &mockQuotaDirtyCache{ + // 第 1 次 Pop 返回 5 keys,之后返回空集(防止 flush 无限循环) + popSequence: [][]UserPlatformQuotaKey{keys}, + getEntries: entries, + } + writer := &mockQuotaSnapshotWriter{} + f := newTestFlusher(cache, writer) + + f.flush() + + // 应收到 5 条 snapshot + if len(writer.receivedSnaps) != 5 { + t.Fatalf("90%% company: expected 5 snaps, got %d", len(writer.receivedSnaps)) + } + + // 验证绝对值语义:第 1 条 snap 的各窗口 usage 应等于 entries[0] 的值 + snap0 := writer.receivedSnaps[0] + entry0 := entries[0] + if snap0.DailyUsageUSD != entry0.DailyUsageUSD { + t.Errorf("snap[0].DailyUsageUSD = %v, want %v", snap0.DailyUsageUSD, entry0.DailyUsageUSD) + } + if snap0.WeeklyUsageUSD != entry0.WeeklyUsageUSD { + t.Errorf("snap[0].WeeklyUsageUSD = %v, want %v", snap0.WeeklyUsageUSD, entry0.WeeklyUsageUSD) + } + if snap0.MonthlyUsageUSD != entry0.MonthlyUsageUSD { + t.Errorf("snap[0].MonthlyUsageUSD = %v, want %v", snap0.MonthlyUsageUSD, entry0.MonthlyUsageUSD) + } + + // FlushBatchSizeTotal 应为 5(本批 keys 数量) + if f.metrics.FlushBatchSizeTotal.Load() != 5 { + t.Errorf("90%% company: FlushBatchSizeTotal = %d, want 5", f.metrics.FlushBatchSizeTotal.Load()) + } + // FlushSuccessTotal 应为 1(1 个批次写成功) + if f.metrics.FlushSuccessTotal.Load() != 1 { + t.Errorf("90%% company: FlushSuccessTotal = %d, want 1", f.metrics.FlushSuccessTotal.Load()) + } + // 无错误、无 Readd + if f.metrics.FlushErrorTotal.Load() != 0 { + t.Errorf("90%% company: FlushErrorTotal = %d, want 0", f.metrics.FlushErrorTotal.Load()) + } + if f.metrics.DirtyReaddTotal.Load() != 0 { + t.Errorf("90%% company: DirtyReaddTotal = %d, want 0", f.metrics.DirtyReaddTotal.Load()) + } + if len(cache.readdCalled) != 0 { + t.Errorf("90%% company: Readd should not be called, got %d calls", len(cache.readdCalled)) + } +} diff --git a/backend/internal/service/user_platform_quota_port.go b/backend/internal/service/user_platform_quota_port.go index cb09542a..0f88eda4 100644 --- a/backend/internal/service/user_platform_quota_port.go +++ b/backend/internal/service/user_platform_quota_port.go @@ -11,6 +11,23 @@ import ( // handler 只需引用 service 包,无需直接依赖 repository 包。 var ErrUserPlatformQuotaNotFound = errors.New("user platform quota not found") +// ErrUserPlatformQuotaFKViolation service 层 sentinel:批量 snapshot UPSERT 时存在 +// user_id 不在 users 表的记录(外键违反)。adapter 负责将 repository 层同名 sentinel 包装为此错误。 +var ErrUserPlatformQuotaFKViolation = errors.New("user platform quota snapshot FK violation") + +// UserPlatformQuotaSnapshot 是 service 层 flusher 向 DB 写入快照时使用的传输结构。 +// 字段语义与 repository.UserPlatformQuotaSnapshot 完全对应,由 adapter 负责转换。 +type UserPlatformQuotaSnapshot struct { + UserID int64 + Platform string + DailyUsageUSD float64 + WeeklyUsageUSD float64 + MonthlyUsageUSD float64 + DailyWindowStart time.Time + WeeklyWindowStart time.Time + MonthlyWindowStart time.Time +} + // UserPlatformQuotaRecord service 层传输结构体(与 repository 层解耦)。 type UserPlatformQuotaRecord struct { UserID int64 @@ -47,4 +64,6 @@ type UserPlatformQuotaRepository interface { // ResetExpiredWindow 重置指定窗口("daily"|"weekly"|"monthly")的用量与起始时间。 // 未命中活跃记录时返回(service-side wrapper of repository.ErrUserPlatformQuotaNotFound)。 ResetExpiredWindow(ctx context.Context, userID int64, platform string, window string, newStart time.Time) error + // BatchSnapshotUsage 绝对值覆盖写入整批 usage 快照。FK 违反返回 ErrUserPlatformQuotaFKViolation。 + BatchSnapshotUsage(ctx context.Context, snapshots []UserPlatformQuotaSnapshot, now time.Time) error } diff --git a/backend/internal/service/user_service_test.go b/backend/internal/service/user_service_test.go index 19aec5d3..1a18e70a 100644 --- a/backend/internal/service/user_service_test.go +++ b/backend/internal/service/user_service_test.go @@ -327,10 +327,22 @@ func (m *mockBillingCache) DeleteUserPlatformQuotaCache(context.Context, int64, return nil } -func (m *mockBillingCache) IncrUserPlatformQuotaUsageCache(context.Context, int64, string, float64, time.Duration) error { +func (m *mockBillingCache) IncrUserPlatformQuotaUsageCache(context.Context, int64, string, float64, time.Duration, bool) error { return nil } +func (m *mockBillingCache) PopDirtyUserPlatformQuotaKeys(context.Context, int) ([]UserPlatformQuotaKey, error) { + return nil, nil +} + +func (m *mockBillingCache) ReaddDirtyUserPlatformQuotaKeys(context.Context, []UserPlatformQuotaKey) error { + return nil +} + +func (m *mockBillingCache) BatchGetUserPlatformQuotaCache(context.Context, []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) { + return nil, nil +} + // --- 测试 --- func TestUpdateBalance_Success(t *testing.T) { diff --git a/backend/internal/service/wire.go b/backend/internal/service/wire.go index d3e4ce51..19bd841d 100644 --- a/backend/internal/service/wire.go +++ b/backend/internal/service/wire.go @@ -571,8 +571,16 @@ var ProviderSet = wire.NewSet( ProvideChannelMonitorService, ProvideChannelMonitorRunner, NewChannelMonitorRequestTemplateService, + ProvideUserPlatformQuotaUsageFlusher, ) +// ProvideUserPlatformQuotaUsageFlusher 创建并启动 UserPlatformQuotaUsageFlusher。 +func ProvideUserPlatformQuotaUsageFlusher(cfg *config.Config, cache BillingCache, quotaRepo UserPlatformQuotaRepository, tw *TimingWheelService) *UserPlatformQuotaUsageFlusher { + svc := NewUserPlatformQuotaUsageFlusher(cfg, cache, quotaRepo, tw) + svc.Start() + return svc +} + // ProvidePaymentConfigService wraps NewPaymentConfigService to accept the named // payment.EncryptionKey type instead of raw []byte, avoiding Wire ambiguity. func ProvidePaymentConfigService(entClient *dbent.Client, settingRepo SettingRepository, key payment.EncryptionKey) *PaymentConfigService { From 5fd9a35093dffe8abc3b72ed6b53b5653b50ee98 Mon Sep 17 00:00:00 2001 From: DaydreamCoding <22166516+DaydreamCoding@users.noreply.github.com> Date: Fri, 29 May 2026 18:47:27 +0800 Subject: [PATCH 43/79] =?UTF-8?q?test(pricing):=20=E4=BF=AE=E5=A4=8D=20cod?= =?UTF-8?q?ex-auto-review=20=E5=AE=9A=E4=BB=B7=E6=B5=8B=E8=AF=95=E6=96=AD?= =?UTF-8?q?=E8=A8=80?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 68901cbff 批量同步定价数据后,codex-auto-review 的 input_cost_per_token 从 2.5e-6 更新为 5e-6,output 从 1.5e-5 更新为 3e-5,cache_read 从 2.5e-7 更新为 5e-7。测试断言需要同步更新以匹配当前定价数据。 Co-Authored-By: Claude Opus 4.8 (1M context) --- backend/internal/service/pricing_service_test.go | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/backend/internal/service/pricing_service_test.go b/backend/internal/service/pricing_service_test.go index cc8b120a..f4252f95 100644 --- a/backend/internal/service/pricing_service_test.go +++ b/backend/internal/service/pricing_service_test.go @@ -124,9 +124,9 @@ func TestDefaultPricingIncludesCodexAutoReview(t *testing.T) { got := svc.GetModelPricing("codex-auto-review") require.NotNil(t, got) - require.InDelta(t, 2.5e-6, got.InputCostPerToken, 1e-12) - require.InDelta(t, 1.5e-5, got.OutputCostPerToken, 1e-12) - require.InDelta(t, 2.5e-7, got.CacheReadInputTokenCost, 1e-12) + require.InDelta(t, 5e-6, got.InputCostPerToken, 1e-12) + require.InDelta(t, 3e-5, got.OutputCostPerToken, 1e-12) + require.InDelta(t, 5e-7, got.CacheReadInputTokenCost, 1e-12) } func TestGetModelPricing_Gpt54MiniUsesDedicatedStaticFallbackWhenRemoteMissing(t *testing.T) { From c256a5441a80a5aade3e54c255867c17db8bfa8c Mon Sep 17 00:00:00 2001 From: wucm667 Date: Fri, 29 May 2026 20:57:29 +0800 Subject: [PATCH 44/79] =?UTF-8?q?feat(admin):=20=E8=B4=A6=E5=8F=B7?= =?UTF-8?q?=E7=94=A8=E9=87=8F=E7=AA=97=E5=8F=A3=205h/7d=20=E5=A2=9E?= =?UTF-8?q?=E5=8A=A0=E8=AF=B4=E6=98=8E=20tooltip?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 在账号管理列表"用量窗口"列表头增加一个说明性 HelpTooltip, 解释 5h / 7d 是上游账号(如 OpenAI ChatGPT、Claude)官方的滚动 用量窗口限制,由上游设定、非 sub2api 配置、与映射模型无关,且 窗口滚动到期后自动重置、无法在 sub2api 端解除。 复用现有 HelpTooltip 组件(teleport 到 body,避免表格裁剪), 单个 ⓘ 图标置于列表头,避免每行重复。新增 i18n key admin.accounts.usageWindowsHint(zh/en 同步)。纯展示说明, 不改用量计算与后端逻辑。 Co-Authored-By: Claude Opus 4.8 --- frontend/src/i18n/locales/en.ts | 1 + frontend/src/i18n/locales/zh.ts | 1 + frontend/src/views/admin/AccountsView.vue | 7 + .../AccountsView.usageWindowsHint.spec.ts | 164 ++++++++++++++++++ 4 files changed, 173 insertions(+) create mode 100644 frontend/src/views/admin/__tests__/AccountsView.usageWindowsHint.spec.ts diff --git a/frontend/src/i18n/locales/en.ts b/frontend/src/i18n/locales/en.ts index 6735029c..2030ddf0 100644 --- a/frontend/src/i18n/locales/en.ts +++ b/frontend/src/i18n/locales/en.ts @@ -3105,6 +3105,7 @@ export default { expiresAt: 'Expires At', actions: 'Actions' }, + usageWindowsHint: '"5h / 7d" are the upstream account\'s official rolling usage windows (e.g. OpenAI ChatGPT, Claude). They are imposed by the upstream provider on the account itself — not configured by sub2api, and unrelated to the models you map. Usage resets automatically once each window rolls over, and the limit cannot be lifted from within sub2api.', allPrivacyModes: 'All Privacy States', privacyUnset: 'Unset', privacyTrainingOff: 'Training data sharing disabled', diff --git a/frontend/src/i18n/locales/zh.ts b/frontend/src/i18n/locales/zh.ts index abb8dff7..6106aa5a 100644 --- a/frontend/src/i18n/locales/zh.ts +++ b/frontend/src/i18n/locales/zh.ts @@ -3143,6 +3143,7 @@ export default { expiresAt: '过期时间', actions: '操作' }, + usageWindowsHint: '“5h / 7d”是上游账号(如 OpenAI ChatGPT、Claude)官方的滚动用量窗口限制,由上游对账号设定,并非 sub2api 配置,也与你映射的模型无关。窗口滚动到期后用量会自动重置,无法在 sub2api 端解除该限制。', allPrivacyModes: '全部Privacy状态', privacyUnset: '未设置', privacyTrainingOff: '已关闭训练数据共享', diff --git a/frontend/src/views/admin/AccountsView.vue b/frontend/src/views/admin/AccountsView.vue index c602225c..04b46a8d 100644 --- a/frontend/src/views/admin/AccountsView.vue +++ b/frontend/src/views/admin/AccountsView.vue @@ -273,6 +273,12 @@ + 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; +}