From 0464856c4aa5deb613dabff662f6ca6bf98fba13 Mon Sep 17 00:00:00 2001 From: Lyonle <214648221+lyon-le@users.noreply.github.com> Date: Fri, 10 Jul 2026 22:27:15 +0800 Subject: [PATCH 1/9] =?UTF-8?q?feat(frontend):=20Fast/Flex=20=E7=AD=96?= =?UTF-8?q?=E7=95=A5=E6=94=AF=E6=8C=81=E6=90=9C=E7=B4=A2=E9=80=89=E6=8B=A9?= =?UTF-8?q?=E7=94=A8=E6=88=B7?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 将手工 User ID 输入替换为带防抖的邮箱模糊搜索和多选标签,同时继续保存 user_ids。 回显已有用户邮箱并保留无法解析的历史 ID,补齐中英文文案路径及组件回归测试。 --- .../__tests__/openaiFastPolicyLocales.spec.ts | 30 +++ .../src/i18n/locales/en/admin/settings.ts | 12 +- .../src/i18n/locales/zh/admin/settings.ts | 12 +- frontend/src/views/admin/SettingsView.vue | 71 +----- .../settings/OpenAIFastPolicyUserSelector.vue | 229 ++++++++++++++++++ .../OpenAIFastPolicyUserSelector.spec.ts | 112 +++++++++ 6 files changed, 390 insertions(+), 76 deletions(-) create mode 100644 frontend/src/i18n/__tests__/openaiFastPolicyLocales.spec.ts create mode 100644 frontend/src/views/admin/settings/OpenAIFastPolicyUserSelector.vue create mode 100644 frontend/src/views/admin/settings/__tests__/OpenAIFastPolicyUserSelector.spec.ts diff --git a/frontend/src/i18n/__tests__/openaiFastPolicyLocales.spec.ts b/frontend/src/i18n/__tests__/openaiFastPolicyLocales.spec.ts new file mode 100644 index 0000000000..7fd550c1ad --- /dev/null +++ b/frontend/src/i18n/__tests__/openaiFastPolicyLocales.spec.ts @@ -0,0 +1,30 @@ +import { describe, expect, it } from 'vitest' + +import en from '../locales/en' +import zh from '../locales/zh' + +describe('OpenAI Fast/Flex policy locale keys', () => { + it('exposes user scope copy at the runtime zh path', () => { + expect(zh.admin.settings.openaiFastPolicy).toMatchObject({ + userIds: '指定用户', + userIdsHint: '输入任意邮箱关键词进行模糊搜索。留空表示对全部 Sub2API 用户生效;选中用户的 API Key 请求优先匹配用户规则。', + userSearchPlaceholder: '输入用户邮箱搜索', + userSearchEmpty: '未找到匹配用户', + userDeleted: '(已删除)', + userIdFallback: '用户 #{id}', + removeUser: '移除用户' + }) + }) + + it('exposes user scope copy at the runtime en path', () => { + expect(en.admin.settings.openaiFastPolicy).toMatchObject({ + userIds: 'Specific users', + userIdsHint: 'Type any part of a user email to search. Leave empty to apply to all Sub2API users. Selected users match requests from their API keys and take precedence over global rules.', + userSearchPlaceholder: 'Search by user email', + userSearchEmpty: 'No matching users found', + userDeleted: '(deleted)', + userIdFallback: 'User #{id}', + removeUser: 'Remove user' + }) + }) +}) diff --git a/frontend/src/i18n/locales/en/admin/settings.ts b/frontend/src/i18n/locales/en/admin/settings.ts index 37dc8aa1c8..e08be9f591 100644 --- a/frontend/src/i18n/locales/en/admin/settings.ts +++ b/frontend/src/i18n/locales/en/admin/settings.ts @@ -979,11 +979,6 @@ export default { scopeOAuth: 'OAuth only', scopeAPIKey: 'API Key only', scopeBedrock: 'Bedrock only', - userIds: 'Specific user IDs', - userIdsHint: 'Leave empty to apply to all Sub2API users. Specified users match requests from their API keys and take precedence over global rules.', - userIdPlaceholder: 'e.g., 1001', - addUserId: 'Add user ID', - removeUserId: 'Remove user ID', errorMessage: 'Error message', errorMessagePlaceholder: 'Custom error message when blocked', errorMessageHint: 'Leave empty for default message', @@ -1024,6 +1019,13 @@ export default { scopeOAuth: 'OAuth only', scopeAPIKey: 'API Key only', scopeBedrock: 'Bedrock only', + userIds: 'Specific users', + userIdsHint: 'Type any part of a user email to search. Leave empty to apply to all Sub2API users. Selected users match requests from their API keys and take precedence over global rules.', + userSearchPlaceholder: 'Search by user email', + userSearchEmpty: 'No matching users found', + userDeleted: '(deleted)', + userIdFallback: 'User #{id}', + removeUser: 'Remove user', errorMessage: 'Error message', errorMessagePlaceholder: 'Custom error message when blocked', errorMessageHint: 'Leave empty for the default message.', diff --git a/frontend/src/i18n/locales/zh/admin/settings.ts b/frontend/src/i18n/locales/zh/admin/settings.ts index 5c0d874b57..08c0dbcd53 100644 --- a/frontend/src/i18n/locales/zh/admin/settings.ts +++ b/frontend/src/i18n/locales/zh/admin/settings.ts @@ -974,11 +974,6 @@ export default { scopeOAuth: '仅 OAuth 账号', scopeAPIKey: '仅 API Key 账号', scopeBedrock: '仅 Bedrock 账号', - userIds: '指定用户 ID', - userIdsHint: '留空表示对全部 Sub2API 用户生效。指定后仅匹配这些用户的 API Key 请求,且优先于全局规则。', - userIdPlaceholder: '例如: 1001', - addUserId: '添加用户 ID', - removeUserId: '移除用户 ID', errorMessage: '错误消息', errorMessagePlaceholder: '拦截时返回的自定义错误消息', errorMessageHint: '留空则使用默认错误消息', @@ -1019,6 +1014,13 @@ export default { scopeOAuth: '仅 OAuth 账号', scopeAPIKey: '仅 API Key 账号', scopeBedrock: '仅 Bedrock 账号', + userIds: '指定用户', + userIdsHint: '输入任意邮箱关键词进行模糊搜索。留空表示对全部 Sub2API 用户生效;选中用户的 API Key 请求优先匹配用户规则。', + userSearchPlaceholder: '输入用户邮箱搜索', + userSearchEmpty: '未找到匹配用户', + userDeleted: '(已删除)', + userIdFallback: '用户 #{id}', + removeUser: '移除用户', errorMessage: '错误消息', errorMessagePlaceholder: '拦截时返回的自定义错误消息', errorMessageHint: '留空则使用默认错误消息。', diff --git a/frontend/src/views/admin/SettingsView.vue b/frontend/src/views/admin/SettingsView.vue index a87cea9060..26253c7eb8 100644 --- a/frontend/src/views/admin/SettingsView.vue +++ b/frontend/src/views/admin/SettingsView.vue @@ -1199,60 +1199,10 @@

{{ t("admin.settings.openaiFastPolicy.userIdsHint") }}

-
- - -
- + @@ -7431,6 +7381,7 @@ import ProxySelector from "@/components/common/ProxySelector.vue"; import ImageUpload from "@/components/common/ImageUpload.vue"; import BackupSettings from "@/views/admin/BackupView.vue"; import EmailTemplateEditor from "@/views/admin/settings/EmailTemplateEditor.vue"; +import OpenAIFastPolicyUserSelector from "@/views/admin/settings/OpenAIFastPolicyUserSelector.vue"; import { useClipboard } from "@/composables/useClipboard"; import { affiliatesAPI, type AffiliateAdminEntry, type SimpleUser as AffiliateSimpleUser } from "@/api/admin/affiliates"; import { extractApiErrorMessage, extractI18nErrorMessage } from "@/utils/apiError"; @@ -10226,18 +10177,6 @@ function removeOpenAIFastPolicyRule(index: number) { openaiFastPolicyForm.rules.splice(index, 1); } -function addOpenAIFastPolicyUserID(rule: OpenAIFastPolicyRule) { - if (!rule.user_ids) rule.user_ids = []; - rule.user_ids.push(0); -} - -function removeOpenAIFastPolicyUserID( - rule: OpenAIFastPolicyRule, - idx: number, -) { - rule.user_ids?.splice(idx, 1); -} - function addOpenAIFastPolicyModelPattern(rule: OpenAIFastPolicyRule) { if (!rule.model_whitelist) rule.model_whitelist = []; rule.model_whitelist.push(""); diff --git a/frontend/src/views/admin/settings/OpenAIFastPolicyUserSelector.vue b/frontend/src/views/admin/settings/OpenAIFastPolicyUserSelector.vue new file mode 100644 index 0000000000..0b25c876ad --- /dev/null +++ b/frontend/src/views/admin/settings/OpenAIFastPolicyUserSelector.vue @@ -0,0 +1,229 @@ + + + diff --git a/frontend/src/views/admin/settings/__tests__/OpenAIFastPolicyUserSelector.spec.ts b/frontend/src/views/admin/settings/__tests__/OpenAIFastPolicyUserSelector.spec.ts new file mode 100644 index 0000000000..c8e1f0e1cf --- /dev/null +++ b/frontend/src/views/admin/settings/__tests__/OpenAIFastPolicyUserSelector.spec.ts @@ -0,0 +1,112 @@ +import { flushPromises, mount } from '@vue/test-utils' +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' + +import OpenAIFastPolicyUserSelector from '../OpenAIFastPolicyUserSelector.vue' + +const messages: Record = { + 'admin.settings.openaiFastPolicy.userDeleted': '(deleted)', + 'admin.settings.openaiFastPolicy.userIdFallback': 'User #{id}', + 'admin.settings.openaiFastPolicy.removeUser': 'Remove user', + 'admin.settings.openaiFastPolicy.userSearchPlaceholder': 'Search users', + 'admin.settings.openaiFastPolicy.userSearchEmpty': 'No users found', + 'common.loading': 'Loading', +} + +vi.mock('vue-i18n', () => ({ + useI18n: () => ({ + t: (key: string, params?: Record) => { + const message = messages[key] ?? key + return params + ? Object.entries(params).reduce( + (value, [name, replacement]) => value.replace(`{${name}}`, String(replacement)), + message, + ) + : message + }, + }), +})) + +const mockSearchUsers = vi.fn() +const mockGetUserById = vi.fn() + +vi.mock('@/api/admin', () => ({ + adminAPI: { + usage: { + searchUsers: (...args: unknown[]) => mockSearchUsers(...args), + }, + users: { + getById: (...args: unknown[]) => mockGetUserById(...args), + }, + }, +})) + +describe('OpenAIFastPolicyUserSelector', () => { + beforeEach(() => { + vi.useFakeTimers() + mockSearchUsers.mockReset() + mockGetUserById.mockReset() + }) + + afterEach(() => { + vi.useRealTimers() + }) + + it('hydrates existing IDs to email labels without changing the saved IDs', async () => { + mockGetUserById.mockResolvedValue({ + id: 7, + email: 'existing@example.com', + deleted_at: null, + }) + + const wrapper = mount(OpenAIFastPolicyUserSelector, { + props: { modelValue: [7] }, + global: { stubs: { Icon: true } }, + }) + await flushPromises() + + expect(mockGetUserById).toHaveBeenCalledWith(7, true) + expect(wrapper.text()).toContain('existing@example.com') + expect(wrapper.text()).toContain('#7') + expect(wrapper.emitted('update:modelValue')).toBeUndefined() + }) + + it('searches after one character and adds the selected user ID', async () => { + mockSearchUsers.mockResolvedValue([ + { id: 9, email: 'alice@example.com', deleted: false }, + ]) + + const wrapper = mount(OpenAIFastPolicyUserSelector, { + props: { modelValue: [] }, + global: { stubs: { Icon: true } }, + }) + const input = wrapper.get('input') + await input.trigger('focus') + await input.setValue('a') + await input.trigger('input') + vi.advanceTimersByTime(300) + await flushPromises() + + expect(mockSearchUsers).toHaveBeenCalledWith('a') + const result = wrapper.findAll('button').find((button) => + button.text().includes('alice@example.com'), + ) + expect(result).toBeDefined() + await result!.trigger('click') + + expect(wrapper.emitted('update:modelValue')).toEqual([[[9]]]) + }) + + it('keeps an unresolved saved ID visible and removable', async () => { + mockGetUserById.mockRejectedValue(new Error('not found')) + + const wrapper = mount(OpenAIFastPolicyUserSelector, { + props: { modelValue: [42] }, + global: { stubs: { Icon: true } }, + }) + await flushPromises() + + expect(wrapper.text()).toContain('User #42') + await wrapper.get('button[aria-label="Remove user"]').trigger('click') + expect(wrapper.emitted('update:modelValue')).toEqual([[[]]]) + }) +}) From 4d4ba64bf7ba110241e0850bee2dd4180a6b3f49 Mon Sep 17 00:00:00 2001 From: wucm667 Date: Fri, 10 Jul 2026 23:32:34 +0800 Subject: [PATCH 2/9] =?UTF-8?q?fix(codex):=20=E5=89=A5=E7=A6=BB=E7=BB=AD?= =?UTF-8?q?=E9=93=BE=20message=20item=20=E7=9A=84=E9=9D=9E=E6=B3=95=20item?= =?UTF-8?q?=5F*=20id?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit OpenAI OAuth 转发续链请求时,type=message 的 item 的 id 被客户端以 item_* 形式回放,但上游要求以 msg 开头,返回 400 "Expected an ID that begins with 'msg'",sub2api 随后向客户端返回 502。 客户端自动重试会不断重放同一份坏上下文,导致连续失败。 filterCodexInputWithOptions 在 PreserveReferences=true 路径下为 type=message 增加与 #3785 (fd64d07e6) 平行的 id 前缀检查:非 msg 开头 即删除。合法的 msg* id 原样保留,不改写 item_* 为 msg_*,因为改写出的 id 未必对应真实上游对象。 原 TestFilterCodexInput_NonToolCallItemKeepsID 以 message + item_msg_001 断言"保留 id",该行为已被上游拒绝,改用 web_search_call 覆盖同一意图。 Fixes #3981 Co-Authored-By: Claude Opus 4.8 --- .../openai_codex_function_call_id_test.go | 15 +- .../openai_codex_message_item_id_test.go | 160 ++++++++++++++++++ .../service/openai_codex_transform.go | 9 + 3 files changed, 177 insertions(+), 7 deletions(-) create mode 100644 backend/internal/service/openai_codex_message_item_id_test.go diff --git a/backend/internal/service/openai_codex_function_call_id_test.go b/backend/internal/service/openai_codex_function_call_id_test.go index 2ac59e0520..edfcac1d71 100644 --- a/backend/internal/service/openai_codex_function_call_id_test.go +++ b/backend/internal/service/openai_codex_function_call_id_test.go @@ -114,14 +114,15 @@ func TestFilterCodexInput_OutputTypeKeepsItemID(t *testing.T) { require.Equal(t, "o1", out["id"], "output item id should be preserved") } -// TestFilterCodexInput_NonToolCallItemKeepsID ensures non-tool-call items -// (e.g. message) still keep their id when PreserveReferences is true. +// TestFilterCodexInput_NonToolCallItemKeepsID ensures items subject to neither +// the fc* (call-input) nor the msg* (message) prefix rule still keep their id +// when PreserveReferences is true. +// message is covered separately in openai_codex_message_item_id_test.go (#3981). func TestFilterCodexInput_NonToolCallItemKeepsID(t *testing.T) { input := []any{ map[string]any{ - "type": "message", - "id": "item_msg_001", - "role": "user", + "type": "web_search_call", + "id": "ws_001", }, } @@ -130,7 +131,7 @@ func TestFilterCodexInput_NonToolCallItemKeepsID(t *testing.T) { }) require.Len(t, filtered, 1) - msg, ok := filtered[0].(map[string]any) + item, ok := filtered[0].(map[string]any) require.True(t, ok) - require.Equal(t, "item_msg_001", msg["id"], "non-tool-call items keep their id in preserve mode") + require.Equal(t, "ws_001", item["id"], "unconstrained items keep their id in preserve mode") } diff --git a/backend/internal/service/openai_codex_message_item_id_test.go b/backend/internal/service/openai_codex_message_item_id_test.go new file mode 100644 index 0000000000..54c2a9ef22 --- /dev/null +++ b/backend/internal/service/openai_codex_message_item_id_test.go @@ -0,0 +1,160 @@ +//go:build unit + +package service + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +// TestFilterCodexInput_StripsMessageItemID_WhenPreservingReferences +// verifies that message items with a non-msg id (e.g. item_*) have their id +// stripped even when PreserveReferences is true. OpenAI upstream requires +// message ids to begin with "msg" and rejects item_* with 400: +// "Expected an ID that begins with 'msg'." (#3981) +func TestFilterCodexInput_StripsMessageItemID_WhenPreservingReferences(t *testing.T) { + input := []any{ + map[string]any{ + "type": "message", + "id": "item_3bc5a3fa8ccde25f1c0000d4", + "role": "user", + "content": []any{ + map[string]any{"type": "input_text", "text": "hello"}, + }, + }, + } + + filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{ + PreserveReferences: true, + }) + + require.Len(t, filtered, 1) + + msg, ok := filtered[0].(map[string]any) + require.True(t, ok) + require.Equal(t, "message", msg["type"]) + _, hasID := msg["id"] + require.False(t, hasID, "item_* id should be stripped from message") + require.Equal(t, "user", msg["role"], "role must be preserved") + require.NotNil(t, msg["content"], "content must be preserved") +} + +// TestFilterCodexInput_KeepsMsgID_WhenPreservingReferences +// verifies that message items with a valid msg* id are kept when +// PreserveReferences is true, so context references are not lost. +func TestFilterCodexInput_KeepsMsgID_WhenPreservingReferences(t *testing.T) { + input := []any{ + map[string]any{ + "type": "message", + "id": "msg_validID123", + "role": "assistant", + }, + } + + filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{ + PreserveReferences: true, + }) + + require.Len(t, filtered, 1) + msg, ok := filtered[0].(map[string]any) + require.True(t, ok) + require.Equal(t, "msg_validID123", msg["id"], "valid msg* id must be preserved") +} + +// TestFilterCodexInput_StripsMessageIDWhenNotPreservingReferences ensures the +// non-continuation path still drops every message id regardless of prefix. +func TestFilterCodexInput_StripsMessageIDWhenNotPreservingReferences(t *testing.T) { + for _, id := range []string{"item_abc", "msg_validID123"} { + input := []any{ + map[string]any{ + "type": "message", + "id": id, + "role": "user", + }, + } + + filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{ + PreserveReferences: false, + }) + + require.Len(t, filtered, 1) + msg, ok := filtered[0].(map[string]any) + require.True(t, ok) + _, hasID := msg["id"] + require.False(t, hasID, "id %q should be stripped when not preserving references", id) + } +} + +// TestFilterCodexInput_MessageIDStripDoesNotMutateInput ensures the original +// input map is not modified in place when the id is stripped. +func TestFilterCodexInput_MessageIDStripDoesNotMutateInput(t *testing.T) { + original := map[string]any{ + "type": "message", + "id": "item_abc", + "role": "user", + } + + filtered := filterCodexInputWithOptions([]any{original}, codexInputFilterOptions{ + PreserveReferences: true, + }) + + require.Len(t, filtered, 1) + require.Equal(t, "item_abc", original["id"], "original input must not be mutated") +} + +// TestFilterCodexInput_MessageStripKeepsFunctionCallBehavior guards against a +// regression of #3785: message and function_call id rules are independent. +func TestFilterCodexInput_MessageStripKeepsFunctionCallBehavior(t *testing.T) { + input := []any{ + map[string]any{ + "type": "message", + "id": "item_msg_001", + "role": "user", + }, + map[string]any{ + "type": "function_call", + "id": "fc_validID123", + "call_id": "fc_validID123", + "name": "bash", + }, + map[string]any{ + "type": "function_call", + "id": "item_A9v0SNfS3VaLrfX0j3y4xhyK", + "call_id": "fc_abc123", + "name": "bash", + }, + map[string]any{ + "type": "function_call_output", + "id": "o1", + "call_id": "fc_abc123", + "output": "done", + }, + } + + filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{ + PreserveReferences: true, + }) + + require.Len(t, filtered, 4) + + msg, ok := filtered[0].(map[string]any) + require.True(t, ok) + _, hasID := msg["id"] + require.False(t, hasID, "message item_* id should be stripped") + + fcValid, ok := filtered[1].(map[string]any) + require.True(t, ok) + require.Equal(t, "fc_validID123", fcValid["id"], "valid fc* id must be preserved") + + fcBad, ok := filtered[2].(map[string]any) + require.True(t, ok) + _, hasID = fcBad["id"] + require.False(t, hasID, "function_call item_* id should still be stripped") + require.Equal(t, "fc_abc123", fcBad["call_id"], "call_id pairing must survive") + + out, ok := filtered[3].(map[string]any) + require.True(t, ok) + require.Equal(t, "o1", out["id"], "output item id should be preserved") + require.Equal(t, "fc_abc123", out["call_id"], "call_id pairing must survive") +} diff --git a/backend/internal/service/openai_codex_transform.go b/backend/internal/service/openai_codex_transform.go index 99355628f2..3869e97c99 100644 --- a/backend/internal/service/openai_codex_transform.go +++ b/backend/internal/service/openai_codex_transform.go @@ -1405,6 +1405,15 @@ func filterCodexInputWithOptions(input []any, opts codexInputFilterOptions) []an ensureCopy() delete(newItem, "id") } + } else if typ == "message" { + // 同理,message 类 item 的 id 必须以 "msg" 开头(上游校验 + // "Expected an ID that begins with 'msg'")。item_* 形式的 id + // 来自客户端回放,需要删除。 + // 注意:不改写成 msg_*,改写出的 id 未必对应真实的上游对象。 + if id, ok := m["id"].(string); ok && id != "" && !strings.HasPrefix(id, "msg") { + ensureCopy() + delete(newItem, "id") + } } filtered = append(filtered, newItem) From 6e2bb312812b214751e7602cf48271ab9efefbcb Mon Sep 17 00:00:00 2001 From: wucm667 Date: Sat, 11 Jul 2026 08:52:46 +0800 Subject: [PATCH 3/9] fix(service): guard compact keepalive writer delegates --- .../service/openai_compact_sse_keepalive.go | 56 ++++++++++ .../openai_compact_sse_keepalive_test.go | 105 ++++++++++++++++++ 2 files changed, 161 insertions(+) diff --git a/backend/internal/service/openai_compact_sse_keepalive.go b/backend/internal/service/openai_compact_sse_keepalive.go index 70ef3fc01a..4d4543f477 100644 --- a/backend/internal/service/openai_compact_sse_keepalive.go +++ b/backend/internal/service/openai_compact_sse_keepalive.go @@ -1,6 +1,9 @@ package service import ( + "bufio" + "errors" + "net" "net/http" "sync" "time" @@ -181,52 +184,105 @@ type openAICompactKeepaliveWriter struct { // suspend 停拍心跳;幂等。任何响应构造(含 Header 访问——写响应必先操作 // 响应头)都视为请求侧接管 ResponseWriter。 func (w *openAICompactKeepaliveWriter) suspend() { + if w.k == nil { + return + } w.k.Stop() } func (w *openAICompactKeepaliveWriter) Header() http.Header { w.suspend() + if w.ResponseWriter == nil { + return http.Header{} + } return w.ResponseWriter.Header() } func (w *openAICompactKeepaliveWriter) Write(data []byte) (int, error) { w.suspend() + if w.ResponseWriter == nil { + return 0, nil + } return w.ResponseWriter.Write(data) } func (w *openAICompactKeepaliveWriter) WriteString(s string) (int, error) { w.suspend() + if w.ResponseWriter == nil { + return 0, nil + } return w.ResponseWriter.WriteString(s) } func (w *openAICompactKeepaliveWriter) WriteHeader(code int) { w.suspend() + if w.ResponseWriter == nil { + return + } w.ResponseWriter.WriteHeader(code) } func (w *openAICompactKeepaliveWriter) WriteHeaderNow() { w.suspend() + if w.ResponseWriter == nil { + return + } w.ResponseWriter.WriteHeaderNow() } func (w *openAICompactKeepaliveWriter) Flush() { w.suspend() + if w.ResponseWriter == nil { + return + } w.ResponseWriter.Flush() } +func (w *openAICompactKeepaliveWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) { + if w.ResponseWriter == nil { + return nil, nil, errors.New("response writer released") + } + return w.ResponseWriter.Hijack() +} + +func (w *openAICompactKeepaliveWriter) CloseNotify() <-chan bool { + if w.ResponseWriter == nil { + ch := make(chan bool) + close(ch) + return ch + } + return w.ResponseWriter.CloseNotify() +} + +func (w *openAICompactKeepaliveWriter) Pusher() http.Pusher { + if w.ResponseWriter == nil { + return nil + } + return w.ResponseWriter.Pusher() +} + func (w *openAICompactKeepaliveWriter) Status() int { + if w.k == nil || w.ResponseWriter == nil { + return 0 + } w.k.mu.Lock() defer w.k.mu.Unlock() return w.ResponseWriter.Status() } func (w *openAICompactKeepaliveWriter) Size() int { + if w.k == nil || w.ResponseWriter == nil { + return 0 + } w.k.mu.Lock() defer w.k.mu.Unlock() return w.ResponseWriter.Size() } func (w *openAICompactKeepaliveWriter) Written() bool { + if w.k == nil || w.ResponseWriter == nil { + return false + } w.k.mu.Lock() defer w.k.mu.Unlock() return w.ResponseWriter.Written() diff --git a/backend/internal/service/openai_compact_sse_keepalive_test.go b/backend/internal/service/openai_compact_sse_keepalive_test.go index 3b217a0718..1efed7e9e4 100644 --- a/backend/internal/service/openai_compact_sse_keepalive_test.go +++ b/backend/internal/service/openai_compact_sse_keepalive_test.go @@ -6,6 +6,7 @@ import ( "testing" "time" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/tidwall/gjson" ) @@ -141,6 +142,110 @@ func TestOpenAICompactKeepaliveWriter_RequestSideWriteSuspendsBeats(t *testing.T require.Contains(t, rec.Body.String(), `{"error":"local reject"}`) } +func TestOpenAICompactKeepaliveWriter_NilInnerWriter_NoPanic(t *testing.T) { + w := &openAICompactKeepaliveWriter{ + k: &openAICompactSSEKeepalive{stop: make(chan struct{})}, + } + w.ResponseWriter = nil + + assert.NotPanics(t, func() { + assert.Equal(t, 0, w.Status()) + }) + assert.NotPanics(t, func() { + assert.Equal(t, 0, w.Size()) + }) + assert.NotPanics(t, func() { + assert.False(t, w.Written()) + }) + assert.NotPanics(t, func() { + assert.NotNil(t, w.Header()) + }) + assert.NotPanics(t, func() { + n, err := w.Write([]byte("test")) + assert.Equal(t, 0, n) + assert.NoError(t, err) + }) + assert.NotPanics(t, func() { + n, err := w.WriteString("test") + assert.Equal(t, 0, n) + assert.NoError(t, err) + }) + assert.NotPanics(t, func() { + w.WriteHeader(http.StatusOK) + }) + assert.NotPanics(t, func() { + w.WriteHeaderNow() + }) + assert.NotPanics(t, func() { + w.Flush() + }) + assert.NotPanics(t, func() { + conn, rw, err := w.Hijack() + assert.Nil(t, conn) + assert.Nil(t, rw) + assert.Error(t, err) + }) + assert.NotPanics(t, func() { + ch := w.CloseNotify() + assert.NotNil(t, ch) + }) + assert.NotPanics(t, func() { + assert.Nil(t, w.Pusher()) + }) +} + +func TestOpenAICompactKeepaliveWriter_NilKeepalive_NoPanic(t *testing.T) { + c, rec := newCompactBridgeTestContext(t, true) + w := &openAICompactKeepaliveWriter{ResponseWriter: c.Writer} + + assert.NotPanics(t, func() { + assert.Equal(t, 0, w.Status()) + }) + assert.NotPanics(t, func() { + assert.Equal(t, 0, w.Size()) + }) + assert.NotPanics(t, func() { + assert.False(t, w.Written()) + }) + assert.NotPanics(t, func() { + w.Header().Set("X-Test", "ok") + }) + assert.NotPanics(t, func() { + w.WriteHeader(http.StatusAccepted) + }) + assert.NotPanics(t, func() { + n, err := w.WriteString("ok") + assert.Equal(t, 2, n) + assert.NoError(t, err) + }) + assert.NotPanics(t, func() { + w.Flush() + }) + require.Equal(t, "ok", rec.Header().Get("X-Test")) + require.Equal(t, "ok", rec.Body.String()) +} + +func TestOpenAICompactKeepaliveWriter_DelegatesWhenReady(t *testing.T) { + c, rec := newCompactBridgeTestContext(t, true) + stop := StartOpenAICompactSSEKeepalive(c, time.Hour) + defer stop() + + w, ok := c.Writer.(*openAICompactKeepaliveWriter) + require.True(t, ok) + + w.Header().Set("X-Test", "ok") + w.WriteHeader(http.StatusAccepted) + n, err := w.WriteString("ready") + require.NoError(t, err) + require.Equal(t, len("ready"), n) + + require.Equal(t, http.StatusAccepted, w.Status()) + require.Equal(t, len("ready"), w.Size()) + require.True(t, w.Written()) + require.Equal(t, "ok", rec.Header().Get("X-Test")) + require.Equal(t, "ready", rec.Body.String()) +} + // fast policy block 在心跳提交后必须降级为 response.failed 终止事件。 func TestWriteOpenAIFastPolicyBlockedResponse_AfterKeepaliveCommit(t *testing.T) { c, rec := newCompactBridgeTestContext(t, true) From 84bb7d070974dc9ee12dcca3d263a87cb4a58430 Mon Sep 17 00:00:00 2001 From: Tian Lee <498756723@qq.com> Date: Fri, 10 Jul 2026 15:32:04 +0800 Subject: [PATCH 4/9] =?UTF-8?q?fix:=20=E4=BF=9D=E7=95=99=20remote=5Fcompac?= =?UTF-8?q?tion=5Fv2=20=E5=8E=9F=E7=94=9F=20Responses=20=E9=93=BE=E8=B7=AF?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- ...openai_gateway_compact_body_signal_test.go | 142 +++++++++------ .../handler/openai_gateway_handler.go | 30 +++- .../service/openai_compact_body_signal.go | 16 +- .../service/openai_gateway_service.go | 2 + .../internal/service/openai_gpt56_max_test.go | 81 ++++++++- .../service/openai_oauth_passthrough_test.go | 4 + .../service/openai_ws_forwarder_payload.go | 5 + .../openai_ws_forwarder_success_test.go | 2 + backend/internal/service/openai_ws_pool.go | 141 ++++++++++++++- .../internal/service/openai_ws_pool_test.go | 165 ++++++++++++++++++ 10 files changed, 505 insertions(+), 83 deletions(-) diff --git a/backend/internal/handler/openai_gateway_compact_body_signal_test.go b/backend/internal/handler/openai_gateway_compact_body_signal_test.go index a4bfb90466..a44d47c856 100644 --- a/backend/internal/handler/openai_gateway_compact_body_signal_test.go +++ b/backend/internal/handler/openai_gateway_compact_body_signal_test.go @@ -23,46 +23,61 @@ func newCompactBodySignalTestContext(t *testing.T, path string, body []byte) *gi return c } -// body-signal 提升后必须与 path-based compact 走同一条链路: -// path 改写、requireCompact 判定、stream/store/prompt_cache_key 归一化删除。 -// 回归防护:若 stream 字段存活,Forward 会用流式 handler 解析 compact 的 -// JSON 响应,导致 "stream ended before a terminal event" 的换号 failover 风暴。 -func TestNormalizeOpenAIResponsesCompactRequest_BodySignalPromoted(t *testing.T) { +func TestNormalizeOpenAIResponsesCompactRequest_RemoteV2StaysOnResponses(t *testing.T) { h := &OpenAIGatewayHandler{} body := []byte(`{ - "model":"gpt-5.5", + "model":"gpt-5.6-sol", "stream":true, "store":true, "prompt_cache_key":"pck-signal-1", + "reasoning":{"effort":"max","context":"all_turns"}, "input":[ {"type":"message","role":"user","content":"hello"}, {"type":"compaction_trigger"} ] }`) c := newCompactBodySignalTestContext(t, "/v1/responses", body) + c.Request.Header.Set("x-codex-beta-features", "responses_websockets_v2, remote_compaction_v2, another_feature") normalized, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body) require.True(t, ok) - require.Equal(t, "/v1/responses/compact", c.Request.URL.Path) - require.True(t, isOpenAIRemoteCompactPath(c)) - - require.False(t, gjson.GetBytes(normalized, "stream").Exists()) - require.False(t, gjson.GetBytes(normalized, "store").Exists()) - require.False(t, gjson.GetBytes(normalized, "prompt_cache_key").Exists()) - require.Equal(t, "gpt-5.5", gjson.GetBytes(normalized, "model").String()) - require.True(t, gjson.GetBytes(normalized, "input").IsArray()) + require.Equal(t, "/v1/responses", c.Request.URL.Path) + require.False(t, isOpenAIRemoteCompactPath(c)) + require.Equal(t, body, normalized) + require.True(t, gjson.GetBytes(normalized, "stream").Bool()) + require.True(t, gjson.GetBytes(normalized, "store").Bool()) + require.Equal(t, "pck-signal-1", gjson.GetBytes(normalized, "prompt_cache_key").String()) + require.Equal(t, "max", gjson.GetBytes(normalized, "reasoning.effort").String()) + require.Equal(t, "all_turns", gjson.GetBytes(normalized, "reasoning.context").String()) reqStream, streamOK := parseOpenAICompatibleStream(normalized) require.True(t, streamOK) - require.False(t, reqStream) + require.True(t, reqStream) - seed, exists := c.Get(service.OpenAICompactSessionSeedKeyForTest()) - require.True(t, exists) - require.Equal(t, "pck-signal-1", seed) + _, seedExists := c.Get(service.OpenAICompactSessionSeedKeyForTest()) + require.False(t, seedExists) + _, streamMarkerExists := c.Get(service.OpenAICompactClientStreamKeyForTest()) + require.False(t, streamMarkerExists) } -func TestNormalizeOpenAIResponsesCompactRequest_BodySignalTrailingSlash(t *testing.T) { +func TestNormalizeOpenAIResponsesCompactRequest_RemoteV2PathAliasesStayOnResponses(t *testing.T) { + h := &OpenAIGatewayHandler{} + body := []byte(`{"model":"gpt-5.6-sol","stream":true,"input":[{"type":"compaction_trigger"}]}`) + for _, path := range []string{"/v1/responses/", "/backend-api/codex/responses"} { + t.Run(path, func(t *testing.T) { + c := newCompactBodySignalTestContext(t, path, body) + c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2") + + normalized, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body) + require.True(t, ok) + require.Equal(t, path, c.Request.URL.Path) + require.Equal(t, body, normalized) + }) + } +} + +func TestNormalizeOpenAIResponsesCompactRequest_BodySignalTrailingSlashPromoted(t *testing.T) { h := &OpenAIGatewayHandler{} body := []byte(`{"model":"gpt-5.5","input":[{"type":"compaction_trigger"}]}`) c := newCompactBodySignalTestContext(t, "/v1/responses/", body) @@ -82,6 +97,64 @@ func TestNormalizeOpenAIResponsesCompactRequest_CodexDirectAliasPromoted(t *test require.Equal(t, "/backend-api/codex/responses/compact", c.Request.URL.Path) } +func TestNormalizeOpenAIResponsesCompactRequest_NonRemoteV2BodySignalPromoted(t *testing.T) { + h := &OpenAIGatewayHandler{} + tests := []struct { + name string + body []byte + betaHeader string + wantMarked bool + }{ + { + name: "no_header", + body: []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"compaction_trigger"}]}`), + wantMarked: true, + }, + { + name: "unrelated_header", + body: []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"compaction_trigger"}]}`), + betaHeader: "responses_websockets_v2", + wantMarked: true, + }, + { + name: "wrong_case_header", + body: []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"compaction_trigger"}]}`), + betaHeader: "REMOTE_COMPACTION_V2", + wantMarked: true, + }, + { + name: "stream_false", + body: []byte(`{"model":"gpt-5.5","stream":false,"input":[{"type":"compaction_trigger"}]}`), + betaHeader: "remote_compaction_v2", + }, + { + name: "stream_absent", + body: []byte(`{"model":"gpt-5.5","input":[{"type":"compaction_trigger"}]}`), + betaHeader: "remote_compaction_v2", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + c := newCompactBodySignalTestContext(t, "/v1/responses", tt.body) + if tt.betaHeader != "" { + c.Request.Header.Set("x-codex-beta-features", tt.betaHeader) + } + + normalized, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), tt.body) + require.True(t, ok) + require.Equal(t, "/v1/responses/compact", c.Request.URL.Path) + require.False(t, gjson.GetBytes(normalized, "stream").Exists()) + + marked, exists := c.Get(service.OpenAICompactClientStreamKeyForTest()) + require.Equal(t, tt.wantMarked, exists) + if tt.wantMarked { + require.Equal(t, true, marked) + } + }) + } +} + func TestNormalizeOpenAIResponsesCompactRequest_NoTriggerUntouched(t *testing.T) { h := &OpenAIGatewayHandler{} body := []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"message","role":"user","content":"hello"}]}`) @@ -99,6 +172,7 @@ func TestNormalizeOpenAIResponsesCompactRequest_PathBasedNoDoubleSuffix(t *testi h := &OpenAIGatewayHandler{} body := []byte(`{"model":"gpt-5.5","stream":true,"store":true,"input":[{"type":"message","role":"user","content":"hello"}]}`) c := newCompactBodySignalTestContext(t, "/v1/responses/compact", body) + c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2") normalized, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body) require.True(t, ok) @@ -118,36 +192,6 @@ func TestNormalizeOpenAIResponsesCompactRequest_SubpathNotPromoted(t *testing.T) require.Equal(t, body, normalized) } -// 回归 #3875:body-signal 原始请求 stream:true 时必须标记 client-stream, -// 供响应写回阶段把上游 unary JSON 合成回 Codex remote compact v2 所需的 SSE。 -func TestNormalizeOpenAIResponsesCompactRequest_BodySignalStreamTrueMarksClientStream(t *testing.T) { - h := &OpenAIGatewayHandler{} - body := []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"compaction_trigger"}]}`) - c := newCompactBodySignalTestContext(t, "/v1/responses", body) - - _, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body) - require.True(t, ok) - - marked, exists := c.Get(service.OpenAICompactClientStreamKeyForTest()) - require.True(t, exists) - require.Equal(t, true, marked) -} - -func TestNormalizeOpenAIResponsesCompactRequest_BodySignalStreamFalseNotMarked(t *testing.T) { - h := &OpenAIGatewayHandler{} - for name, body := range map[string][]byte{ - "stream_false": []byte(`{"model":"gpt-5.5","stream":false,"input":[{"type":"compaction_trigger"}]}`), - "stream_absent": []byte(`{"model":"gpt-5.5","input":[{"type":"compaction_trigger"}]}`), - } { - c := newCompactBodySignalTestContext(t, "/v1/responses", body) - _, ok := h.normalizeOpenAIResponsesCompactRequest(c, zap.NewNop(), body) - require.True(t, ok, name) - require.Equal(t, "/v1/responses/compact", c.Request.URL.Path, name) - _, exists := c.Get(service.OpenAICompactClientStreamKeyForTest()) - require.False(t, exists, "case %s 不应标记 client-stream", name) - } -} - // path-based compact(Codex v1 unary 协议)即使 body 带 stream:true 也不标记, // 保持 JSON 写回行为不变。 func TestNormalizeOpenAIResponsesCompactRequest_PathBasedStreamTrueNotMarked(t *testing.T) { diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 83e644d857..a4d7ee7b18 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -580,21 +580,33 @@ func isBareOpenAIResponsesPath(c *gin.Context) bool { return strings.HasSuffix(normalizedPath, "/responses") } -// normalizeOpenAIResponsesCompactRequest 统一处理两种入站 compact 形态: -// path-based(POST /v1/responses/compact)与 Codex remote compact v2 的 -// body-signal(普通 POST /v1/responses 的 input 中携带 type=compaction_trigger, -// 见 #3777)。body-signal 命中时在 stream 解析、compact body 归一化与 -// requireCompact 调度判定之前改写 URL path,使后续全部链路(含 passthrough -// 分支与上游 URL 构建)与 path-based 完全一致。 +func isOpenAIRemoteCompactionV2Request(c *gin.Context, body []byte) bool { + stream, valid := parseOpenAICompatibleStream(body) + if !valid || !stream || c == nil || c.Request == nil { + return false + } + for _, header := range c.Request.Header.Values("x-codex-beta-features") { + for _, feature := range strings.Split(header, ",") { + if strings.TrimSpace(feature) == "remote_compaction_v2" { + return true + } + } + } + return false +} + +// normalizeOpenAIResponsesCompactRequest keeps Codex remote compaction v2 on +// its native streaming /responses wire and preserves the legacy body-signal +// promotion for clients that do not explicitly advertise that protocol. // 返回归一化后的 body;ok=false 表示错误响应已写出,调用方应直接 return。 func (h *OpenAIGatewayHandler) normalizeOpenAIResponsesCompactRequest(c *gin.Context, reqLog *zap.Logger, body []byte) ([]byte, bool) { isCompactRequest := service.IsOpenAIResponsesCompactPathForTest(c) if !isCompactRequest && isBareOpenAIResponsesPath(c) && service.HasCompactionTriggerInInput(body) { + if isOpenAIRemoteCompactionV2Request(c, body) { + return body, true + } c.Request.URL.Path = strings.TrimRight(c.Request.URL.Path, "/") + "/compact" isCompactRequest = true - // Codex remote compact v2 的原始请求是流式 /responses:白名单归一化会删除 - // stream 并让上游走 unary JSON,但客户端仍按 SSE 消费响应。记录原始 - // stream 意图,响应写回阶段据此把 JSON 合成回 SSE(#3875)。 clientStream := gjson.GetBytes(body, "stream").Bool() if clientStream { service.MarkOpenAICompactClientStream(c) diff --git a/backend/internal/service/openai_compact_body_signal.go b/backend/internal/service/openai_compact_body_signal.go index fce62046c1..ce561b0c5a 100644 --- a/backend/internal/service/openai_compact_body_signal.go +++ b/backend/internal/service/openai_compact_body_signal.go @@ -2,18 +2,10 @@ package service import "github.com/tidwall/gjson" -// HasCompactionTriggerInInput detects the Codex remote compact v2 body signal: -// an input item with type "compaction_trigger". When the client sends this -// inside a normal POST /v1/responses (instead of POST /v1/responses/compact), -// the request must still be treated as a compact request — otherwise the -// upstream path, model mapping, and body normalization are all wrong, causing -// Codex to receive a non-compact response and fail with: -// -// "remote compaction v2 expected exactly one compaction output item, got 0" -// -// The gateway handler promotes such requests by rewriting the URL path to the -// compact form before stream parsing, compact body normalization, and -// compact-capable account scheduling, so both inbound forms share one code path. +// HasCompactionTriggerInInput detects an input item with +// type="compaction_trigger". The handler combines this body signal with the +// request path, stream flag, and Codex beta feature header to distinguish the +// native remote compaction v2 wire from the legacy /responses/compact bridge. func HasCompactionTriggerInInput(body []byte) bool { if len(body) == 0 { return false diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index 0d03d83dcf..c3b1b996f1 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -66,6 +66,7 @@ var openaiAllowedHeaders = map[string]bool{ "user-agent": true, "originator": true, "session_id": true, + "x-codex-beta-features": true, "x-codex-turn-state": true, "x-codex-turn-metadata": true, } @@ -81,6 +82,7 @@ var openaiPassthroughAllowedHeaders = map[string]bool{ "user-agent": true, "originator": true, "session_id": true, + "x-codex-beta-features": true, "x-codex-turn-state": true, "x-codex-turn-metadata": true, } diff --git a/backend/internal/service/openai_gpt56_max_test.go b/backend/internal/service/openai_gpt56_max_test.go index 272eb16ff0..cbca2ff3ee 100644 --- a/backend/internal/service/openai_gpt56_max_test.go +++ b/backend/internal/service/openai_gpt56_max_test.go @@ -223,13 +223,17 @@ func TestOpenAIGatewayServiceForwardOAuthCompactDowngradesMaxEffort(t *testing.T require.Equal(t, "xhigh", *result.ReasoningEffort) } -func TestOpenAIGatewayServiceForwardOAuthResponsesPreservesMaxEffort(t *testing.T) { +func TestOpenAIGatewayServiceForwardOAuthRemoteCompactV2PreservesResponsesWire(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}}`)), + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader( + "data: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"compaction\",\"encrypted_content\":\"summary\"}}\n\n" + + "data: {\"type\":\"response.completed\",\"response\":{\"usage\":{\"input_tokens\":1,\"output_tokens\":2}}}\n\n" + + "data: [DONE]\n\n", + )), }, } cfg := &config.Config{} @@ -244,6 +248,9 @@ func TestOpenAIGatewayServiceForwardOAuthResponsesPreservesMaxEffort(t *testing. Credentials: map[string]any{ "access_token": "oauth-token", "chatgpt_account_id": "chatgpt-acc", + "compact_model_mapping": map[string]any{ + "gpt-5.6-sol": "gpt-5.6-sol-openai-compact", + }, }, Status: StatusActive, Schedulable: true, @@ -251,16 +258,82 @@ func TestOpenAIGatewayServiceForwardOAuthResponsesPreservesMaxEffort(t *testing. rec := httptest.NewRecorder() c, _ := gin.CreateTestContext(rec) c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2") SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) - body := []byte(`{"model":"gpt-5.6-sol","instructions":"response-test","input":"hello","reasoning":{"effort":"max"}}`) + body := []byte(`{"model":"gpt-5.6-sol","stream":true,"instructions":"response-test","input":[{"type":"message","role":"user","content":"hello"},{"type":"compaction_trigger"}],"reasoning":{"effort":"max","context":"all_turns"}}`) result, err := svc.Forward(context.Background(), c, account, body) require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, upstream.lastReq) require.Equal(t, chatgptCodexURL, upstream.lastReq.URL.String()) + require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(upstream.lastBody, "model").String()) + require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool()) + require.Equal(t, "compaction_trigger", gjson.GetBytes(upstream.lastBody, "input.#(type==\"compaction_trigger\").type").String()) require.Equal(t, "max", gjson.GetBytes(upstream.lastBody, "reasoning.effort").String()) + require.Equal(t, "all_turns", gjson.GetBytes(upstream.lastBody, "reasoning.context").String()) + require.Equal(t, "remote_compaction_v2", upstream.lastReq.Header.Get("x-codex-beta-features")) + require.Contains(t, rec.Body.String(), `"type":"compaction"`) + require.Contains(t, rec.Body.String(), `"encrypted_content":"summary"`) + require.NotNil(t, result.ReasoningEffort) + require.Equal(t, "max", *result.ReasoningEffort) +} + +func TestOpenAIGatewayServiceForwardAPIKeyRemoteCompactV2PreservesResponsesWire(t *testing.T) { + gin.SetMode(gin.TestMode) + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader( + "data: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"compaction\",\"encrypted_content\":\"summary\"}}\n\n" + + "data: {\"type\":\"response.completed\",\"response\":{\"usage\":{\"input_tokens\":1,\"output_tokens\":2}}}\n\n" + + "data: [DONE]\n\n", + )), + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream} + account := &Account{ + ID: 11, + Name: "openai-apikey-responses", + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://example.com/v1", + "compact_model_mapping": map[string]any{ + "gpt-5.6-sol": "gpt-5.6-sol-openai-compact", + }, + }, + Extra: map[string]any{"use_responses_api": true}, + Status: StatusActive, + Schedulable: true, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2") + SetOpenAIClientTransport(c, OpenAIClientTransportHTTP) + + body := []byte(`{"model":"gpt-5.6-sol","stream":true,"instructions":"response-test","input":[{"type":"message","role":"user","content":"hello"},{"type":"compaction_trigger"}],"reasoning":{"effort":"max","context":"all_turns"}}`) + result, err := svc.Forward(context.Background(), c, account, body) + + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, upstream.lastReq) + require.Equal(t, "https://example.com/v1/responses", upstream.lastReq.URL.String()) + require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(upstream.lastBody, "model").String()) + require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool()) + require.Equal(t, "compaction_trigger", gjson.GetBytes(upstream.lastBody, "input.#(type==\"compaction_trigger\").type").String()) + require.Equal(t, "max", gjson.GetBytes(upstream.lastBody, "reasoning.effort").String()) + require.Equal(t, "all_turns", gjson.GetBytes(upstream.lastBody, "reasoning.context").String()) + require.Equal(t, "remote_compaction_v2", upstream.lastReq.Header.Get("x-codex-beta-features")) + require.Contains(t, rec.Body.String(), `"type":"compaction"`) + require.Contains(t, rec.Body.String(), `"encrypted_content":"summary"`) require.NotNil(t, result.ReasoningEffort) require.Equal(t, "max", *result.ReasoningEffort) } diff --git a/backend/internal/service/openai_oauth_passthrough_test.go b/backend/internal/service/openai_oauth_passthrough_test.go index de8ecf030b..60790b9e4c 100644 --- a/backend/internal/service/openai_oauth_passthrough_test.go +++ b/backend/internal/service/openai_oauth_passthrough_test.go @@ -347,6 +347,7 @@ func TestOpenAIGatewayService_OAuthPassthrough_StreamKeepsToolNameAndBodyNormali c.Request.Header.Set("Accept-Encoding", "gzip") c.Request.Header.Set("Proxy-Authorization", "Basic abc") c.Request.Header.Set("X-Test", "keep") + c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2") originalBody := []byte(`{"model":"gpt-5.2","stream":true,"store":true,"instructions":"local-test-instructions","input":[{"type":"text","text":"hi"}]}`) @@ -409,6 +410,7 @@ func TestOpenAIGatewayService_OAuthPassthrough_StreamKeepsToolNameAndBodyNormali require.Empty(t, upstream.lastReq.Header.Get("Accept-Encoding")) require.Empty(t, upstream.lastReq.Header.Get("Proxy-Authorization")) require.Empty(t, upstream.lastReq.Header.Get("X-Test")) + require.Equal(t, "remote_compaction_v2", upstream.lastReq.Header.Get("x-codex-beta-features")) // 3) required OAuth headers are present require.Equal(t, "chatgpt.com", upstream.lastReq.Host) @@ -1373,6 +1375,7 @@ func TestOpenAIGatewayService_APIKeyPassthrough_PreservesBodyAndUsesResponsesEnd c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(nil)) c.Request.Header.Set("User-Agent", "curl/8.0") c.Request.Header.Set("X-Test", "keep") + c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2") originalBody := []byte(`{"model":"gpt-5.2","stream":false,"service_tier":"flex","max_output_tokens":128,"input":[{"type":"text","text":"hi"}]}`) resp := &http.Response{ @@ -1410,6 +1413,7 @@ func TestOpenAIGatewayService_APIKeyPassthrough_PreservesBodyAndUsesResponsesEnd require.Equal(t, "https://api.openai.com/v1/responses", upstream.lastReq.URL.String()) require.Equal(t, "Bearer sk-api-key", upstream.lastReq.Header.Get("Authorization")) require.Equal(t, "curl/8.0", upstream.lastReq.Header.Get("User-Agent")) + require.Equal(t, "remote_compaction_v2", upstream.lastReq.Header.Get("x-codex-beta-features")) require.Empty(t, upstream.lastReq.Header.Get("X-Test")) } diff --git a/backend/internal/service/openai_ws_forwarder_payload.go b/backend/internal/service/openai_ws_forwarder_payload.go index a4d47218e7..5830444815 100644 --- a/backend/internal/service/openai_ws_forwarder_payload.go +++ b/backend/internal/service/openai_ws_forwarder_payload.go @@ -74,6 +74,11 @@ func (s *OpenAIGatewayService) buildOpenAIWSHeaders( if v := strings.TrimSpace(c.Request.Header.Get("accept-language")); v != "" { headers.Set("accept-language", v) } + for _, value := range c.Request.Header.Values("x-codex-beta-features") { + if value = strings.TrimSpace(value); value != "" { + headers.Add("x-codex-beta-features", value) + } + } } // OAuth 账号:将 apiKeyID 混入 session 标识符,防止跨用户会话碰撞。 if account != nil && account.Type == AccountTypeOAuth { diff --git a/backend/internal/service/openai_ws_forwarder_success_test.go b/backend/internal/service/openai_ws_forwarder_success_test.go index adae109e09..bb4ac2242c 100644 --- a/backend/internal/service/openai_ws_forwarder_success_test.go +++ b/backend/internal/service/openai_ws_forwarder_success_test.go @@ -602,6 +602,7 @@ func TestOpenAIGatewayService_Forward_WSv2_OAuthStoreFalseByDefault(t *testing.T c.Request.Header.Set("User-Agent", "codex_cli_rs/0.98.0") c.Request.Header.Set("session_id", "sess-oauth-1") c.Request.Header.Set("conversation_id", "conv-oauth-1") + c.Request.Header.Set("x-codex-beta-features", "remote_compaction_v2") cfg := &config.Config{} cfg.Security.URLAllowlist.Enabled = false @@ -661,6 +662,7 @@ func TestOpenAIGatewayService_Forward_WSv2_OAuthStoreFalseByDefault(t *testing.T require.True(t, gjson.Get(requestJSON, "stream").Exists(), "WSv2 payload 应保留 stream 字段") require.True(t, gjson.Get(requestJSON, "stream").Bool(), "OAuth Codex 规范化后应强制 stream=true") require.Equal(t, openAIWSBetaV2Value, captureDialer.lastHeaders.Get("OpenAI-Beta")) + require.Equal(t, "remote_compaction_v2", captureDialer.lastHeaders.Get("x-codex-beta-features")) // OAuth 账号的 session_id/conversation_id 应被 isolateOpenAISessionID 隔离, // 测试中未设置 api_key 到 context,apiKeyID=0。 require.Equal(t, isolateOpenAISessionID(0, "sess-oauth-1"), captureDialer.lastHeaders.Get("session_id")) diff --git a/backend/internal/service/openai_ws_pool.go b/backend/internal/service/openai_ws_pool.go index 5950e02841..329908e762 100644 --- a/backend/internal/service/openai_ws_pool.go +++ b/backend/internal/service/openai_ws_pool.go @@ -218,6 +218,9 @@ func (l *openAIWSConnLease) Release() { return } l.conn.release() + if l.pool != nil { + l.pool.notifyAccountPoolChanged(l.accountID) + } } type openAIWSConn struct { @@ -225,6 +228,7 @@ type openAIWSConn struct { ws openAIWSClientConn handshakeHeaders http.Header + betaFeatures string leaseCh chan struct{} closedCh chan struct{} @@ -498,6 +502,10 @@ func (c *openAIWSConn) handshakeHeader(name string) string { return strings.TrimSpace(c.handshakeHeaders.Get(strings.TrimSpace(name))) } +func (c *openAIWSConn) matchesBetaFeatures(betaFeatures string) bool { + return c != nil && c.betaFeatures == betaFeatures +} + func (c *openAIWSConn) isPrewarmed() bool { if c == nil { return false @@ -516,6 +524,7 @@ type openAIWSAccountPool struct { mu sync.Mutex conns map[string]*openAIWSConn pinnedConns map[string]int + changedCh chan struct{} creating int lastCleanupAt time.Time lastAcquire *openAIWSAcquireRequest @@ -525,6 +534,23 @@ type openAIWSAccountPool struct { prewarmFailAt time.Time } +func (ap *openAIWSAccountPool) changeChannelLocked() chan struct{} { + if ap.changedCh == nil { + ap.changedCh = make(chan struct{}) + } + return ap.changedCh +} + +func (ap *openAIWSAccountPool) signalChangedLocked() { + if ap == nil { + return + } + if ap.changedCh != nil { + close(ap.changedCh) + } + ap.changedCh = make(chan struct{}) +} + type OpenAIWSPoolMetricsSnapshot struct { AcquireTotal int64 AcquireReuseTotal int64 @@ -786,7 +812,9 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque return nil, errors.New("ws url is empty") } +retryAcquire: accountID := req.Account.ID + betaFeatures := normalizeOpenAIWSBetaFeatures(req.Headers) effectiveMaxConns := p.effectiveMaxConnsByAccount(req.Account) if effectiveMaxConns <= 0 { return nil, errOpenAIWSConnQueueFull @@ -814,7 +842,7 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque return nil, errOpenAIWSPreferredConnUnavailable } preferredConn, ok := ap.conns[preferredConnID] - if !ok || preferredConn == nil { + if !ok || !preferredConn.matchesBetaFeatures(betaFeatures) { p.recordConnPickDuration(time.Since(pickStartedAt)) ap.mu.Unlock() closeOpenAIWSConns(evicted) @@ -895,7 +923,7 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque } if preferredConnID != "" { - if conn, ok := ap.conns[preferredConnID]; ok && conn.tryAcquire() { + if conn, ok := ap.conns[preferredConnID]; ok && conn.matchesBetaFeatures(betaFeatures) && conn.tryAcquire() { connPick := time.Since(pickStartedAt) p.recordConnPickDuration(connPick) ap.mu.Unlock() @@ -917,7 +945,7 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque } } - best := p.pickLeastBusyConnLocked(ap, "") + best := p.pickLeastBusyConnLocked(ap, "", betaFeatures) if best != nil && best.tryAcquire() { connPick := time.Since(pickStartedAt) p.recordConnPickDuration(connPick) @@ -939,7 +967,7 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque return lease, nil } for _, conn := range ap.conns { - if conn == nil || conn == best { + if conn == nil || conn == best || !conn.matchesBetaFeatures(betaFeatures) { continue } if conn.tryAcquire() { @@ -965,6 +993,37 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque } } + if !req.ForceNewConn && len(ap.conns)+ap.creating >= effectiveMaxConns { + compatible := p.pickLeastBusyConnLocked(ap, "", betaFeatures) + if idle := p.pickOldestIdleConnWithDifferentBetaFeaturesLocked(ap, betaFeatures); idle != nil { + delete(ap.conns, idle.id) + evicted = append(evicted, idle) + p.metrics.scaleDownTotal.Add(1) + } else if compatible == nil { + hasConnection := false + for _, conn := range ap.conns { + if conn != nil { + hasConnection = true + break + } + } + if !hasConnection && ap.creating == 0 { + ap.mu.Unlock() + closeOpenAIWSConns(evicted) + return nil, errOpenAIWSConnClosed + } + changedCh := ap.changeChannelLocked() + ap.mu.Unlock() + closeOpenAIWSConns(evicted) + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-changedCh: + goto retryAcquire + } + } + } + if req.ForceNewConn && len(ap.conns)+ap.creating >= effectiveMaxConns { if idle := p.pickOldestIdleConnLocked(ap); idle != nil { delete(ap.conns, idle.id) @@ -988,6 +1047,7 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque if dialErr != nil { ap.prewarmFails++ ap.prewarmFailAt = time.Now() + ap.signalChangedLocked() ap.mu.Unlock() return nil, dialErr } @@ -1016,7 +1076,7 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque return nil, errOpenAIWSConnQueueFull } - target := p.pickLeastBusyConnLocked(ap, req.PreferredConnID) + target := p.pickLeastBusyConnLocked(ap, req.PreferredConnID, betaFeatures) connPick := time.Since(pickStartedAt) p.recordConnPickDuration(connPick) if target == nil { @@ -1089,6 +1149,22 @@ func (p *openAIWSConnPool) pickOldestIdleConnLocked(ap *openAIWSAccountPool) *op return oldest } +func (p *openAIWSConnPool) pickOldestIdleConnWithDifferentBetaFeaturesLocked(ap *openAIWSAccountPool, betaFeatures string) *openAIWSConn { + if ap == nil || len(ap.conns) == 0 { + return nil + } + var oldest *openAIWSConn + for _, conn := range ap.conns { + if conn == nil || conn.matchesBetaFeatures(betaFeatures) || conn.isLeased() || conn.waiters.Load() > 0 || p.isConnPinnedLocked(ap, conn.id) { + continue + } + if oldest == nil || conn.lastUsedAt().Before(oldest.lastUsedAt()) { + oldest = conn + } + } + return oldest +} + func (p *openAIWSConnPool) getOrCreateAccountPool(accountID int64) *openAIWSAccountPool { if p == nil || accountID <= 0 { return nil @@ -1101,6 +1177,7 @@ func (p *openAIWSConnPool) getOrCreateAccountPool(accountID int64) *openAIWSAcco ap := &openAIWSAccountPool{ conns: make(map[string]*openAIWSConn), pinnedConns: make(map[string]int), + changedCh: make(chan struct{}), } actual, _ := p.accounts.LoadOrStore(accountID, ap) if typed, ok := actual.(*openAIWSAccountPool); ok && typed != nil { @@ -1126,6 +1203,16 @@ func (p *openAIWSConnPool) getAccountPool(accountID int64) (*openAIWSAccountPool return ap, typed && ap != nil } +func (p *openAIWSConnPool) notifyAccountPoolChanged(accountID int64) { + ap, ok := p.getAccountPool(accountID) + if !ok || ap == nil { + return + } + ap.mu.Lock() + ap.signalChangedLocked() + ap.mu.Unlock() +} + func (p *openAIWSConnPool) isConnPinnedLocked(ap *openAIWSAccountPool, connID string) bool { if ap == nil || connID == "" || len(ap.pinnedConns) == 0 { return false @@ -1212,17 +1299,20 @@ func (p *openAIWSConnPool) cleanupAccountLocked(ap *openAIWSAccountPool, now tim p.metrics.scaleDownTotal.Add(int64(redundant)) } } + if len(evicted) > 0 { + ap.signalChangedLocked() + } return evicted } -func (p *openAIWSConnPool) pickLeastBusyConnLocked(ap *openAIWSAccountPool, preferredConnID string) *openAIWSConn { +func (p *openAIWSConnPool) pickLeastBusyConnLocked(ap *openAIWSAccountPool, preferredConnID, betaFeatures string) *openAIWSConn { if ap == nil || len(ap.conns) == 0 { return nil } preferredConnID = stringsTrim(preferredConnID) if preferredConnID != "" { - if conn, ok := ap.conns[preferredConnID]; ok { + if conn, ok := ap.conns[preferredConnID]; ok && conn.matchesBetaFeatures(betaFeatures) { return conn } } @@ -1230,7 +1320,7 @@ func (p *openAIWSConnPool) pickLeastBusyConnLocked(ap *openAIWSAccountPool, pref var bestWaiters int32 var bestLastUsed time.Time for _, conn := range ap.conns { - if conn == nil { + if conn == nil || !conn.matchesBetaFeatures(betaFeatures) { continue } waiters := conn.waiters.Load() @@ -1395,10 +1485,12 @@ func (p *openAIWSConnPool) prewarmConns(accountID int64, req openAIWSAcquireRequ if err != nil { ap.prewarmFails++ ap.prewarmFailAt = time.Now() + ap.signalChangedLocked() ap.mu.Unlock() continue } if len(ap.conns) >= p.effectiveMaxConnsByAccount(req.Account) { + ap.signalChangedLocked() ap.mu.Unlock() conn.close() continue @@ -1406,6 +1498,7 @@ func (p *openAIWSConnPool) prewarmConns(accountID int64, req openAIWSAcquireRequ ap.conns[conn.id] = conn ap.prewarmFails = 0 ap.prewarmFailAt = time.Time{} + ap.signalChangedLocked() ap.mu.Unlock() } } @@ -1424,6 +1517,7 @@ func (p *openAIWSConnPool) evictConn(accountID int64, connID string) { if len(ap.pinnedConns) > 0 { delete(ap.pinnedConns, connID) } + ap.signalChangedLocked() } ap.mu.Unlock() } @@ -1476,9 +1570,11 @@ func (p *openAIWSConnPool) UnpinConn(accountID int64, connID string) { count := ap.pinnedConns[connID] if count <= 1 { delete(ap.pinnedConns, connID) + ap.signalChangedLocked() return } ap.pinnedConns[connID] = count - 1 + ap.signalChangedLocked() } func (p *openAIWSConnPool) dialConn(ctx context.Context, req openAIWSAcquireRequest) (*openAIWSConn, error) { @@ -1501,7 +1597,9 @@ func (p *openAIWSConnPool) dialConn(ctx context.Context, req openAIWSAcquireRequ } } id := p.nextConnID(req.Account.ID) - return newOpenAIWSConn(id, req.Account.ID, conn, handshakeHeaders), nil + pooledConn := newOpenAIWSConn(id, req.Account.ID, conn, handshakeHeaders) + pooledConn.betaFeatures = normalizeOpenAIWSBetaFeatures(req.Headers) + return pooledConn, nil } func (p *openAIWSConnPool) nextConnID(accountID int64) string { @@ -1679,6 +1777,31 @@ func cloneOpenAIWSAcquireRequestPtr(req *openAIWSAcquireRequest) *openAIWSAcquir return &copied } +func normalizeOpenAIWSBetaFeatures(headers http.Header) string { + features := make(map[string]struct{}) + for name, values := range headers { + if !strings.EqualFold(strings.TrimSpace(name), "x-codex-beta-features") { + continue + } + for _, value := range values { + for _, feature := range strings.Split(value, ",") { + if feature = strings.TrimSpace(feature); feature != "" { + features[feature] = struct{}{} + } + } + } + } + if len(features) == 0 { + return "" + } + normalized := make([]string, 0, len(features)) + for feature := range features { + normalized = append(normalized, feature) + } + sort.Strings(normalized) + return strings.Join(normalized, ",") +} + func cloneHeader(src http.Header) http.Header { if src == nil { return nil diff --git a/backend/internal/service/openai_ws_pool_test.go b/backend/internal/service/openai_ws_pool_test.go index b2683ee041..ae9b94ce4a 100644 --- a/backend/internal/service/openai_ws_pool_test.go +++ b/backend/internal/service/openai_ws_pool_test.go @@ -342,6 +342,171 @@ func TestOpenAIWSConnPool_ForceNewConnSkipsReuse(t *testing.T) { require.Equal(t, 2, dialer.DialCount(), "ForceNewConn=true 时应跳过空闲连接复用并新建连接") } +func TestOpenAIWSConnPool_AcquireReusesOnlyMatchingBetaFeatures(t *testing.T) { + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 2 + cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0 + cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 2 + + pool := newOpenAIWSConnPool(cfg) + dialer := &openAIWSCountingDialer{} + pool.setClientDialerForTest(dialer) + + account := &Account{ID: 128, Platform: PlatformOpenAI, Type: AccountTypeAPIKey} + baseReq := openAIWSAcquireRequest{ + Account: account, + WSURL: "wss://example.com/v1/responses", + } + + plainLease, err := pool.Acquire(context.Background(), baseReq) + require.NoError(t, err) + plainConnID := plainLease.ConnID() + plainLease.Release() + + betaReq := baseReq + betaReq.Headers = http.Header{"X-Codex-Beta-Features": {" remote_compaction_v2 ", " responses_websockets_v2 "}} + betaLease, err := pool.Acquire(context.Background(), betaReq) + require.NoError(t, err) + require.False(t, betaLease.Reused()) + require.NotEqual(t, plainConnID, betaLease.ConnID()) + betaConnID := betaLease.ConnID() + betaLease.Release() + + reorderedReq := baseReq + reorderedReq.Headers = http.Header{"X-Codex-Beta-Features": {"responses_websockets_v2,remote_compaction_v2"}} + reorderedLease, err := pool.Acquire(context.Background(), reorderedReq) + require.NoError(t, err) + require.True(t, reorderedLease.Reused()) + require.Equal(t, betaConnID, reorderedLease.ConnID()) + reorderedLease.Release() + + _, err = pool.Acquire(context.Background(), openAIWSAcquireRequest{ + Account: account, + WSURL: baseReq.WSURL, + Headers: betaReq.Headers, + PreferredConnID: plainConnID, + ForcePreferredConn: true, + }) + require.ErrorIs(t, err, errOpenAIWSPreferredConnUnavailable) + + plainLease, err = pool.Acquire(context.Background(), baseReq) + require.NoError(t, err) + require.True(t, plainLease.Reused()) + require.Equal(t, plainConnID, plainLease.ConnID()) + plainLease.Release() + + require.Equal(t, 2, dialer.DialCount()) +} + +func TestOpenAIWSConnPool_AcquireReplacesIdleConnWithDifferentBetaFeatures(t *testing.T) { + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1 + cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0 + cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1 + + pool := newOpenAIWSConnPool(cfg) + dialer := &openAIWSCountingDialer{} + pool.setClientDialerForTest(dialer) + + account := &Account{ID: 129, Platform: PlatformOpenAI, Type: AccountTypeAPIKey} + plainLease, err := pool.Acquire(context.Background(), openAIWSAcquireRequest{ + Account: account, + WSURL: "wss://example.com/v1/responses", + }) + require.NoError(t, err) + plainConnID := plainLease.ConnID() + plainLease.Release() + + betaLease, err := pool.Acquire(context.Background(), openAIWSAcquireRequest{ + Account: account, + WSURL: "wss://example.com/v1/responses", + Headers: http.Header{"X-Codex-Beta-Features": {"remote_compaction_v2"}}, + }) + require.NoError(t, err) + require.False(t, betaLease.Reused()) + require.NotEqual(t, plainConnID, betaLease.ConnID()) + betaLease.Release() + + require.Equal(t, 2, dialer.DialCount()) +} + +func TestOpenAIWSConnPool_AcquireWaitsForBusyIncompatibleConnection(t *testing.T) { + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1 + cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0 + cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1 + + pool := newOpenAIWSConnPool(cfg) + dialer := &openAIWSCountingDialer{} + pool.setClientDialerForTest(dialer) + account := &Account{ID: 130, Platform: PlatformOpenAI, Type: AccountTypeAPIKey} + baseReq := openAIWSAcquireRequest{Account: account, WSURL: "wss://example.com/v1/responses"} + + plainLease, err := pool.Acquire(context.Background(), baseReq) + require.NoError(t, err) + plainConnID := plainLease.ConnID() + + type acquireResult struct { + lease *openAIWSConnLease + err error + } + resultCh := make(chan acquireResult, 1) + var done atomic.Bool + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + go func() { + betaReq := baseReq + betaReq.Headers = http.Header{"X-Codex-Beta-Features": {"remote_compaction_v2"}} + lease, acquireErr := pool.Acquire(ctx, betaReq) + resultCh <- acquireResult{lease: lease, err: acquireErr} + done.Store(true) + }() + + require.Never(t, done.Load, 50*time.Millisecond, 5*time.Millisecond) + plainLease.Release() + + result := <-resultCh + require.NoError(t, result.err) + require.NotNil(t, result.lease) + require.False(t, result.lease.Reused()) + require.NotEqual(t, plainConnID, result.lease.ConnID()) + result.lease.Release() + require.Equal(t, 2, dialer.DialCount()) +} + +func TestOpenAIWSConnPool_AcquireReplacesIncompatibleIdleWhenMatchingBusy(t *testing.T) { + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 2 + cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0 + cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 2 + + pool := newOpenAIWSConnPool(cfg) + dialer := &openAIWSCountingDialer{} + pool.setClientDialerForTest(dialer) + account := &Account{ID: 131, Platform: PlatformOpenAI, Type: AccountTypeAPIKey} + baseReq := openAIWSAcquireRequest{Account: account, WSURL: "wss://example.com/v1/responses"} + + plainLease, err := pool.Acquire(context.Background(), baseReq) + require.NoError(t, err) + plainConnID := plainLease.ConnID() + plainLease.Release() + + betaReq := baseReq + betaReq.Headers = http.Header{"X-Codex-Beta-Features": {"remote_compaction_v2"}} + busyBetaLease, err := pool.Acquire(context.Background(), betaReq) + require.NoError(t, err) + + secondBetaLease, err := pool.Acquire(context.Background(), betaReq) + require.NoError(t, err) + require.False(t, secondBetaLease.Reused()) + require.NotEqual(t, plainConnID, secondBetaLease.ConnID()) + require.NotEqual(t, busyBetaLease.ConnID(), secondBetaLease.ConnID()) + + secondBetaLease.Release() + busyBetaLease.Release() + require.Equal(t, 3, dialer.DialCount()) +} + func TestOpenAIWSConnPool_AcquireForcePreferredConnUnavailable(t *testing.T) { cfg := &config.Config{} cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 2 From 52071d391b5b2a4e4e0940aea85fc731857c6d07 Mon Sep 17 00:00:00 2001 From: Tian Lee <498756723@qq.com> Date: Sat, 11 Jul 2026 20:31:33 +0800 Subject: [PATCH 5/9] =?UTF-8?q?fix(openai):=20=E8=BD=AC=E5=8F=91=20Codex?= =?UTF-8?q?=20alpha/search=20=E7=8B=AC=E7=AB=8B=E6=90=9C=E7=B4=A2=E7=AB=AF?= =?UTF-8?q?=E7=82=B9?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/internal/handler/endpoint.go | 9 +- backend/internal/handler/endpoint_test.go | 4 + .../internal/handler/openai_alpha_search.go | 194 ++++++++++++++++++ backend/internal/server/routes/gateway.go | 3 + .../internal/server/routes/gateway_test.go | 30 +++ .../internal/service/openai_alpha_search.go | 151 ++++++++++++++ .../service/openai_alpha_search_test.go | 139 +++++++++++++ 7 files changed, 527 insertions(+), 3 deletions(-) create mode 100644 backend/internal/handler/openai_alpha_search.go create mode 100644 backend/internal/service/openai_alpha_search.go create mode 100644 backend/internal/service/openai_alpha_search_test.go diff --git a/backend/internal/handler/endpoint.go b/backend/internal/handler/endpoint.go index 0b9930c5cc..55d138d3b4 100644 --- a/backend/internal/handler/endpoint.go +++ b/backend/internal/handler/endpoint.go @@ -18,6 +18,7 @@ const ( EndpointMessages = "/v1/messages" EndpointChatCompletions = "/v1/chat/completions" EndpointEmbeddings = "/v1/embeddings" + EndpointAlphaSearch = "/v1/alpha/search" EndpointResponses = "/v1/responses" EndpointResponsesCompact = "/v1/responses/compact" EndpointImagesGenerations = "/v1/images/generations" @@ -75,6 +76,8 @@ func NormalizeInboundEndpoint(path string) string { switch { case strings.Contains(path, EndpointEmbeddings): return EndpointEmbeddings + case strings.Contains(path, EndpointAlphaSearch) || isBareOrSubpathOf(strings.TrimRight(path, "/"), "/alpha/search") || isBareOrSubpathOf(strings.TrimRight(path, "/"), "/backend-api/codex/alpha/search"): + return EndpointAlphaSearch case strings.Contains(path, EndpointChatCompletions): return EndpointChatCompletions case strings.Contains(path, EndpointMessages): @@ -155,8 +158,8 @@ func isBareOrSubpathOf(path, root string) bool { // account platform and the normalized inbound endpoint. // // Platform-specific rules: -// - OpenAI always forwards to /v1/responses (with optional subpath -// such as /v1/responses/compact preserved from the raw URL). +// - OpenAI text compatibility routes forward to /v1/responses; native +// endpoints such as embeddings and alpha search retain their paths. // - Anthropic → /v1/messages // - Gemini → /v1beta/models // - Antigravity → /v1/messages (Claude) or gemini (Gemini) @@ -167,7 +170,7 @@ func DeriveUpstreamEndpoint(inbound, rawRequestPath, platform string) string { switch platform { case service.PlatformOpenAI, service.PlatformGrok: - if inbound == EndpointEmbeddings || inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits || inbound == EndpointVideosGenerations || inbound == EndpointVideos { + if inbound == EndpointEmbeddings || inbound == EndpointAlphaSearch || inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits || inbound == EndpointVideosGenerations || inbound == EndpointVideos { 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 96ed1292b3..e0e26805f7 100644 --- a/backend/internal/handler/endpoint_test.go +++ b/backend/internal/handler/endpoint_test.go @@ -25,6 +25,7 @@ func TestNormalizeInboundEndpoint(t *testing.T) { {"/v1/messages", EndpointMessages}, {"/v1/chat/completions", EndpointChatCompletions}, {"/v1/embeddings", EndpointEmbeddings}, + {"/v1/alpha/search", EndpointAlphaSearch}, {"/v1/responses", EndpointResponses}, {"/v1/responses/compact", EndpointResponsesCompact}, {"/v1/responses/compact/detail", EndpointResponsesCompact}, @@ -50,11 +51,13 @@ func TestNormalizeInboundEndpoint(t *testing.T) { {"/responses", EndpointResponses}, {"/responses/compact", EndpointResponsesCompact}, {"/responses/compact/detail", EndpointResponsesCompact}, + {"/alpha/search", EndpointAlphaSearch}, // Bare Codex direct alias route — root vs. compact. {"/backend-api/codex/responses", EndpointResponses}, {"/backend-api/codex/responses/compact", EndpointResponsesCompact}, {"/backend-api/codex/responses/compact/detail", EndpointResponsesCompact}, + {"/backend-api/codex/alpha/search", EndpointAlphaSearch}, // Must NOT generalize to arbitrary paths merely ending in // "/responses" (or "/responses/compact") that are unrelated to @@ -119,6 +122,7 @@ func TestDeriveUpstreamEndpoint(t *testing.T) { {"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 alpha search", EndpointAlphaSearch, "/backend-api/codex/alpha/search", service.PlatformOpenAI, EndpointAlphaSearch}, {"openai image generations", EndpointImagesGenerations, "/v1/images/generations", service.PlatformOpenAI, EndpointImagesGenerations}, {"openai image edits", EndpointImagesEdits, "/openai/v1/images/edits", service.PlatformOpenAI, EndpointImagesEdits}, {"grok video generations", EndpointVideosGenerations, "/v1/videos/generations", service.PlatformGrok, EndpointVideosGenerations}, diff --git a/backend/internal/handler/openai_alpha_search.go b/backend/internal/handler/openai_alpha_search.go new file mode 100644 index 0000000000..3532e42363 --- /dev/null +++ b/backend/internal/handler/openai_alpha_search.go @@ -0,0 +1,194 @@ +package handler + +import ( + "errors" + "net/http" + "strconv" + "strings" + "time" + + pkghttputil "github.com/Wei-Shaw/sub2api/internal/pkg/httputil" + 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" +) + +// AlphaSearch proxies the standalone search endpoint used by Codex Responses Lite. +func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) { + streamStarted := false + defer h.recoverResponsesPanic(c, &streamStarted) + setOpenAIClientTransportHTTP(c) + requestStart := time.Now() + + apiKey, ok := middleware2.GetAPIKeyFromContext(c) + if !ok || apiKey.Group == nil { + h.errorResponse(c, http.StatusUnauthorized, "authentication_error", "Invalid API key") + return + } + if apiKey.Group.Platform != service.PlatformOpenAI { + h.errorResponse(c, http.StatusNotFound, "not_found_error", "Codex alpha search is only available for OpenAI groups") + 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.alpha_search", + 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) { + logRequestBodyParseFailure(reqLog, body, nil) + 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 + } + requestedModel := strings.TrimSpace(modelResult.String()) + reqLog = reqLog.With(zap.String("model", requestedModel)) + setOpsRequestContext(c, requestedModel, false) + setOpsEndpointContext(c, "", int16(service.RequestTypeSync)) + + channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, requestedModel) + forwardBody := openAIModelMappedBody(body, channelMapping.Mapped, channelMapping.MappedModel, h.gatewayService.ReplaceModelInBody) + subscription, _ := middleware2.GetSubscriptionFromContext(c) + service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds()) + + userRelease, acquired := h.acquireResponsesUserSlot(c, subject.UserID, subject.Concurrency, false, &streamStarted, reqLog) + if !acquired { + return + } + if userRelease != nil { + defer userRelease() + } + + if err := h.billingCacheService.CheckBillingEligibility(c.Request.Context(), apiKey.User, apiKey, apiKey.Group, subscription, service.QuotaPlatform(c.Request.Context(), apiKey)); err != nil { + status, code, message, retryAfter := billingErrorDetails(err) + if retryAfter > 0 { + c.Header("Retry-After", strconv.Itoa(retryAfter)) + } + h.errorResponse(c, status, code, message) + return + } + + searchID := strings.TrimSpace(gjson.GetBytes(body, "id").String()) + sessionHash := h.gatewayService.GenerateSessionHashWithFallback(c, nil, searchID) + failedAccountIDs := make(map[int64]struct{}) + var lastFailoverErr *service.UpstreamFailoverError + switchCount := 0 + routingStart := time.Now() + + for { + selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability( + c.Request.Context(), + apiKey.GroupID, + "", + sessionHash, + requestedModel, + failedAccountIDs, + service.OpenAIUpstreamTransportHTTPSSE, + service.OpenAIEndpointCapabilityChatCompletions, + false, + false, + service.PlatformOpenAI, + ) + if err != nil || selection == nil || selection.Account == nil { + if len(failedAccountIDs) == 0 { + cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, requestedModel, requestedModel, service.PlatformOpenAI) + if !cls.ModelNotFound { + markOpsRoutingCapacityLimitedIfNoAvailable(c, err) + } + h.errorResponse(c, cls.Status, cls.ErrType, cls.Message) + return + } + if lastFailoverErr != nil { + h.handleFailoverExhausted(c, lastFailoverErr, false) + } else { + h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Upstream request failed") + } + return + } + + account := selection.Account + setOpsSelectedAccount(c, account.ID, account.Platform) + accountRelease, acquired := h.acquireResponsesAccountSlot(c, apiKey.GroupID, sessionHash, selection, false, &streamStarted, reqLog) + if !acquired { + return + } + service.SetOpsLatencyMs(c, service.OpsRoutingLatencyMsKey, time.Since(routingStart).Milliseconds()) + writerSizeBeforeForward := c.Writer.Size() + forwardStart := time.Now() + err = func() error { + if accountRelease != nil { + defer accountRelease() + } + return h.gatewayService.ForwardAlphaSearch(c.Request.Context(), c, account, forwardBody) + }() + service.SetOpsLatencyMs(c, service.OpsResponseLatencyMsKey, time.Since(forwardStart).Milliseconds()) + + if err == nil { + h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, true, nil) + return + } + + var failoverErr *service.UpstreamFailoverError + if !errors.As(err, &failoverErr) { + 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_alpha_search.forward_failed", zap.Int64("account_id", account.ID), zap.Error(err)) + return + } + + h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) + if c.Writer.Size() != writerSizeBeforeForward { + h.handleFailoverExhausted(c, failoverErr, true) + return + } + h.gatewayService.RecordOpenAIAccountSwitch() + failedAccountIDs[account.ID] = struct{}{} + lastFailoverErr = failoverErr + if switchCount >= h.maxAccountSwitches { + h.handleFailoverExhausted(c, failoverErr, false) + return + } + switchCount++ + if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount) { + h.handleFailoverExhausted(c, failoverErr, false) + return + } + reqLog.Warn("openai_alpha_search.upstream_failover_switching", + zap.Int64("account_id", account.ID), + zap.Int("upstream_status", failoverErr.StatusCode), + zap.Int("switch_count", switchCount), + ) + } +} diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index ba5b4f61d1..7960137604 100644 --- a/backend/internal/server/routes/gateway.go +++ b/backend/internal/server/routes/gateway.go @@ -147,6 +147,7 @@ func RegisterGatewayRoutes( } h.Gateway.Responses(c) }) + gateway.POST("/alpha/search", h.OpenAIGateway.AlphaSearch) gateway.GET("/responses", func(c *gin.Context) { h.OpenAIGateway.ResponsesWebSocket(c) }) @@ -212,6 +213,7 @@ func RegisterGatewayRoutes( } r.POST("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, responsesHandler) r.POST("/responses/*subpath", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, responsesHandler) + r.POST("/alpha/search", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.OpenAIGateway.AlphaSearch) r.GET("/responses", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, func(c *gin.Context) { h.OpenAIGateway.ResponsesWebSocket(c) }) @@ -220,6 +222,7 @@ func RegisterGatewayRoutes( { codexDirect.POST("/responses", responsesHandler) codexDirect.POST("/responses/*subpath", responsesHandler) + codexDirect.POST("/alpha/search", h.OpenAIGateway.AlphaSearch) codexDirect.GET("/responses", func(c *gin.Context) { h.OpenAIGateway.ResponsesWebSocket(c) }) diff --git a/backend/internal/server/routes/gateway_test.go b/backend/internal/server/routes/gateway_test.go index 2779dd8f01..6b15fbfa9b 100644 --- a/backend/internal/server/routes/gateway_test.go +++ b/backend/internal/server/routes/gateway_test.go @@ -65,6 +65,36 @@ func TestGatewayRoutesOpenAIResponsesCompactPathIsRegistered(t *testing.T) { } } +func TestGatewayRoutesOpenAIAlphaSearchPathsAreRegistered(t *testing.T) { + router := newGatewayRoutesTestRouter() + registered := make(map[string]bool) + for _, route := range router.Routes() { + if route.Method == http.MethodPost { + registered[route.Path] = true + } + } + + for _, path := range []string{ + "/v1/alpha/search", + "/alpha/search", + "/backend-api/codex/alpha/search", + } { + require.True(t, registered[path], "POST %s should be registered", path) + } +} + +func TestGatewayRoutesAlphaSearchRejectsNonOpenAIGroup(t *testing.T) { + router := newGatewayRoutesTestRouter(service.PlatformGrok) + req := httptest.NewRequest(http.MethodPost, "/v1/alpha/search", strings.NewReader(`{"model":"gpt-5.6-sol"}`)) + req.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + + router.ServeHTTP(w, req) + + require.Equal(t, http.StatusNotFound, w.Code) + require.Contains(t, w.Body.String(), "only available for OpenAI groups") +} + func TestGatewayRoutesOpenAIImagesPathsAreRegistered(t *testing.T) { router := newGatewayRoutesTestRouter() diff --git a/backend/internal/service/openai_alpha_search.go b/backend/internal/service/openai_alpha_search.go new file mode 100644 index 0000000000..ecc4496e66 --- /dev/null +++ b/backend/internal/service/openai_alpha_search.go @@ -0,0 +1,151 @@ +package service + +import ( + "bytes" + "context" + "fmt" + "io" + "net/http" + "net/url" + "strings" + "time" + + "github.com/gin-gonic/gin" + "github.com/tidwall/gjson" +) + +const ( + chatgptCodexAlphaSearchURL = "https://chatgpt.com/backend-api/codex/alpha/search" + openAIPlatformAlphaSearchURL = "https://api.openai.com/v1/alpha/search" +) + +// ForwardAlphaSearch proxies Codex standalone web search without binding the +// evolving alpha request or response schema. +func (s *OpenAIGatewayService) ForwardAlphaSearch(ctx context.Context, c *gin.Context, account *Account, body []byte) error { + if s == nil || c == nil || account == nil { + return fmt.Errorf("service, context, and account are required") + } + modelResult := gjson.GetBytes(body, "model") + requestedModel := strings.TrimSpace(modelResult.String()) + if modelResult.Type != gjson.String || requestedModel == "" { + return fmt.Errorf("model is required") + } + + upstreamModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(requestedModel)) + if upstreamModel != "" && upstreamModel != requestedModel { + body = ReplaceModelInBody(body, upstreamModel) + } + + token, _, err := s.GetAccessToken(ctx, account) + if err != nil { + return err + } + + req, err := s.buildOpenAIAlphaSearchRequest(ctx, c, account, body, token) + if err != nil { + return err + } + + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + upstreamStart := time.Now() + resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, account.Concurrency) + SetOpsLatencyMs(c, OpsUpstreamLatencyMsKey, time.Since(upstreamStart).Milliseconds()) + if err != nil { + return s.handleOpenAIUpstreamTransportError(ctx, c, account, err, true) + } + defer func() { _ = resp.Body.Close() }() + + respBody, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError) + if err != nil { + return fmt.Errorf("read alpha search response: %w", err) + } + + if resp.StatusCode >= http.StatusBadRequest { + upstreamMessage := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(respBody))) + if s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMessage, respBody) { + resp.Body = io.NopCloser(bytes.NewReader(respBody)) + s.handleFailoverSideEffects(ctx, resp, account, respBody, upstreamModel) + return &UpstreamFailoverError{ + StatusCode: resp.StatusCode, + ResponseBody: respBody, + RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), + } + } + } + + if !account.IsShadow() { + s.UpdateCodexUsageSnapshotFromHeaders(ctx, account.ID, resp.Header) + } + writeOpenAIPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) + contentType := resp.Header.Get("Content-Type") + if contentType == "" { + contentType = "application/json" + } + c.Data(resp.StatusCode, contentType, respBody) + return nil +} + +func (s *OpenAIGatewayService) buildOpenAIAlphaSearchRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token string) (*http.Request, error) { + clientBeta := "" + if c != nil { + clientBeta = c.GetHeader("OpenAI-Beta") + } + req, err := s.buildUpstreamRequestOpenAIPassthrough(ctx, c, account, body, token) + if err != nil { + return nil, err + } + + targetURL, err := s.openAIAlphaSearchURL(account) + if err != nil { + return nil, err + } + parsedURL, err := url.Parse(targetURL) + if err != nil { + return nil, fmt.Errorf("parse alpha search URL: %w", err) + } + if c != nil && c.Request != nil && c.Request.URL != nil { + query := parsedURL.Query() + for key, values := range c.Request.URL.Query() { + for _, value := range values { + query.Add(key, value) + } + } + parsedURL.RawQuery = query.Encode() + } + req.URL = parsedURL + req.Header.Set("Accept", "application/json") + if clientBeta == "" { + req.Header.Del("OpenAI-Beta") + } + if version := strings.TrimSpace(c.GetHeader("Version")); version != "" { + req.Header.Set("Version", version) + } else if account.Type == AccountTypeOAuth { + req.Header.Set("Version", codexCLIVersion) + } + return req, nil +} + +func (s *OpenAIGatewayService) openAIAlphaSearchURL(account *Account) (string, error) { + if account == nil { + return "", fmt.Errorf("account is required") + } + switch account.Type { + case AccountTypeOAuth: + return chatgptCodexAlphaSearchURL, nil + case AccountTypeAPIKey: + baseURL := account.GetOpenAIBaseURL() + if baseURL == "" { + return openAIPlatformAlphaSearchURL, nil + } + validatedURL, err := s.validateUpstreamBaseURL(baseURL) + if err != nil { + return "", err + } + return buildOpenAIEndpointURL(validatedURL, "/v1/alpha/search"), nil + default: + return "", fmt.Errorf("unsupported OpenAI account type: %s", account.Type) + } +} diff --git a/backend/internal/service/openai_alpha_search_test.go b/backend/internal/service/openai_alpha_search_test.go new file mode 100644 index 0000000000..52e5bc36ca --- /dev/null +++ b/backend/internal/service/openai_alpha_search_test.go @@ -0,0 +1,139 @@ +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 TestForwardAlphaSearchOAuthPreservesWire(t *testing.T) { + gin.SetMode(gin.TestMode) + body := []byte(`{ + "id":"search-session", + "model":"gpt-5.6-sol", + "reasoning":{"effort":"max","context":"all_turns"}, + "input":[{"type":"message","role":"user","content":[{"type":"input_text","text":"latest news"}]}], + "commands":{"search_query":[{"q":"OpenAI news","recency":1}]}, + "settings":{"allowed_callers":["direct"],"external_web_access":true}, + "max_output_tokens":2000, + "future_field":{"keep":true} + }`) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/alpha/search?feature=standalone", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + c.Request.Header.Set("User-Agent", codexCLIUserAgent) + c.Request.Header.Set("Originator", "codex_cli_rs") + c.Request.Header.Set("Version", "0.144.1") + + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"encrypted_output":"ciphertext","output":"search result"}`)), + }} + service := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} + account := &Account{ + ID: 42, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "oauth-token", + "chatgpt_account_id": "chatgpt-account", + }, + } + + err := service.ForwardAlphaSearch(context.Background(), c, account, body) + + require.NoError(t, err) + require.Equal(t, http.StatusOK, recorder.Code) + require.JSONEq(t, `{"encrypted_output":"ciphertext","output":"search result"}`, recorder.Body.String()) + require.Equal(t, chatgptCodexAlphaSearchURL+"?feature=standalone", upstream.lastReq.URL.String()) + require.Equal(t, "chatgpt.com", upstream.lastReq.Host) + require.Equal(t, "Bearer oauth-token", upstream.lastReq.Header.Get("Authorization")) + require.Equal(t, "chatgpt-account", upstream.lastReq.Header.Get("chatgpt-account-id")) + require.Equal(t, "application/json", upstream.lastReq.Header.Get("Accept")) + require.Equal(t, "0.144.1", upstream.lastReq.Header.Get("Version")) + require.Empty(t, upstream.lastReq.Header.Get("OpenAI-Beta")) + require.JSONEq(t, string(body), string(upstream.lastBody)) +} + +func TestForwardAlphaSearchAPIKeyMapsModelAndPassesThroughError(t *testing.T) { + gin.SetMode(gin.TestMode) + body := []byte(`{"id":"search-session","model":"gpt-5.6-sol","commands":{"search_query":[{"q":"news"}]}}`) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/alpha/search", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + + upstreamBody := `{"error":{"type":"invalid_request_error","message":"bad search"}}` + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusBadRequest, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(upstreamBody)), + }} + service := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} + account := &Account{ + ID: 7, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_key": "sk-test", + "base_url": "https://compat.example/v4", + "model_mapping": map[string]any{ + "gpt-5.6-sol": "upstream-5.6", + }, + }, + } + + err := service.ForwardAlphaSearch(context.Background(), c, account, body) + + require.NoError(t, err) + require.Equal(t, http.StatusBadRequest, recorder.Code) + require.JSONEq(t, upstreamBody, recorder.Body.String()) + require.Equal(t, "https://compat.example/v4/alpha/search", upstream.lastReq.URL.String()) + require.Equal(t, "Bearer sk-test", upstream.lastReq.Header.Get("Authorization")) + require.Equal(t, "upstream-5.6", gjson.GetBytes(upstream.lastBody, "model").String()) + require.True(t, gjson.GetBytes(upstream.lastBody, "commands.search_query").IsArray()) +} + +func TestForwardAlphaSearchReturnsFailoverBeforeWriting(t *testing.T) { + gin.SetMode(gin.TestMode) + body := []byte(`{"id":"search-session","model":"gpt-5.6-sol","commands":{}}`) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/alpha/search", bytes.NewReader(body)) + + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusTooManyRequests, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"error":{"message":"rate limited"}}`)), + }} + service := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} + account := &Account{ + ID: 8, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_key": "sk-test", + }, + } + + err := service.ForwardAlphaSearch(context.Background(), c, account, body) + + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, http.StatusTooManyRequests, failoverErr.StatusCode) + require.Equal(t, openAIPlatformAlphaSearchURL, upstream.lastReq.URL.String()) + require.False(t, c.Writer.Written()) + require.Empty(t, recorder.Body.String()) +} From d5b47c21429e405c4142c61c5d37620b09a67d4d Mon Sep 17 00:00:00 2001 From: visa2 Date: Sat, 11 Jul 2026 20:06:07 +0800 Subject: [PATCH 6/9] fix(openai): restore Codex identity for OAuth Messages --- .../internal/service/openai_codex_identity.go | 22 ++++- .../service/openai_codex_identity_test.go | 37 ++++++++- .../service/openai_compat_model_test.go | 82 +++++++++++++++++-- .../service/openai_gateway_grok_test.go | 5 ++ .../service/openai_gateway_messages.go | 21 +++-- ...nai_gateway_messages_chat_fallback_test.go | 6 ++ 6 files changed, 154 insertions(+), 19 deletions(-) diff --git a/backend/internal/service/openai_codex_identity.go b/backend/internal/service/openai_codex_identity.go index 68d4105d84..47bc19ce0c 100644 --- a/backend/internal/service/openai_codex_identity.go +++ b/backend/internal/service/openai_codex_identity.go @@ -11,12 +11,32 @@ import ( // 若请求携带 version 且低于该值,上游直接 404(issue #3901,2026-07 实测)。 const codexUpstreamMinVersion = "0.144.0" +// ensureCodexIdentityHeaders 补齐 OAuth(ChatGPT 内部接口)出站请求所需的 Codex 身份头。 +// 已有 User-Agent 与 version 保持不变,交给紧随其后的 enforceCodexIdentityHeaders +// 做官方身份配对与最低版本校正。 +func ensureCodexIdentityHeaders(h http.Header) { + if h == nil { + return + } + if strings.TrimSpace(h.Get("user-agent")) == "" { + h.Set("user-agent", codexCLIUserAgent) + } + if strings.TrimSpace(h.Get("originator")) == "" { + h.Set("originator", "codex_cli_rs") + } + if strings.TrimSpace(h.Get("version")) == "" { + h.Set("version", codexCLIVersion) + } + h.Set("OpenAI-Beta", "responses=experimental") +} + // enforceCodexIdentityHeaders 收口 OAuth(ChatGPT 内部接口)出站请求的客户端身份头。 // 上游要求 originator 与 User-Agent 首段配套且为官方客户端标识,version 头(若携带) // 不低于 0.144.0,任一不满足即 404(issue #3901)。以最终 User-Agent 为准推导配套 // originator;推导不出官方身份(第三方 UA / UA 缺失)时整体回退为默认 Codex CLI 身份。 // -// 仅对携带 originator 的请求生效——compat messages bridge 故意不带 originator,保持原样。 +// 仅对携带 originator 的请求生效;需要从缺失身份头恢复的调用方应先调用 +// ensureCodexIdentityHeaders。 // 必须在所有 User-Agent 改写(自定义 UA / ForceCodexCLI / 浏览器 UA 兜底)之后调用。 func enforceCodexIdentityHeaders(h http.Header) { if h == nil || h.Get("originator") == "" { diff --git a/backend/internal/service/openai_codex_identity_test.go b/backend/internal/service/openai_codex_identity_test.go index 7d2c6d8520..ecb6eb7f6f 100644 --- a/backend/internal/service/openai_codex_identity_test.go +++ b/backend/internal/service/openai_codex_identity_test.go @@ -7,6 +7,36 @@ import ( "github.com/stretchr/testify/require" ) +func TestEnsureCodexIdentityHeaders(t *testing.T) { + t.Run("补齐缺失身份头", func(t *testing.T) { + h := make(http.Header) + + ensureCodexIdentityHeaders(h) + enforceCodexIdentityHeaders(h) + + require.Equal(t, "codex_cli_rs", h.Get("originator")) + require.Equal(t, codexCLIUserAgent, h.Get("user-agent")) + require.Equal(t, codexCLIVersion, h.Get("version")) + require.Equal(t, "responses=experimental", h.Get("OpenAI-Beta")) + }) + + t.Run("保留已有官方UA和合法version并重新配对", func(t *testing.T) { + const tuiUA = "codex-tui/9.9.9 (Mac OS X 14.0; arm64) iTerm (codex-tui; 9.9.9)" + h := make(http.Header) + h.Set("user-agent", tuiUA) + h.Set("version", "9.9.9") + h.Set("OpenAI-Beta", "assistants=v2") + + ensureCodexIdentityHeaders(h) + enforceCodexIdentityHeaders(h) + + require.Equal(t, "codex-tui", h.Get("originator")) + require.Equal(t, tuiUA, h.Get("user-agent")) + require.Equal(t, "9.9.9", h.Get("version")) + require.Equal(t, "responses=experimental", h.Get("OpenAI-Beta")) + }) +} + func TestEnforceCodexIdentityHeaders(t *testing.T) { const tuiUA = "codex-tui/0.140.2 (Mac OS X 14.0; arm64) iTerm (codex-tui; 0.140.2)" @@ -102,13 +132,14 @@ func TestEnforceCodexIdentityHeaders(t *testing.T) { } } -// compat messages bridge 故意不带 originator:收口必须保持 no-op,不得注入身份头。 +// enforce 本身仍只负责收口:缺少 originator 时必须保持 no-op,由需要恢复身份的 +// 调用方先显式调用 ensureCodexIdentityHeaders。 func TestEnforceCodexIdentityHeaders_NoOriginatorIsNoop(t *testing.T) { h := make(http.Header) - h.Set("user-agent", "luna/1.0.0") + h.Set("user-agent", "third-party-client/1.0.0") enforceCodexIdentityHeaders(h) require.Empty(t, h.Get("originator")) - require.Equal(t, "luna/1.0.0", h.Get("user-agent")) + require.Equal(t, "third-party-client/1.0.0", h.Get("user-agent")) } diff --git a/backend/internal/service/openai_compat_model_test.go b/backend/internal/service/openai_compat_model_test.go index 7c7ac1b94f..69b6ddbca2 100644 --- a/backend/internal/service/openai_compat_model_test.go +++ b/backend/internal/service/openai_compat_model_test.go @@ -837,8 +837,7 @@ func TestForwardAsAnthropic_ReusesOAuthCodexTurnState(t *testing.T) { require.NoError(t, err) require.NotNil(t, firstResult) require.Empty(t, upstream.requests[0].Header.Get("x-codex-turn-state")) - require.Empty(t, upstream.requests[0].Header.Get("OpenAI-Beta")) - require.Empty(t, upstream.requests[0].Header.Get("originator")) + requireOpenAIMessagesCodexIdentity(t, upstream.requests[0], codexCLIUserAgent, "codex_cli_rs") secondBody := []byte(`{"model":"claude-sonnet-4-5","max_tokens":16,"messages":[{"role":"user","content":"first"},{"role":"assistant","content":"ok"},{"role":"user","content":"second"}],"stream":false}`) secondRec := httptest.NewRecorder() @@ -852,12 +851,73 @@ func TestForwardAsAnthropic_ReusesOAuthCodexTurnState(t *testing.T) { require.Equal(t, "turn_state_first", upstream.requests[1].Header.Get("x-codex-turn-state")) require.Equal(t, generateSessionUUID(isolateOpenAISessionID(0, "stable-cache-key")), upstream.requests[1].Header.Get("session_id")) require.Empty(t, upstream.requests[1].Header.Get("conversation_id")) - require.Empty(t, upstream.requests[1].Header.Get("OpenAI-Beta")) - require.Empty(t, upstream.requests[1].Header.Get("originator")) + requireOpenAIMessagesCodexIdentity(t, upstream.requests[1], codexCLIUserAgent, "codex_cli_rs") require.False(t, gjson.GetBytes(upstream.bodies[1], "prompt_cache_key").Exists()) require.False(t, gjson.GetBytes(upstream.bodies[1], "previous_response_id").Exists()) } +func TestForwardAsAnthropic_OAuthRestoresCodexIdentityHeaders(t *testing.T) { + gin.SetMode(gin.TestMode) + + const tuiUA = "codex-tui/9.9.9 (Mac OS X 14.0; arm64) iTerm (codex-tui; 9.9.9)" + tests := []struct { + name string + userAgent string + originator string + wantUserAgent string + wantOriginator string + }{ + { + name: "官方UA逐字保留并重新配对", + userAgent: tuiUA, + originator: "opencode", + wantUserAgent: tuiUA, + wantOriginator: "codex-tui", + }, + { + name: "第三方UA回退为默认Codex身份", + userAgent: "third-party-client/1.0.0", + originator: "opencode", + wantUserAgent: codexCLIUserAgent, + wantOriginator: "codex_cli_rs", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + body := []byte(`{"model":"claude-sonnet-4-5","max_tokens":16,"messages":[{"role":"user","content":"hello"}],"stream":false}`) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + c.Request.Header.Set("User-Agent", tt.userAgent) + c.Request.Header.Set("originator", tt.originator) + + upstream := &httpUpstreamRecorder{resp: openAICompatSSECompletedResponse("resp_identity", "gpt-5.4")} + svc := &OpenAIGatewayService{ + httpUpstream: upstream, + cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{Enabled: false}}}, + } + account := &Account{ + ID: 1, + Name: "openai-oauth", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "oauth-token", + "chatgpt_account_id": "chatgpt-acc", + }, + } + + result, err := svc.ForwardAsAnthropic(context.Background(), c, account, body, "", "gpt-5.4") + require.NoError(t, err) + require.NotNil(t, result) + requireOpenAIMessagesCodexIdentity(t, upstream.lastReq, tt.wantUserAgent, tt.wantOriginator) + }) + } +} + func TestForwardAsAnthropic_OAuthDigestFallbackReusesTurnStateWithoutExplicitKey(t *testing.T) { t.Parallel() gin.SetMode(gin.TestMode) @@ -896,6 +956,7 @@ func TestForwardAsAnthropic_OAuthDigestFallbackReusesTurnStateWithoutExplicitKey firstSessionID := upstream.requests[0].Header.Get("session_id") require.NotEmpty(t, firstSessionID) require.Empty(t, upstream.requests[0].Header.Get("x-codex-turn-state")) + requireOpenAIMessagesCodexIdentity(t, upstream.requests[0], codexCLIUserAgent, "codex_cli_rs") require.False(t, gjson.GetBytes(upstream.bodies[0], "prompt_cache_key").Exists()) secondBody := []byte(`{"model":"claude-sonnet-4-5","max_tokens":16,"messages":[{"role":"user","content":"first"},{"role":"assistant","content":"ok"},{"role":"user","content":"second"}],"stream":false}`) @@ -910,6 +971,7 @@ func TestForwardAsAnthropic_OAuthDigestFallbackReusesTurnStateWithoutExplicitKey require.Equal(t, firstSessionID, upstream.requests[1].Header.Get("session_id")) require.Equal(t, "turn_state_digest_first", upstream.requests[1].Header.Get("x-codex-turn-state")) require.Empty(t, upstream.requests[1].Header.Get("conversation_id")) + requireOpenAIMessagesCodexIdentity(t, upstream.requests[1], codexCLIUserAgent, "codex_cli_rs") require.False(t, gjson.GetBytes(upstream.bodies[1], "prompt_cache_key").Exists()) require.False(t, gjson.GetBytes(upstream.bodies[1], "previous_response_id").Exists()) } @@ -1064,8 +1126,7 @@ func TestForwardAsAnthropic_OAuthKeepsSystemAsDeveloperInput(t *testing.T) { instructions := gjson.GetBytes(upstream.lastBody, "instructions") require.True(t, instructions.Exists()) require.Empty(t, instructions.String()) - require.Empty(t, upstream.requests[0].Header.Get("OpenAI-Beta")) - require.Empty(t, upstream.requests[0].Header.Get("originator")) + requireOpenAIMessagesCodexIdentity(t, upstream.requests[0], codexCLIUserAgent, "codex_cli_rs") } func TestForwardAsAnthropic_OAuthAddsClaudeCodeTodoGuardForCompatModel(t *testing.T) { @@ -1202,6 +1263,15 @@ func openAICompatSSECompletedResponse(responseID, model string) *http.Response { } } +func requireOpenAIMessagesCodexIdentity(t *testing.T, req *http.Request, wantUserAgent, wantOriginator string) { + t.Helper() + require.NotNil(t, req) + require.Equal(t, wantUserAgent, req.Header.Get("User-Agent")) + require.Equal(t, wantOriginator, req.Header.Get("originator")) + require.Equal(t, codexCLIVersion, req.Header.Get("version")) + require.Equal(t, "responses=experimental", req.Header.Get("OpenAI-Beta")) +} + func openAICompatSSEResponseWithoutUsage(responseID, model string) *http.Response { body := strings.Join([]string{ `data: {"type":"response.completed","response":{"id":"` + responseID + `","object":"response","model":"` + model + `","status":"completed","output":[{"type":"message","id":"msg_1","role":"assistant","status":"completed","content":[{"type":"output_text","text":"ok"}]}]}}`, diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index 8dbd9ddad0..5952469409 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -878,6 +878,8 @@ func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) { c, _ := gin.CreateTestContext(recorder) body := []byte(`{"model":"grok","max_tokens":32,"stream":false,"messages":[{"role":"user","content":"hi"}]}`) c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body)) + c.Request.Header.Set("OpenAI-Beta", "grok-experimental") + c.Request.Header.Set("originator", "opencode") account := &Account{ ID: 54, @@ -908,6 +910,9 @@ func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) { require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String()) require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) require.Equal(t, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent")) + require.Equal(t, "grok-experimental", upstream.lastReq.Header.Get("OpenAI-Beta")) + require.Empty(t, upstream.lastReq.Header.Get("originator")) + require.Empty(t, upstream.lastReq.Header.Get("version")) require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String()) require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool()) require.NotContains(t, string(upstream.lastBody), "chatgpt.com") diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go index 219b5e4be4..de332f1404 100644 --- a/backend/internal/service/openai_gateway_messages.go +++ b/backend/internal/service/openai_gateway_messages.go @@ -261,9 +261,8 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( // 6. Build upstream request if account.Type == AccountTypeOAuth && account.Platform != PlatformGrok { // Messages 兼容桥即使 body 未带 todo-guard/prompt_cache_key 标记(如映射到非 - // gpt-5/codex 模型),也必须让 buildUpstreamRequest 走 bridge 分支:不带 - // originator、User-Agent 逐字透传,避免身份收口(issue #3901)误改本路径 - // 刻意最小化的请求形态(下方的 Del(OpenAI-Beta/originator) 兜底保持不变)。 + // gpt-5/codex 模型),也必须让 buildUpstreamRequest 走 bridge 分支,以保留 + // 既有 body/session/conversation 行为。身份头在 post-build 阶段统一恢复。 setOpenAICompatMessagesBridgeContext(c, true) } upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) @@ -288,12 +287,16 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( } } if account.Type == AccountTypeOAuth && account.Platform != PlatformGrok { - // Anthropic Messages compatibility uses the ChatGPT Codex SSE endpoint. - // Match airgate-openai's request shape: the SSE endpoint does not need - // the Responses experimental beta header, and forcing originator can make - // ChatGPT select a different internal continuation path. - upstreamReq.Header.Del("OpenAI-Beta") - upstreamReq.Header.Del("originator") + // buildUpstreamRequest 保留 Messages bridge 的 body/session 兼容行为,并会先 + // 清除身份头。真正发送前恢复完整 Codex 身份,避免 ChatGPT Codex 上游因缺失 + // originator/OpenAI-Beta 返回 404(issue #3901)。 + ensureCodexIdentityHeaders(upstreamReq.Header) + enforceCodexIdentityHeaders(upstreamReq.Header) + logger.L().Debug("openai messages: upstream identity restored", + zap.Int64("account_id", account.ID), + zap.String("upstream_model", upstreamModel), + zap.Bool("compat_identity_restored", true), + ) } if account.Type == AccountTypeOAuth && promptCacheKey != "" && strings.TrimSpace(c.GetHeader("conversation_id")) == "" { upstreamReq.Header.Del("conversation_id") diff --git a/backend/internal/service/openai_gateway_messages_chat_fallback_test.go b/backend/internal/service/openai_gateway_messages_chat_fallback_test.go index 8bb15c81aa..b36b76ffa8 100644 --- a/backend/internal/service/openai_gateway_messages_chat_fallback_test.go +++ b/backend/internal/service/openai_gateway_messages_chat_fallback_test.go @@ -356,6 +356,8 @@ func TestForwardAsAnthropic_ResponsesSupportedAccountStillUsesResponsesEndpoint( c, _ := gin.CreateTestContext(rec) c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body)) c.Request.Header.Set("Content-Type", "application/json") + c.Request.Header.Set("User-Agent", "third-party-client/1.0.0") + c.Request.Header.Set("originator", "opencode") upstreamBody := strings.Join([]string{ `data: {"type":"response.completed","response":{"id":"resp_native","object":"response","model":"gpt-5.4","status":"completed","output":[{"type":"message","id":"msg_1","role":"assistant","status":"completed","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":5,"output_tokens":2,"total_tokens":7}}}`, @@ -385,5 +387,9 @@ func TestForwardAsAnthropic_ResponsesSupportedAccountStillUsesResponsesEndpoint( "responses-capable account must stay on /v1/responses, got %s", upstream.lastReq.URL.String()) require.True(t, gjson.GetBytes(upstream.lastBody, "input").Exists()) require.False(t, gjson.GetBytes(upstream.lastBody, "messages").Exists()) + require.Equal(t, "third-party-client/1.0.0", upstream.lastReq.Header.Get("User-Agent")) + require.Equal(t, "opencode", upstream.lastReq.Header.Get("originator")) + require.Empty(t, upstream.lastReq.Header.Get("version")) + require.Empty(t, upstream.lastReq.Header.Get("OpenAI-Beta")) require.Equal(t, "ok", gjson.Get(rec.Body.String(), "content.0.text").String()) } From 5015b7a1c174583ce4b31b0deee85f576850146a Mon Sep 17 00:00:00 2001 From: jjaw Date: Sun, 12 Jul 2026 04:53:34 +0800 Subject: [PATCH 7/9] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=20tool=5Fsearch=20?= =?UTF-8?q?=E5=8F=82=E6=95=B0=E5=AF=B9=E8=B1=A1=E5=8F=8D=E5=BA=8F=E5=88=97?= =?UTF-8?q?=E5=8C=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../responses_stream_event_wire_test.go | 53 +++++++++++++++++++ backend/internal/pkg/apicompat/types.go | 50 +++++++++++++++++ 2 files changed, 103 insertions(+) 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 f44f3e7770..fb138a1469 100644 --- a/backend/internal/pkg/apicompat/responses_stream_event_wire_test.go +++ b/backend/internal/pkg/apicompat/responses_stream_event_wire_test.go @@ -131,3 +131,56 @@ func TestWire_UnknownEventFallsBackToDefault(t *testing.T) { }) require.Contains(t, m, "response") } + +func TestResponsesOutputUnmarshal_ToolSearchObjectArguments(t *testing.T) { + var item ResponsesOutput + require.NoError(t, json.Unmarshal([]byte(`{ + "type":"tool_search_call", + "id":"item_1", + "call_id":"call_1", + "execution":"client", + "arguments":{"query":"gmail","limit":2} + }`), &item)) + require.Equal(t, "tool_search_call", item.Type) + require.Equal(t, `{"query":"gmail","limit":2}`, item.Arguments) + + wire, err := json.Marshal(item) + require.NoError(t, err) + var decoded map[string]any + require.NoError(t, json.Unmarshal(wire, &decoded)) + args, ok := decoded["arguments"].(map[string]any) + require.True(t, ok, "tool_search_call arguments must remain an object") + require.Equal(t, "gmail", args["query"]) +} + +func TestResponsesResponseUnmarshal_ToolSearchObjectArguments(t *testing.T) { + var response ResponsesResponse + require.NoError(t, json.Unmarshal([]byte(`{ + "id":"response_1", + "object":"response", + "status":"completed", + "output":[{ + "type":"tool_search_call", + "id":"item_1", + "call_id":"call_1", + "arguments":{"query":"gmail"} + }] + }`), &response)) + require.Len(t, response.Output, 1) + require.Equal(t, `{"query":"gmail"}`, response.Output[0].Arguments) +} + +func TestResponsesStreamEventUnmarshal_ToolSearchObjectArguments(t *testing.T) { + var event ResponsesStreamEvent + require.NoError(t, json.Unmarshal([]byte(`{ + "type":"response.output_item.done", + "item":{ + "type":"tool_search_call", + "id":"item_1", + "call_id":"call_1", + "arguments":{"query":"gmail"} + } + }`), &event)) + require.NotNil(t, event.Item) + require.Equal(t, `{"query":"gmail"}`, event.Item.Arguments) +} diff --git a/backend/internal/pkg/apicompat/types.go b/backend/internal/pkg/apicompat/types.go index 9f3f2daa66..6cf9a2be31 100644 --- a/backend/internal/pkg/apicompat/types.go +++ b/backend/internal/pkg/apicompat/types.go @@ -353,6 +353,56 @@ func (o ResponsesOutput) MarshalJSON() ([]byte, error) { return json.Marshal(m) } +// UnmarshalJSON accepts both the Responses function-call string form and the +// tool_search_call object form for arguments. The bridge stores arguments as a +// string internally, so object arguments are retained as their raw JSON. +func (o *ResponsesOutput) UnmarshalJSON(data []byte) error { + type responsesOutputAlias ResponsesOutput + + var kind struct { + Type string `json:"type"` + } + if err := json.Unmarshal(data, &kind); err != nil { + return err + } + if kind.Type != "tool_search_call" { + var decoded responsesOutputAlias + if err := json.Unmarshal(data, &decoded); err != nil { + return err + } + *o = ResponsesOutput(decoded) + return nil + } + + var fields map[string]json.RawMessage + if err := json.Unmarshal(data, &fields); err != nil { + return err + } + arguments, hasArguments := fields["arguments"] + delete(fields, "arguments") + normalized, err := json.Marshal(fields) + if err != nil { + return err + } + + var decoded responsesOutputAlias + if err := json.Unmarshal(normalized, &decoded); err != nil { + return err + } + *o = ResponsesOutput(decoded) + if !hasArguments || string(arguments) == "null" { + return nil + } + + var argumentString string + if err := json.Unmarshal(arguments, &argumentString); err == nil { + o.Arguments = argumentString + } else { + o.Arguments = string(arguments) + } + return nil +} + // WebSearchAction describes the search action in a web_search_call output item. type WebSearchAction struct { Type string `json:"type,omitempty"` // "search" From 06af8115f7fda82c70075a675bb581a25c3ed4d7 Mon Sep 17 00:00:00 2001 From: jjaw Date: Sun, 12 Jul 2026 04:53:41 +0800 Subject: [PATCH 8/9] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=20compact=20=E5=BF=83?= =?UTF-8?q?=E8=B7=B3=20writer=20=E7=94=9F=E5=91=BD=E5=91=A8=E6=9C=9F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../handler/ops_capture_writer_nil_test.go | 31 +++++++++++++++++++ .../service/openai_compact_sse_keepalive.go | 15 +++++++-- 2 files changed, 43 insertions(+), 3 deletions(-) diff --git a/backend/internal/handler/ops_capture_writer_nil_test.go b/backend/internal/handler/ops_capture_writer_nil_test.go index 4e96333f9d..88aa7c043f 100644 --- a/backend/internal/handler/ops_capture_writer_nil_test.go +++ b/backend/internal/handler/ops_capture_writer_nil_test.go @@ -1,9 +1,15 @@ package handler import ( + "net/http" + "net/http/httptest" "testing" + "time" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestOpsCaptureWriter_NilInnerWriter_NoPanic(t *testing.T) { @@ -57,3 +63,28 @@ func TestOpsCaptureWriter_NilInnerWriter_NoPanic(t *testing.T) { assert.Nil(t, p) }) } + +func TestOpsCaptureWriter_CompactKeepaliveRestoresOriginalWriter(t *testing.T) { + gin.SetMode(gin.TestMode) + router := gin.New() + outerStatus := -1 + router.Use(func(c *gin.Context) { + c.Next() + outerStatus = c.Writer.Status() + }) + router.Use(OpsErrorLoggerMiddleware(nil)) + router.GET("/compact", func(c *gin.Context) { + service.MarkOpenAICompactClientStream(c) + stop := service.StartOpenAICompactSSEKeepalive(c, time.Hour) + defer stop() + c.Status(http.StatusOK) + }) + + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodGet, "/compact", nil) + require.NotPanics(t, func() { + router.ServeHTTP(recorder, request) + }) + require.Equal(t, http.StatusOK, outerStatus) + require.Equal(t, http.StatusOK, recorder.Code) +} diff --git a/backend/internal/service/openai_compact_sse_keepalive.go b/backend/internal/service/openai_compact_sse_keepalive.go index 70ef3fc01a..776fde6962 100644 --- a/backend/internal/service/openai_compact_sse_keepalive.go +++ b/backend/internal/service/openai_compact_sse_keepalive.go @@ -43,12 +43,14 @@ func StartOpenAICompactSSEKeepalive(c *gin.Context, interval time.Duration) func if c == nil || c.Writer == nil || interval <= 0 || !openAICompactClientWantsStream(c) { return func() {} } + originalWriter := c.Writer k := &openAICompactSSEKeepalive{ - writer: c.Writer, + writer: originalWriter, stop: make(chan struct{}), } c.Set(openAICompactSSEKeepaliveKey, k) - c.Writer = &openAICompactKeepaliveWriter{ResponseWriter: c.Writer, k: k} + wrappedWriter := &openAICompactKeepaliveWriter{ResponseWriter: originalWriter, k: k} + c.Writer = wrappedWriter var reqDone <-chan struct{} if c.Request != nil { @@ -71,7 +73,14 @@ func StartOpenAICompactSSEKeepalive(c *gin.Context, interval time.Duration) func timer.Reset(interval) } }() - return k.Stop + return func() { + k.Stop() + // Do not leave a pooled middleware writer reachable through the compact + // wrapper after the request finishes. + if current, ok := c.Writer.(*openAICompactKeepaliveWriter); ok && current == wrappedWriter { + c.Writer = originalWriter + } + } } // beat 在锁内提交(首次)响应头并写出一条 SSE 注释行;返回 false 表示心跳已 From 73ffd134301190ffd27c6b6ab5749a21d87be0df Mon Sep 17 00:00:00 2001 From: jjaw Date: Sun, 12 Jul 2026 05:12:20 +0800 Subject: [PATCH 9/9] =?UTF-8?q?=E5=85=B3=E8=81=94=E4=B8=8A=E6=B8=B8=20issu?= =?UTF-8?q?e=20#3818=20#3887=20#3961?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 相关修复 PR:#3989 #3994