mirror of
https://github.com/router-for-me/CLIProxyAPI.git
synced 2026-09-30 04:40:37 +08:00
fix(openai): preserve prewarm input and allow named tool outputs in responses websocket
- Track pending synthetic prewarm response IDs to merge warmup inputs into subsequent delta followups. - Normalize transcript replacements when followups do not reference the prewarm parent response ID. - Validate that the `input` field is an array for `response.create` requests. - Allow `function_call_output` items without a `call_id` when a non-empty tool name is present. Closes: #5631
This commit is contained in:
@@ -3,6 +3,7 @@ package openai
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
@@ -328,6 +329,8 @@ func (h *OpenAIResponsesAPIHandler) ResponsesWebsocket(c *gin.Context) {
|
||||
var lastRequest []byte
|
||||
lastResponseOutput := []byte("[]")
|
||||
lastResponseID := ""
|
||||
// Remains pending until a generating request commits successfully.
|
||||
pendingPrewarmID := ""
|
||||
var lastResponsePendingToolCallIDs []string
|
||||
pinnedAuthID := ""
|
||||
// Preserve independent upstream auth affinity when a downstream session switches providers.
|
||||
@@ -495,7 +498,25 @@ func (h *OpenAIResponsesAPIHandler) ResponsesWebsocket(c *gin.Context) {
|
||||
var requestJSON []byte
|
||||
var updatedLastRequest []byte
|
||||
var errMsg *interfaces.ErrorMessage
|
||||
if nativeWebsocketPassthrough {
|
||||
previousResponseID := strings.TrimSpace(gjson.GetBytes(payload, "previous_response_id").String())
|
||||
if pendingPrewarmID != "" && previousResponseID != "" {
|
||||
if previousResponseID != pendingPrewarmID {
|
||||
errMsg = responsesWebsocketPreviousResponseNotFoundError()
|
||||
} else {
|
||||
requestJSON, updatedLastRequest, errMsg = normalizeResponsesWebsocketPrewarmFollowup(payload, lastRequest)
|
||||
}
|
||||
} else if pendingPrewarmID != "" && gjson.GetBytes(payload, "type").String() == wsRequestTypeCreate {
|
||||
input := gjson.GetBytes(payload, "input")
|
||||
if input.Exists() && !input.IsArray() {
|
||||
errMsg = &interfaces.ErrorMessage{
|
||||
StatusCode: http.StatusBadRequest,
|
||||
Error: fmt.Errorf("websocket request requires array field: input"),
|
||||
}
|
||||
} else {
|
||||
// No parent reference means a self-contained replacement, not a delta.
|
||||
requestJSON, updatedLastRequest, errMsg = normalizeResponseCreateRequest(normalizeResponseTranscriptReplacement(payload, lastRequest))
|
||||
}
|
||||
} else if nativeWebsocketPassthrough {
|
||||
requestJSON, errMsg = normalizeResponsesWebsocketPassthroughRequest(payload, requestModelName)
|
||||
} else if len(lastRequest) == 0 && strings.TrimSpace(gjson.GetBytes(payload, "previous_response_id").String()) != "" {
|
||||
errMsg = responsesWebsocketPreviousResponseNotFoundError()
|
||||
@@ -547,10 +568,12 @@ func (h *OpenAIResponsesAPIHandler) ResponsesWebsocket(c *gin.Context) {
|
||||
lastResponseOutput = []byte("[]")
|
||||
lastResponseID = ""
|
||||
lastResponsePendingToolCallIDs = nil
|
||||
if errWrite := writeResponsesWebsocketSyntheticPrewarm(c, writer, requestJSON, wsTimelineLog, passthroughSessionID); errWrite != nil {
|
||||
prewarmID, errWrite := writeResponsesWebsocketSyntheticPrewarm(c, writer, requestJSON, wsTimelineLog, passthroughSessionID)
|
||||
if errWrite != nil {
|
||||
wsTerminateErr = errWrite
|
||||
return
|
||||
}
|
||||
pendingPrewarmID = prewarmID
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -651,6 +674,7 @@ func (h *OpenAIResponsesAPIHandler) ResponsesWebsocket(c *gin.Context) {
|
||||
}
|
||||
|
||||
toolCacheTurn.commit()
|
||||
pendingPrewarmID = ""
|
||||
upstreamMode = attemptedUpstreamMode
|
||||
if upstreamMode == responsesWebsocketUpstreamModeWS {
|
||||
upstreamWebsocketAuthID = lastAttemptedAuthID
|
||||
|
||||
@@ -1,11 +1,14 @@
|
||||
package openai
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/google/uuid"
|
||||
"github.com/router-for-me/CLIProxyAPI/v7/internal/interfaces"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/tidwall/gjson"
|
||||
"github.com/tidwall/sjson"
|
||||
@@ -28,10 +31,10 @@ func writeResponsesWebsocketSyntheticPrewarm(
|
||||
requestJSON []byte,
|
||||
wsTimelineLog websocketTimelineAppender,
|
||||
sessionID string,
|
||||
) error {
|
||||
) (string, error) {
|
||||
payloads, errPayloads := syntheticResponsesWebsocketPrewarmPayloads(requestJSON)
|
||||
if errPayloads != nil {
|
||||
return errPayloads
|
||||
return "", errPayloads
|
||||
}
|
||||
for i := 0; i < len(payloads); i++ {
|
||||
markAPIResponseTimestamp(c)
|
||||
@@ -49,10 +52,35 @@ func writeResponsesWebsocketSyntheticPrewarm(
|
||||
websocketPayloadEventType(payloads[i]),
|
||||
errWrite,
|
||||
)
|
||||
return errWrite
|
||||
return "", errWrite
|
||||
}
|
||||
}
|
||||
return nil
|
||||
return gjson.GetBytes(payloads[0], "response.id").String(), nil
|
||||
}
|
||||
|
||||
// A synthetic warm-up acknowledges input that never reached the upstream.
|
||||
// Materialize that input before compacted-history detection can mistake the
|
||||
// client's remaining delta for a complete replacement transcript.
|
||||
func normalizeResponsesWebsocketPrewarmFollowup(rawJSON, warmupRequest []byte) ([]byte, []byte, *interfaces.ErrorMessage) {
|
||||
requestType := strings.TrimSpace(gjson.GetBytes(rawJSON, "type").String())
|
||||
if requestType != wsRequestTypeCreate && requestType != wsRequestTypeAppend {
|
||||
return nil, warmupRequest, &interfaces.ErrorMessage{StatusCode: http.StatusBadRequest, Error: fmt.Errorf("unsupported websocket request type: %s", requestType)}
|
||||
}
|
||||
input := gjson.GetBytes(rawJSON, "input")
|
||||
if !input.IsArray() {
|
||||
return nil, warmupRequest, &interfaces.ErrorMessage{StatusCode: http.StatusBadRequest, Error: fmt.Errorf("websocket request requires array field: input")}
|
||||
}
|
||||
merged, errMerge := mergeResponsesWebsocketInput(warmupRequest, []byte("[]"), input.Raw)
|
||||
if errMerge != nil {
|
||||
return nil, warmupRequest, &interfaces.ErrorMessage{StatusCode: http.StatusBadRequest, Error: errMerge}
|
||||
}
|
||||
normalized := normalizeResponseTranscriptReplacement(rawJSON, warmupRequest)
|
||||
var errSet error
|
||||
normalized, errSet = sjson.SetRawBytes(normalized, "input", merged)
|
||||
if errSet != nil {
|
||||
return nil, warmupRequest, &interfaces.ErrorMessage{StatusCode: http.StatusBadRequest, Error: errSet}
|
||||
}
|
||||
return normalized, normalized, nil
|
||||
}
|
||||
|
||||
func syntheticResponsesWebsocketPrewarmPayloads(requestJSON []byte) ([][]byte, error) {
|
||||
|
||||
@@ -48,6 +48,14 @@ func normalizeResponsesWebsocketRequestWithIncrementalState(rawJSON []byte, last
|
||||
}
|
||||
|
||||
func normalizeResponseCreateRequest(rawJSON []byte) ([]byte, []byte, *interfaces.ErrorMessage) {
|
||||
input := gjson.GetBytes(rawJSON, "input")
|
||||
if input.Exists() && !input.IsArray() {
|
||||
return nil, nil, &interfaces.ErrorMessage{
|
||||
StatusCode: http.StatusBadRequest,
|
||||
Error: fmt.Errorf("websocket request requires array field: input"),
|
||||
}
|
||||
}
|
||||
|
||||
normalized, errDelete := sjson.DeleteBytes(rawJSON, "type")
|
||||
if errDelete != nil {
|
||||
normalized = bytes.Clone(rawJSON)
|
||||
|
||||
@@ -600,6 +600,23 @@ type websocketProviderCaptureExecutor struct {
|
||||
websocketCaptureExecutor
|
||||
}
|
||||
|
||||
type websocketPrewarmRetryExecutor struct {
|
||||
websocketProviderCaptureExecutor
|
||||
failFirst bool
|
||||
}
|
||||
|
||||
func (e *websocketPrewarmRetryExecutor) ExecuteStream(ctx context.Context, auth *coreauth.Auth, req coreexecutor.Request, opts coreexecutor.Options) (*coreexecutor.StreamResult, error) {
|
||||
if e.failFirst && e.streamCalls == 0 {
|
||||
e.streamCalls++
|
||||
e.payloads = append(e.payloads, bytes.Clone(req.Payload))
|
||||
chunks := make(chan coreexecutor.StreamChunk, 1)
|
||||
chunks <- coreexecutor.StreamChunk{Err: websocketPinnedFailoverStatusError{status: http.StatusBadRequest, msg: "retry diagnostic"}}
|
||||
close(chunks)
|
||||
return &coreexecutor.StreamResult{Chunks: chunks}, nil
|
||||
}
|
||||
return e.websocketCaptureExecutor.ExecuteStream(ctx, auth, req, opts)
|
||||
}
|
||||
|
||||
type websocketProviderRouteHost struct{}
|
||||
|
||||
func (*websocketProviderRouteHost) HasModelRouters() bool { return true }
|
||||
@@ -4360,6 +4377,171 @@ func TestWebsocketUpstreamSupportsCompactionReplayForModelFalseWhenMixedBackends
|
||||
}
|
||||
}
|
||||
|
||||
func TestResponsesWebsocketPrewarmPreservesCompactedFollowup(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
for _, tc := range []struct {
|
||||
name, input string
|
||||
parent, failFirst, wrongParent, invalidFirst, invalidType, omitModel bool
|
||||
wantPrefix bool
|
||||
}{
|
||||
{name: "compacted_delta", input: `[{"type":"compaction","encrypted_content":"opaque-checkpoint"},{"type":"function_call_output","name":"automation_update","output":"current heartbeat"}]`, parent: true, wantPrefix: true},
|
||||
{name: "ordinary_delta", input: `[{"type":"function_call_output","name":"automation_update","output":"current heartbeat"}]`, parent: true, wantPrefix: true},
|
||||
{name: "failed_attempt_reconnect", input: `[{"type":"compaction","encrypted_content":"opaque-checkpoint"},{"type":"function_call_output","name":"automation_update","output":"current heartbeat"}]`, parent: true, failFirst: true, wantPrefix: true},
|
||||
{name: "invalid_delta_retry", input: `[{"type":"compaction","encrypted_content":"opaque-checkpoint"},{"type":"function_call_output","name":"automation_update","output":"current heartbeat"}]`, parent: true, invalidFirst: true, wantPrefix: true},
|
||||
{name: "invalid_type_retry", input: `[{"type":"function_call_output","name":"automation_update","output":"current heartbeat"}]`, parent: true, invalidType: true, wantPrefix: true},
|
||||
{name: "replacement_inherits_defaults", input: `[{"type":"additional_tools","role":"developer","tools":[]}]`, omitModel: true},
|
||||
{name: "invalid_replacement_retry", input: `[{"type":"additional_tools","role":"developer","tools":[]}]`, invalidFirst: true},
|
||||
{name: "replacement_empty_tools", input: `[{"type":"additional_tools","role":"developer","tools":[]},{"type":"message","role":"user","content":"replacement"}]`},
|
||||
{name: "replacement_new_tools", input: `[{"type":"additional_tools","role":"developer","tools":[{"type":"function","name":"replacement_tool"}]},{"type":"message","role":"user","content":"replacement"}]`},
|
||||
{name: "unrelated_parent", input: `[{"type":"message","role":"user","content":"not this warmup"}]`, parent: true, wrongParent: true},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
executor := &websocketPrewarmRetryExecutor{websocketProviderCaptureExecutor: websocketProviderCaptureExecutor{provider: "codex"}, failFirst: tc.failFirst}
|
||||
manager := coreauth.NewManager(nil, nil, nil)
|
||||
manager.RegisterExecutor(executor)
|
||||
auth := &coreauth.Auth{ID: "prewarm-prefix-" + tc.name, Provider: "codex", Status: coreauth.StatusActive, Attributes: map[string]string{"websockets": "false"}}
|
||||
if _, errRegister := manager.Register(context.Background(), auth); errRegister != nil {
|
||||
t.Fatal(errRegister)
|
||||
}
|
||||
registry.GetGlobalRegistry().RegisterClient(auth.ID, auth.Provider, []*registry.ModelInfo{{ID: "prewarm-prefix-model"}})
|
||||
t.Cleanup(func() { registry.GetGlobalRegistry().UnregisterClient(auth.ID) })
|
||||
h := NewOpenAIResponsesAPIHandler(handlers.NewBaseAPIHandlers(&sdkconfig.SDKConfig{}, manager))
|
||||
router := gin.New()
|
||||
router.GET("/v1/responses/ws", h.ResponsesWebsocket)
|
||||
server := httptest.NewServer(router)
|
||||
defer server.Close()
|
||||
conn, _, errDial := websocket.DefaultDialer.Dial("ws"+strings.TrimPrefix(server.URL, "http")+"/v1/responses/ws", http.Header{"Session_id": []string{auth.ID}})
|
||||
if errDial != nil {
|
||||
t.Fatal(errDial)
|
||||
}
|
||||
defer func() {
|
||||
_ = conn.Close()
|
||||
}()
|
||||
send := func(raw string) {
|
||||
t.Helper()
|
||||
if errSend := conn.WriteMessage(websocket.TextMessage, []byte(raw)); errSend != nil {
|
||||
t.Fatal(errSend)
|
||||
}
|
||||
}
|
||||
read := func() []byte {
|
||||
t.Helper()
|
||||
_, b, errRead := conn.ReadMessage()
|
||||
if errRead != nil {
|
||||
t.Fatal(errRead)
|
||||
}
|
||||
return b
|
||||
}
|
||||
warmup := `{"type":"response.create","model":"prewarm-prefix-model","instructions":"legacy base","generate":false,"input":[{"type":"additional_tools","id":"warm-tools","role":"developer","tools":[{"type":"namespace","name":"functions","tools":[{"type":"custom","name":"exec"}]}]},{"type":"message","id":"warm-base","role":"developer","content":"base instructions"}]}`
|
||||
send(warmup)
|
||||
created := read()
|
||||
parent := gjson.GetBytes(created, "response.id").String()
|
||||
if !strings.HasPrefix(parent, "resp_prewarm_") {
|
||||
t.Fatalf("expected synthetic warmup: %s", created)
|
||||
}
|
||||
if got := gjson.GetBytes(read(), "type").String(); got != wsEventTypeCompleted {
|
||||
t.Fatalf("warmup event=%s", got)
|
||||
}
|
||||
if executor.streamCalls != 0 {
|
||||
t.Fatal("warmup reached upstream")
|
||||
}
|
||||
parentField := ""
|
||||
if tc.parent {
|
||||
if tc.wrongParent {
|
||||
parent = "resp_prewarm_unrelated"
|
||||
}
|
||||
parentField = fmt.Sprintf(`,"previous_response_id":%q`, parent)
|
||||
}
|
||||
followup := fmt.Sprintf(`{"type":"response.create","model":"prewarm-prefix-model"%s,"input":%s,"client_metadata":{"source":"automation_heartbeat","keep":"unchanged"}}`, parentField, tc.input)
|
||||
if tc.omitModel {
|
||||
followup = strings.Replace(followup, `,"model":"prewarm-prefix-model"`, "", 1)
|
||||
}
|
||||
if tc.invalidType {
|
||||
send(fmt.Sprintf(`{"type":"unsupported","previous_response_id":%q,"input":[]}`, parent))
|
||||
if gjson.GetBytes(read(), "type").String() != "error" || executor.streamCalls != 0 {
|
||||
t.Fatal("invalid request type reached upstream")
|
||||
}
|
||||
}
|
||||
if tc.invalidFirst {
|
||||
if tc.parent {
|
||||
send(fmt.Sprintf(`{"type":"response.create","previous_response_id":%q,"input":{}}`, parent))
|
||||
} else {
|
||||
send(`{"type":"response.create","input":{}}`)
|
||||
}
|
||||
if gjson.GetBytes(read(), "type").String() != "error" || executor.streamCalls != 0 {
|
||||
t.Fatal("invalid delta did not fail before upstream")
|
||||
}
|
||||
}
|
||||
send(followup)
|
||||
result := read()
|
||||
if tc.wrongParent {
|
||||
if gjson.GetBytes(result, "type").String() != "error" || executor.streamCalls != 0 {
|
||||
t.Fatalf("unrelated parent inherited warmup: %s", result)
|
||||
}
|
||||
return
|
||||
}
|
||||
if tc.failFirst {
|
||||
if gjson.GetBytes(result, "type").String() != "error" {
|
||||
t.Fatalf("expected first failure: %s", result)
|
||||
}
|
||||
// Terminal upstream errors close the connection. Codex reconnects
|
||||
// and establishes a new warm-up before retrying the full request.
|
||||
_ = conn.Close()
|
||||
var errReconnect error
|
||||
conn, _, errReconnect = websocket.DefaultDialer.Dial("ws"+strings.TrimPrefix(server.URL, "http")+"/v1/responses/ws", http.Header{"Session_id": []string{auth.ID}})
|
||||
if errReconnect != nil {
|
||||
t.Fatal(errReconnect)
|
||||
}
|
||||
send(warmup)
|
||||
newParent := gjson.GetBytes(read(), "response.id").String()
|
||||
if gjson.GetBytes(read(), "type").String() != wsEventTypeCompleted {
|
||||
t.Fatal("retry warmup failed")
|
||||
}
|
||||
send(strings.ReplaceAll(followup, parent, newParent))
|
||||
result = read()
|
||||
}
|
||||
if gjson.GetBytes(result, "type").String() != wsEventTypeCompleted {
|
||||
t.Fatalf("followup failed: %s", result)
|
||||
}
|
||||
for _, forwarded := range executor.payloads {
|
||||
if gjson.GetBytes(forwarded, "model").String() != "prewarm-prefix-model" || gjson.GetBytes(forwarded, "instructions").String() != "legacy base" {
|
||||
t.Fatalf("request defaults lost: %s", forwarded)
|
||||
}
|
||||
input := gjson.GetBytes(forwarded, "input").Array()
|
||||
want := gjson.Parse(tc.input).Array()
|
||||
if tc.wantPrefix {
|
||||
if len(input) != len(want)+2 || input[0].Get("id").String() != "warm-tools" || input[1].Get("id").String() != "warm-base" {
|
||||
t.Fatalf("acknowledged tools/base prefix lost or duplicated: %s", forwarded)
|
||||
}
|
||||
input = input[2:]
|
||||
} else if len(input) != len(want) {
|
||||
t.Fatalf("stale prefix inherited by replacement: %s", forwarded)
|
||||
}
|
||||
for i := range want {
|
||||
if input[i].Raw != want[i].Raw {
|
||||
t.Fatalf("delta changed: got %s want %s", input[i].Raw, want[i].Raw)
|
||||
}
|
||||
}
|
||||
if gjson.GetBytes(forwarded, "previous_response_id").Exists() || gjson.GetBytes(forwarded, "generate").Exists() {
|
||||
t.Fatalf("synthetic state leaked upstream: %s", forwarded)
|
||||
}
|
||||
if gjson.GetBytes(forwarded, "client_metadata.keep").String() != "unchanged" {
|
||||
t.Fatal("metadata changed")
|
||||
}
|
||||
}
|
||||
if tc.name == "compacted_delta" {
|
||||
send(`{"type":"response.create","model":"prewarm-prefix-model","previous_response_id":"resp-upstream","input":[{"type":"function_call_output","name":"automation_update","output":"next heartbeat"}]}`)
|
||||
if gjson.GetBytes(read(), "type").String() != wsEventTypeCompleted {
|
||||
t.Fatal("real-response continuation failed")
|
||||
}
|
||||
last := executor.payloads[len(executor.payloads)-1]
|
||||
if len(gjson.GetBytes(last, `input.#(type=="additional_tools")#`).Array()) != 1 || gjson.GetBytes(last, "input").Array()[len(gjson.GetBytes(last, "input").Array())-1].Get("output").String() != "next heartbeat" {
|
||||
t.Fatalf("continuation lost or duplicated state: %s", last)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestResponsesWebsocketPrewarmHandledLocallyForSSEUpstream(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
|
||||
@@ -532,7 +532,12 @@ func repairResponsesToolCallItems(
|
||||
for _, item := range items {
|
||||
if isResponsesToolCallOutputType(item.itemType) {
|
||||
if item.callID == "" {
|
||||
// Upstream rejects tool outputs without a call_id; drop it.
|
||||
// Codex sends standalone named results for heartbeat and delegation
|
||||
// input. These intentionally have no preceding call or call_id.
|
||||
name := gjson.GetBytes(item.raw, "name")
|
||||
if item.itemType == "function_call_output" && name.Type == gjson.String && strings.TrimSpace(name.String()) != "" {
|
||||
filtered = append(filtered, item)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user