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:
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user