Merge pull request #4832 from ListenCodes/fix/openai-apikey-responses-item-id-v2

fix(openai): sanitize API-key responses item IDs
This commit is contained in:
Wesley Liddick
2026-07-25 13:42:35 +08:00
committed by GitHub
4 changed files with 184 additions and 19 deletions
@@ -1494,25 +1494,9 @@ func filterCodexInputWithOptions(input []any, opts codexInputFilterOptions) []an
if !opts.PreserveReferences {
ensureCopy()
delete(newItem, "id")
} else if isCodexToolCallInputType(typ) {
// 续链模式下保留 id 以维持上下文引用,但 function_call 等
// call-input 类 item 的 id 必须以 "fc" 开头(上游校验
// "Expected an ID that begins with 'fc'")。item_* 形式的 id
// 来自客户端回放,需要删除。
// 注意:function_call_output 等 output 类的 id 无此约束,不动。
if id, ok := m["id"].(string); ok && id != "" && !strings.HasPrefix(id, "fc") {
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")
}
} else if id, ok := m["id"].(string); ok && shouldStripOpenAIResponsesInputItemID(typ, id) {
ensureCopy()
delete(newItem, "id")
}
filtered = append(filtered, newItem)
@@ -0,0 +1,90 @@
//go:build unit
package service
import (
"context"
"fmt"
"io"
"net/http"
"runtime"
"strings"
"testing"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
func TestOpenAIGatewayService_APIKeyPassthrough_StripsInvalidInputItemIDs(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(
`{"id":"resp_test","model":"gpt-5.6-sol","output":[],"usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}`,
)),
}}
svc := newOpenAIImageGenerationControlTestService(upstream)
c, _ := newOpenAIImageGenerationControlTestContext(true, "codex_cli_rs/0.144.1")
account := newOpenAIImageGenerationControlTestAccount()
account.Extra = map[string]any{"openai_passthrough": true}
body := []byte(`{
"model":"gpt-5.6-sol",
"stream":false,
"input":[
{"type":"message","id":"item_bad_message","role":"assistant","content":[{"type":"output_text","text":"hello"}]},
{"type":"function_call","id":"item_bad_call","call_id":"call_123","name":"exec_command","arguments":"{}"},
{"type":"message","id":"msg_valid","role":"user","content":[{"type":"input_text","text":"continue"}]},
{"type":"function_call","id":"fc_valid","call_id":"call_456","name":"apply_patch","arguments":"{}"},
{"type":"function_call_output","id":"item_output","call_id":"call_123","output":"done"},
{"type":"web_search_call","id":"item_unconstrained"}
]
}`)
result, err := svc.Forward(context.Background(), c, account, body)
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, upstream.lastReq)
forwarded := upstream.lastBody
require.False(t, gjson.GetBytes(forwarded, "input.0.id").Exists())
require.Equal(t, "hello", gjson.GetBytes(forwarded, "input.0.content.0.text").String())
require.False(t, gjson.GetBytes(forwarded, "input.1.id").Exists())
require.Equal(t, "call_123", gjson.GetBytes(forwarded, "input.1.call_id").String())
require.Equal(t, "exec_command", gjson.GetBytes(forwarded, "input.1.name").String())
require.Equal(t, "{}", gjson.GetBytes(forwarded, "input.1.arguments").String())
require.Equal(t, "msg_valid", gjson.GetBytes(forwarded, "input.2.id").String())
require.Equal(t, "fc_valid", gjson.GetBytes(forwarded, "input.3.id").String())
require.Equal(t, "item_output", gjson.GetBytes(forwarded, "input.4.id").String())
require.Equal(t, "call_123", gjson.GetBytes(forwarded, "input.4.call_id").String())
require.Equal(t, "item_unconstrained", gjson.GetBytes(forwarded, "input.5.id").String())
}
func TestSanitizeOpenAIResponsesInputItemIDs_AllocationGrowthIsLinear(t *testing.T) {
makeBody := func(itemCount int) []byte {
items := make([]string, itemCount)
for i := range items {
items[i] = fmt.Sprintf(`{"type":"message","id":"item_%d","role":"user","content":[{"type":"input_text","text":"hello"}]}`, i)
}
return []byte(`{"model":"gpt-5.6-sol","input":[` + strings.Join(items, ",") + `]}`)
}
allocatedBytes := func(body []byte) uint64 {
runtime.GC()
var before, after runtime.MemStats
runtime.ReadMemStats(&before)
sanitized, changed, err := sanitizeOpenAIResponsesInputItemIDs(body)
runtime.ReadMemStats(&after)
require.NoError(t, err)
require.True(t, changed)
require.NotEmpty(t, sanitized)
return after.TotalAlloc - before.TotalAlloc
}
smallAllocated := allocatedBytes(makeBody(20))
largeAllocated := allocatedBytes(makeBody(200))
require.Less(t, largeAllocated, smallAllocated*30,
"10x more input items must not cause quadratic whole-body allocation growth")
}
@@ -95,6 +95,19 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
if account.Type == AccountTypeAPIKey && !openai_compat.ShouldUseResponsesAPI(account.Extra) {
return s.forwardResponsesViaRawChatCompletions(ctx, c, account, body)
}
if account.Platform == PlatformOpenAI && account.Type == AccountTypeAPIKey {
sanitizedBody, changed, sanitizeErr := sanitizeOpenAIResponsesInputItemIDs(body)
if sanitizeErr != nil {
return nil, fmt.Errorf("sanitize OpenAI Responses input item IDs: %w", sanitizeErr)
}
if changed {
body = sanitizedBody
originalBody = sanitizedBody
requestView = newOpenAIRequestView(sanitizedBody)
reqModel, reqStream, promptCacheKey = requestView.Model, requestView.Stream, requestView.PromptCacheKey
originalModel = reqModel
}
}
compatMessagesBridge := isOpenAICompatMessagesBridgeBody(body)
setOpenAICompatMessagesBridgeContext(c, compatMessagesBridge)
@@ -0,0 +1,78 @@
package service
import (
"fmt"
"strings"
"github.com/tidwall/gjson"
"github.com/tidwall/sjson"
)
// Invalid replayed IDs are removed rather than rewritten because a fabricated
// msg/fc ID may point at a different upstream object.
func shouldStripOpenAIResponsesInputItemID(itemType, id string) bool {
if id == "" {
return false
}
if itemType == "message" {
return !strings.HasPrefix(id, "msg")
}
if isCodexToolCallInputType(itemType) {
return !strings.HasPrefix(id, "fc")
}
return false
}
func sanitizeOpenAIResponsesInputItemIDs(body []byte) ([]byte, bool, error) {
input := gjson.GetBytes(body, "input")
if !input.IsArray() {
return body, false, nil
}
items := make([][]byte, 0)
changed := false
var sanitizeErr error
index := 0
input.ForEach(func(_, item gjson.Result) bool {
currentIndex := index
index++
itemBody := []byte(item.Raw)
if item.IsObject() {
itemType := item.Get("type")
id := item.Get("id")
if itemType.Type == gjson.String && id.Type == gjson.String &&
shouldStripOpenAIResponsesInputItemID(itemType.String(), id.String()) {
itemBody, sanitizeErr = sjson.DeleteBytes(itemBody, "id")
if sanitizeErr != nil {
sanitizeErr = fmt.Errorf("delete input.%d.id: %w", currentIndex, sanitizeErr)
return false
}
changed = true
}
}
items = append(items, itemBody)
return true
})
if sanitizeErr != nil {
return nil, false, sanitizeErr
}
if !changed {
return body, false, nil
}
rebuiltInput := make([]byte, 0, len(input.Raw))
rebuiltInput = append(rebuiltInput, '[')
for i, item := range items {
if i > 0 {
rebuiltInput = append(rebuiltInput, ',')
}
rebuiltInput = append(rebuiltInput, item...)
}
rebuiltInput = append(rebuiltInput, ']')
sanitized, err := sjson.SetRawBytes(body, "input", rebuiltInput)
if err != nil {
return nil, false, fmt.Errorf("replace sanitized input: %w", err)
}
return sanitized, true, nil
}