Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions internal/runtime/executor/codex_websockets_connection.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package executor

import (
"bytes"
"context"
"errors"
"fmt"
Expand All @@ -17,6 +18,7 @@ import (
cliproxyexecutor "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/executor"
"github.com/router-for-me/CLIProxyAPI/v7/sdk/proxyutil"
log "github.com/sirupsen/logrus"
"github.com/tidwall/gjson"
"github.com/tidwall/sjson"
"golang.org/x/net/proxy"
)
Expand Down Expand Up @@ -135,6 +137,10 @@ func buildCodexWebsocketRequestBody(body []byte) []byte {
return nil
}

if isResponsesSteerPayload(body) {
return bytes.Clone(body)
}

// Match codex-rs websocket v2 semantics: every request is `response.create`.
// Incremental follow-up turns continue on the same websocket using
// `previous_response_id` + incremental `input`, not `response.append`.
Expand All @@ -146,6 +152,10 @@ func buildCodexWebsocketRequestBody(body []byte) []byte {
return body
}

func isResponsesSteerPayload(body []byte) bool {
return strings.TrimSpace(gjson.GetBytes(body, "type").String()) == "response.steer"
}

func readCodexWebsocketMessage(ctx context.Context, sess *codexWebsocketSession, conn *websocket.Conn, readCh chan codexWebsocketRead) (int, []byte, error) {
if sess == nil {
if conn == nil {
Expand Down
16 changes: 16 additions & 0 deletions internal/runtime/executor/codex_websockets_executor_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,22 @@ func TestBuildCodexWebsocketRequestBodyPreservesPreviousResponseID(t *testing.T)
}
}

func TestBuildCodexWebsocketRequestBodyPreservesSteer(t *testing.T) {
body := []byte(`{"type":"response.steer","previous_response_id":"resp-1","input":"Keep the scope small."}`)

wsReqBody := buildCodexWebsocketRequestBody(body)

if got := gjson.GetBytes(wsReqBody, "type").String(); got != "response.steer" {
t.Fatalf("type = %s, want response.steer", got)
}
if gjson.GetBytes(wsReqBody, "model").Exists() {
t.Fatalf("steer body must not gain a model field: %s", wsReqBody)
}
if gjson.GetBytes(wsReqBody, "stream").Exists() {
t.Fatalf("steer body must not gain a stream field: %s", wsReqBody)
}
}

func BenchmarkBuildCodexWebsocketRequestBodyLargePayload(b *testing.B) {
body := []byte(`{"model":"gpt-5.6","input":[{"type":"message","id":"msg_1","role":"user","content":"` + strings.Repeat("x", 8<<20) + `"}]}`)
b.ReportAllocs()
Expand Down
85 changes: 53 additions & 32 deletions internal/runtime/executor/codex_websockets_stream.go
Original file line number Diff line number Diff line change
Expand Up @@ -44,29 +44,40 @@ func (e *CodexWebsocketsExecutor) ExecuteStream(ctx context.Context, auth *clipr
originalPayloadSource = opts.OriginalRequest
}
originalPayload := originalPayloadSource
originalTranslated, body := translateCodexRequestPair(from, to, baseModel, originalPayload, req.Payload, true)

body, err = helps.ApplyRequestThinking(body, req, opts, from.String(), to.String(), e.Identifier())
if err != nil {
return nil, err
}
steerRequest := isResponsesSteerPayload(originalPayload) || isResponsesSteerPayload(req.Payload)

var body []byte
var originalTranslated []byte
var multiAgentV2Conflict bool
var optimizeMultiAgentV2 bool
var replayScope codexReasoningReplayScope
if !steerRequest {
originalTranslated, body = translateCodexRequestPair(from, to, baseModel, originalPayload, req.Payload, true)

body, err = helps.ApplyRequestThinking(body, req, opts, from.String(), to.String(), e.Identifier())
if err != nil {
return nil, err
}

requestedModel := helps.PayloadRequestedModel(opts, req.Model)
requestPath := helps.PayloadRequestPath(opts)
body = helps.ApplyPayloadConfigWithRequest(e.cfg, baseModel, to.String(), from.String(), "", body, originalTranslated, requestedModel, requestPath, opts.Headers)
body = helps.SetStringIfDifferent(body, "model", baseModel)
body = normalizeCodexInstructions(body)
if e.cfg == nil || e.cfg.DisableImageGeneration == config.DisableImageGenerationOff {
body = ensureImageGenerationTool(body, baseModel, auth, opts.Headers)
}
body = sanitizeOpenAIResponsesReasoningEncryptedContent(ctx, "codex websockets executor", body)
body = normalizeCodexWebsocketParallelToolCalls(body, opts.Headers)
body = helps.NormalizeCodexToolSchemas(body)
multiAgentV2Conflict := helps.HasCodexMultiAgentV2NamespaceConflict(body)
body, optimizeMultiAgentV2 := helps.OptimizeCodexMultiAgentV2RequestForAuth(ctx, opts.Headers, body, e.cfg, auth, baseModel)
body, replayScope, errReplay := applyCodexReasoningReplayCacheRequired(ctx, from, req, opts, body)
if errReplay != nil {
return nil, errReplay
requestedModel := helps.PayloadRequestedModel(opts, req.Model)
requestPath := helps.PayloadRequestPath(opts)
body = helps.ApplyPayloadConfigWithRequest(e.cfg, baseModel, to.String(), from.String(), "", body, originalTranslated, requestedModel, requestPath, opts.Headers)
body = helps.SetStringIfDifferent(body, "model", baseModel)
body = normalizeCodexInstructions(body)
if e.cfg == nil || e.cfg.DisableImageGeneration == config.DisableImageGenerationOff {
body = ensureImageGenerationTool(body, baseModel, auth, opts.Headers)
}
body = sanitizeOpenAIResponsesReasoningEncryptedContent(ctx, "codex websockets executor", body)
body = normalizeCodexWebsocketParallelToolCalls(body, opts.Headers)
body = helps.NormalizeCodexToolSchemas(body)
multiAgentV2Conflict = helps.HasCodexMultiAgentV2NamespaceConflict(body)
body, optimizeMultiAgentV2 = helps.OptimizeCodexMultiAgentV2RequestForAuth(ctx, opts.Headers, body, e.cfg, auth, baseModel)
body, replayScope, err = applyCodexReasoningReplayCacheRequired(ctx, from, req, opts, body)
if err != nil {
return nil, err
}
} else {
body = bytes.Clone(req.Payload)
}

httpURL := strings.TrimSuffix(baseURL, "/") + "/responses"
Expand All @@ -75,17 +86,27 @@ func (e *CodexWebsocketsExecutor) ExecuteStream(ctx context.Context, auth *clipr
return nil, err
}

body, wsHeaders, errPromptCache := applyCodexPromptCacheHeadersWithContext(ctx, from, req, body, opts.Headers)
if errPromptCache != nil {
return nil, errPromptCache
}
clientBody := body
var wsHeaders http.Header
var identityState codexIdentityConfuseState
upstreamBody, identityState := applyCodexIdentityConfuseBody(e.cfg, auth, originalPayloadSource, body)
reporter.SetTranslatedReasoningEffort(clientBody, to.String())
wsHeaders = applyCodexWebsocketHeaders(ctx, wsHeaders, auth, apiKey, e.cfg, opts.Headers)
applyModelHeaderOverrides(wsHeaders, baseModel)
applyCodexIdentityConfuseHeaders(wsHeaders, &identityState)
var upstreamBody []byte
clientBody := body
if !steerRequest {
var errPromptCache error
body, wsHeaders, errPromptCache = applyCodexPromptCacheHeadersWithContext(ctx, from, req, body, opts.Headers)
if errPromptCache != nil {
return nil, errPromptCache
}
clientBody = body
upstreamBody, identityState = applyCodexIdentityConfuseBody(e.cfg, auth, originalPayloadSource, body)
reporter.SetTranslatedReasoningEffort(clientBody, to.String())
wsHeaders = applyCodexWebsocketHeaders(ctx, wsHeaders, auth, apiKey, e.cfg, opts.Headers)
applyModelHeaderOverrides(wsHeaders, baseModel)
applyCodexIdentityConfuseHeaders(wsHeaders, &identityState)
} else {
wsHeaders = applyCodexWebsocketHeaders(ctx, http.Header{}, auth, apiKey, e.cfg, opts.Headers)
applyModelHeaderOverrides(wsHeaders, baseModel)
upstreamBody = body
}

var authID, authLabel, authType, authValue string
authID = auth.ID
Expand Down
21 changes: 17 additions & 4 deletions internal/runtime/executor/xai_websockets_executor.go
Original file line number Diff line number Diff line change
Expand Up @@ -470,9 +470,19 @@ func (e *XAIWebsocketsExecutor) ExecuteStream(ctx context.Context, auth *cliprox
baseURL = xaiauth.DefaultAPIBaseURL
}

prepared, err := e.prepareResponsesWebsocketRequest(ctx, req, opts)
if err != nil {
return nil, err
var prepared *xaiPreparedRequest
steerRequest := isResponsesSteerPayload(req.Payload)
if steerRequest {
prepared = &xaiPreparedRequest{
baseModel: strings.TrimSpace(req.Model),
body: bytes.Clone(req.Payload),
}
} else {
var errPrepare error
prepared, errPrepare = e.prepareResponsesWebsocketRequest(ctx, req, opts)
if errPrepare != nil {
return nil, errPrepare
}
}

reporter := helps.NewExecutorUsageReporter(ctx, e, prepared.baseModel, auth)
Expand Down Expand Up @@ -510,7 +520,7 @@ func (e *XAIWebsocketsExecutor) ExecuteStream(ctx context.Context, auth *cliprox
}
}
idMapper := newXAIWebsocketRequestIDMapper(e.idStore, stateSessionID, req.Payload)
if idMapper != nil {
if idMapper != nil && !steerRequest {
if websocketSessionTargetChanged(sess, authID, wsURL) {
idMapper.upstreamPreviousID = ""
}
Expand Down Expand Up @@ -1487,6 +1497,9 @@ func buildXAIWebsocketRequestBody(body []byte) []byte {
if len(body) == 0 {
return nil
}
if isResponsesSteerPayload(body) {
return bytes.Clone(body)
}
wsReqBody := bytes.Clone(body)
wsReqBody, _ = sjson.SetBytes(wsReqBody, "type", "response.create")
wsReqBody, _ = sjson.DeleteBytes(wsReqBody, "stream")
Expand Down
16 changes: 16 additions & 0 deletions internal/runtime/executor/xai_websockets_executor_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -1502,6 +1502,22 @@ func TestBuildXAIWebsocketRequestBodySetsStoreAndKeepsPromptCacheKey(t *testing.
}
}

func TestBuildXAIWebsocketRequestBodyPreservesSteer(t *testing.T) {
body := []byte(`{"type":"response.steer","previous_response_id":"resp-1","input":"Keep the scope small."}`)

payload := buildXAIWebsocketRequestBody(body)

if got := gjson.GetBytes(payload, "type").String(); got != "response.steer" {
t.Fatalf("type = %q, want response.steer; payload=%s", got, payload)
}
if gjson.GetBytes(payload, "store").Exists() {
t.Fatalf("steer body must not gain store: %s", payload)
}
if gjson.GetBytes(payload, "model").Exists() {
t.Fatalf("steer body must not gain model: %s", payload)
}
}

func TestXAIWebsocketsExecuteStreamCompletesGenerateFalseWarmup(t *testing.T) {
upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }}
capturedPayload := make(chan []byte, 1)
Expand Down
47 changes: 43 additions & 4 deletions sdk/api/handlers/openai/openai_responses_websocket.go
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@ import (
const (
wsRequestTypeCreate = "response.create"
wsRequestTypeAppend = "response.append"
wsRequestTypeSteer = "response.steer"
wsEventTypeError = "error"
wsEventTypeCompleted = "response.completed"
wsEventTypeDone = "response.done"
Expand Down Expand Up @@ -554,8 +555,32 @@ func (h *OpenAIResponsesAPIHandler) ResponsesWebsocket(c *gin.Context) {
continue
}

requestJSON = h.prepareCodexMultiAgentV2Tools(c, requestJSON)
requestJSON = h.prepareCodexOrphanDelegation(c, requestJSON)
if !responsesWebsocketSteerRequest(requestJSON) {
requestJSON = h.prepareCodexMultiAgentV2Tools(c, requestJSON)
requestJSON = h.prepareCodexOrphanDelegation(c, requestJSON)
} else if !nativeWebsocketPassthrough {
errMsg = responsesWebsocketSteeringNotSupportedError()
h.LoggingAPIResponseError(context.WithValue(context.Background(), "gin", c), errMsg)
markAPIResponseTimestamp(c)
errorPayload, errWrite := writeResponsesWebsocketError(writer, wsTimelineLog, errMsg)
log.Infof(
"responses websocket: downstream_out id=%s type=%d event=%s payload=%s",
passthroughSessionID,
websocket.TextMessage,
websocketPayloadEventType(errorPayload),
websocketPayloadPreview(errorPayload),
)
if errWrite != nil {
log.Warnf(
"responses websocket: downstream_out write failed id=%s event=%s error=%v",
passthroughSessionID,
websocketPayloadEventType(errorPayload),
errWrite,
)
return
}
continue
}

if !useUpstreamWebsocketPassthrough && shouldHandleResponsesWebsocketPrewarmLocally(payload, lastRequest, false) {
if updated, errDelete := sjson.DeleteBytes(requestJSON, "generate"); errDelete == nil {
Expand Down Expand Up @@ -588,7 +613,10 @@ func (h *OpenAIResponsesAPIHandler) ResponsesWebsocket(c *gin.Context) {
nextLastRequest = requestJSON
}

modelName := gjson.GetBytes(requestJSON, "model").String()
modelName := strings.TrimSpace(gjson.GetBytes(requestJSON, "model").String())
if modelName == "" {
modelName = requestModelName
}
lastAttemptedAuthID := pinnedAuthID
attemptedUpstreamMode := responsesWebsocketUpstreamModeUnknown
selectedAuthObserved := false
Expand Down Expand Up @@ -701,8 +729,10 @@ func responsesWebsocketHTTPReplayRequiredError() error {
}

func responsesWebsocketRequestRequiresCurrentUpstream(payload []byte) bool {
requestType := strings.TrimSpace(gjson.GetBytes(payload, "type").String())
return strings.TrimSpace(gjson.GetBytes(payload, "previous_response_id").String()) != "" ||
strings.TrimSpace(gjson.GetBytes(payload, "type").String()) == wsRequestTypeAppend
requestType == wsRequestTypeAppend ||
requestType == wsRequestTypeSteer
}

func responsesWebsocketNativePassthroughAllowed(upstreamMode string, useUpstreamWebsocket bool, pinnedAuthID string, upstreamAuthID string) bool {
Expand Down Expand Up @@ -731,6 +761,15 @@ func websocketUpgradeHeaders(req *http.Request) http.Header {
return headers
}

func responsesWebsocketSteeringNotSupportedError() *interfaces.ErrorMessage {
return &interfaces.ErrorMessage{
StatusCode: http.StatusBadRequest,
Error: errors.New(
`{"error":{"message":"response.steer requires an upstream Responses WebSocket on the current connection","type":"invalid_request_error","code":"steering_not_supported"}}`,
),
}
}

func responsesWebsocketPreviousResponseNotFoundError() *interfaces.ErrorMessage {
return &interfaces.ErrorMessage{
StatusCode: http.StatusConflict,
Expand Down
61 changes: 61 additions & 0 deletions sdk/api/handlers/openai/openai_responses_websocket_requests.go
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,12 @@ func normalizeResponsesWebsocketRequestWithIncrementalState(rawJSON []byte, last
case wsRequestTypeAppend:
// log.Infof("responses websocket: response.append request")
return normalizeResponseSubsequentRequest(rawJSON, lastRequest, lastResponseOutput, lastResponseID, lastResponsePendingToolCallIDs, allowIncrementalInputWithPreviousResponseID, allowCompactionReplayBypass)
case wsRequestTypeSteer:
normalized, errMsg := normalizeResponsesWebsocketSteerRequest(rawJSON)
if errMsg != nil {
return nil, lastRequest, errMsg
}
return normalized, lastRequest, nil
default:
return nil, lastRequest, &interfaces.ErrorMessage{
StatusCode: http.StatusBadRequest,
Expand All @@ -47,6 +53,59 @@ func normalizeResponsesWebsocketRequestWithIncrementalState(rawJSON []byte, last
}
}

func responsesWebsocketSteerRequest(rawJSON []byte) bool {
return strings.TrimSpace(gjson.GetBytes(rawJSON, "type").String()) == wsRequestTypeSteer
}

// Official mid-turn steer accepts only type / previous_response_id / input.
// Extra fields are left for upstream to reject; we must not add stream or model.
func normalizeResponsesWebsocketSteerRequest(rawJSON []byte) ([]byte, *interfaces.ErrorMessage) {
if !json.Valid(rawJSON) {
return nil, &interfaces.ErrorMessage{
StatusCode: http.StatusBadRequest,
Error: fmt.Errorf("invalid websocket request JSON"),
}
}
if !responsesWebsocketSteerRequest(rawJSON) {
return nil, &interfaces.ErrorMessage{
StatusCode: http.StatusBadRequest,
Error: fmt.Errorf("unsupported websocket request type: %s", strings.TrimSpace(gjson.GetBytes(rawJSON, "type").String())),
}
}
if strings.TrimSpace(gjson.GetBytes(rawJSON, "previous_response_id").String()) == "" {
return nil, &interfaces.ErrorMessage{
StatusCode: http.StatusBadRequest,
Error: fmt.Errorf("response.steer requires previous_response_id"),
}
}
input := gjson.GetBytes(rawJSON, "input")
if !input.Exists() {
return nil, &interfaces.ErrorMessage{
StatusCode: http.StatusBadRequest,
Error: fmt.Errorf("response.steer requires input"),
}
}
if input.IsArray() && len(input.Array()) == 0 {
return nil, &interfaces.ErrorMessage{
StatusCode: http.StatusBadRequest,
Error: fmt.Errorf("response.steer requires nonempty input"),
}
}
if input.Type == gjson.String && strings.TrimSpace(input.String()) == "" {
return nil, &interfaces.ErrorMessage{
StatusCode: http.StatusBadRequest,
Error: fmt.Errorf("response.steer requires nonempty input"),
}
}
if !input.IsArray() && input.Type != gjson.String {
return nil, &interfaces.ErrorMessage{
StatusCode: http.StatusBadRequest,
Error: fmt.Errorf("response.steer input must be a string or a nonempty array"),
}
}
return bytes.Clone(rawJSON), nil
}

func normalizeResponseCreateRequest(rawJSON []byte) ([]byte, []byte, *interfaces.ErrorMessage) {
input := gjson.GetBytes(rawJSON, "input")
if input.Exists() && !input.IsArray() {
Expand Down Expand Up @@ -722,6 +781,8 @@ func normalizeResponsesWebsocketPassthroughRequest(rawJSON []byte, modelName str
requestType := strings.TrimSpace(gjson.GetBytes(rawJSON, "type").String())
switch requestType {
case wsRequestTypeCreate, wsRequestTypeAppend:
case wsRequestTypeSteer:
return normalizeResponsesWebsocketSteerRequest(rawJSON)
default:
return nil, &interfaces.ErrorMessage{
StatusCode: http.StatusBadRequest,
Expand Down
Loading
Loading