**fix(xai): improve namespace-specific tool parameter handling and simplify Codex automation schema**

- Refined tool normalization logic to pass namespace context (`namespaceName`) to `normalizeXAITool`.
- Limited schema simplification to `codex_app.automation_update` to avoid unintended modifications to unrelated tools or namespaces.
- Updated `xaiFunctionParametersNeedSimplification` to match tools by exact namespace and name.
- Enhanced tests to validate correct schema simplification for `codex_app.automation_update` while preserving unrelated tool definitions.
This commit is contained in:
Luis Pater
2026-07-10 13:11:55 +08:00
parent bf25331c77
commit ca67caf087
4 changed files with 207 additions and 44 deletions

View File

@@ -36,17 +36,18 @@ var (
) )
const ( const (
xaiImageHandlerType = "openai-image" xaiImageHandlerType = "openai-image"
xaiVideoHandlerType = "openai-video" xaiVideoHandlerType = "openai-video"
xaiCustomToolType = "custom" xaiCustomToolType = "custom"
xaiFunctionToolType = "function" xaiFunctionToolType = "function"
xaiImageGenerationToolType = "image_generation" xaiImageGenerationToolType = "image_generation"
xaiNamespaceToolType = "namespace" xaiNamespaceToolType = "namespace"
xaiToolSearchType = "tool_search" xaiToolSearchType = "tool_search"
xaiWebSearchToolType = "web_search" xaiWebSearchToolType = "web_search"
// Codex Desktop injects codex_app.automation_update with a large oneOf+$ref // Codex Desktop injects codex_app.automation_update with a large oneOf+$ref
// schema. xAI's free/build Responses path accepts the HTTP request but never // schema. xAI's free/build Responses path accepts the HTTP request but never
// emits SSE when that schema is present, so Desktop hangs on "thinking". // emits SSE when that schema is present, so Desktop hangs on "thinking".
xaiCodexAppNamespaceName = "codex_app"
xaiAutomationUpdateToolName = "automation_update" xaiAutomationUpdateToolName = "automation_update"
// Permissive placeholder schema: keeps the tool callable without the hang. // Permissive placeholder schema: keeps the tool callable without the hang.
xaiSafeFunctionParameters = `{"type":"object","properties":{},"additionalProperties":true}` xaiSafeFunctionParameters = `{"type":"object","properties":{},"additionalProperties":true}`
@@ -1049,9 +1050,10 @@ func normalizeXAITools(body []byte) []byte {
toolType := tool.Get("type").String() toolType := tool.Get("type").String()
if toolType == xaiNamespaceToolType { if toolType == xaiNamespaceToolType {
changed = true changed = true
namespaceName := tool.Get("name").String()
if namespaceTools := tool.Get("tools"); namespaceTools.IsArray() { if namespaceTools := tool.Get("tools"); namespaceTools.IsArray() {
for _, nestedTool := range namespaceTools.Array() { for _, nestedTool := range namespaceTools.Array() {
nestedRaw, nestedChanged, ok := normalizeXAITool(nestedTool) nestedRaw, nestedChanged, ok := normalizeXAITool(nestedTool, namespaceName)
if !ok { if !ok {
return body return body
} }
@@ -1068,7 +1070,7 @@ func normalizeXAITools(body []byte) []byte {
} }
continue continue
} }
raw, toolChanged, ok := normalizeXAITool(tool) raw, toolChanged, ok := normalizeXAITool(tool, "")
if !ok { if !ok {
return body return body
} }
@@ -1114,7 +1116,7 @@ func normalizeXAIToolChoiceForTools(body []byte) []byte {
return body return body
} }
func normalizeXAITool(tool gjson.Result) ([]byte, bool, bool) { func normalizeXAITool(tool gjson.Result, namespaceName string) ([]byte, bool, bool) {
toolType := tool.Get("type").String() toolType := tool.Get("type").String()
changed := false changed := false
if toolType == xaiToolSearchType || toolType == xaiImageGenerationToolType { if toolType == xaiToolSearchType || toolType == xaiImageGenerationToolType {
@@ -1149,41 +1151,34 @@ func normalizeXAITool(tool gjson.Result) ([]byte, bool, bool) {
raw = updatedTool raw = updatedTool
changed = true changed = true
} }
// Codex Desktop's automation_update schema (and similar large oneOf+$ref // Codex Desktop's codex_app.automation_update schema hangs xAI free/build
// function schemas) hang xAI free/build streaming. Simplify parameters so // streaming. Limit the workaround to that exact namespaced tool so unrelated
// the request still carries the tool name but does not stall the stream. // tools keep their parameter contracts.
if toolType == xaiFunctionToolType && xaiFunctionParametersNeedSimplification(tool) { if toolType == xaiFunctionToolType && xaiFunctionParametersNeedSimplification(tool, namespaceName) {
updatedTool, errSet := sjson.SetRawBytes(raw, "parameters", []byte(xaiSafeFunctionParameters)) updatedTool, errSet := sjson.SetRawBytes(raw, "parameters", []byte(xaiSafeFunctionParameters))
if errSet != nil { if errSet != nil {
return nil, false, false return nil, false, false
} }
raw = updatedTool raw = updatedTool
if strict := tool.Get("strict"); strict.Exists() && strict.Bool() {
updatedTool, errSet = sjson.SetBytes(raw, "strict", false)
if errSet != nil {
return nil, false, false
}
raw = updatedTool
}
changed = true changed = true
log.Debugf("xai: simplified parameters for tool %s to avoid upstream hang", tool.Get("name").String()) log.Debugf("xai: simplified parameters for tool %s.%s to avoid upstream hang", namespaceName, tool.Get("name").String())
} }
return raw, changed, true return raw, changed, true
} }
// xaiFunctionParametersNeedSimplification reports whether a function tool's // xaiFunctionParametersNeedSimplification reports whether a function tool is
// JSON Schema is known/likely to hang xAI Responses streaming. // the Codex Desktop automation tool known to hang xAI Responses streaming.
func xaiFunctionParametersNeedSimplification(tool gjson.Result) bool { func xaiFunctionParametersNeedSimplification(tool gjson.Result, namespaceName string) bool {
name := strings.TrimSpace(tool.Get("name").String()) return strings.EqualFold(strings.TrimSpace(tool.Get("type").String()), xaiFunctionToolType) &&
if strings.EqualFold(name, xaiAutomationUpdateToolName) { strings.EqualFold(strings.TrimSpace(namespaceName), xaiCodexAppNamespaceName) &&
return true strings.EqualFold(strings.TrimSpace(tool.Get("name").String()), xaiAutomationUpdateToolName)
}
params := tool.Get("parameters")
if !params.Exists() {
return false
}
raw := params.Raw
// Heuristic: large schemas combining oneOf with $ref/$defs hang the free
// Grok Responses path (no SSE events until client cancel).
if len(raw) < 1500 {
return false
}
hasOneOf := strings.Contains(raw, `"oneOf"`)
hasRef := strings.Contains(raw, `"$ref"`) || strings.Contains(raw, `"$defs"`)
return hasOneOf && hasRef
} }
func sanitizeXAIInputEncryptedContent(body []byte) []byte { func sanitizeXAIInputEncryptedContent(body []byte) []byte {

View File

@@ -1060,10 +1060,10 @@ func TestXAIExecutorExecuteVideosUsesNativeEndpointFromRequestPath(t *testing.T)
} }
} }
func TestNormalizeXAITools_SimplifiesAutomationUpdateSchema(t *testing.T) { func TestNormalizeXAITools_SimplifiesCodexAppAutomationUpdateSchema(t *testing.T) {
// Large oneOf+$ref schema mimicking Codex Desktop codex_app.automation_update. // Large oneOf+$ref schema mimicking Codex Desktop codex_app.automation_update.
params := `{"oneOf":[{"type":"object","properties":{"mode":{"type":"string"}}}],"$defs":{"a":{"type":"string"}},"x":"` + strings.Repeat("y", 1600) + `"}` params := `{"oneOf":[{"type":"object","properties":{"mode":{"type":"string"}}}],"$defs":{"a":{"type":"string"}},"x":"` + strings.Repeat("y", 1600) + `"}`
body := []byte(`{"model":"grok-4.5","tools":[{"type":"namespace","name":"codex_app","tools":[{"type":"function","name":"automation_update","description":"sched","strict":false,"parameters":` + params + `}]},{"type":"function","name":"exec_command","parameters":{"type":"object","properties":{"cmd":{"type":"string"}}}}]}`) body := []byte(`{"model":"grok-4.5","tools":[{"type":"namespace","name":"codex_app","tools":[{"type":"function","name":"automation_update","description":"sched","strict":true,"parameters":` + params + `}]},{"type":"function","name":"exec_command","parameters":{"type":"object","properties":{"cmd":{"type":"string"}}}}]}`)
out := normalizeXAITools(body) out := normalizeXAITools(body)
tools := gjson.GetBytes(out, "tools") tools := gjson.GetBytes(out, "tools")
@@ -1086,6 +1086,9 @@ func TestNormalizeXAITools_SimplifiesAutomationUpdateSchema(t *testing.T) {
if tool.Get("parameters.additionalProperties").Type != gjson.True { if tool.Get("parameters.additionalProperties").Type != gjson.True {
t.Fatalf("automation_update parameters should allow additionalProperties: %s", paramsRaw) t.Fatalf("automation_update parameters should allow additionalProperties: %s", paramsRaw)
} }
if tool.Get("strict").Type != gjson.False {
t.Fatalf("automation_update strict = %s, want false", tool.Get("strict").Raw)
}
case "exec_command": case "exec_command":
foundExec = true foundExec = true
if got := tool.Get("parameters.properties.cmd.type").String(); got != "string" { if got := tool.Get("parameters.properties.cmd.type").String(); got != "string" {
@@ -1101,14 +1104,72 @@ func TestNormalizeXAITools_SimplifiesAutomationUpdateSchema(t *testing.T) {
} }
} }
func TestNormalizeXAITools_PreservesUnrelatedSchemas(t *testing.T) {
largeParams := `{"oneOf":[{"type":"object","properties":{"mode":{"type":"string"}}}],"$defs":{"a":{"type":"string"}},"x":"` + strings.Repeat("y", 1600) + `"}`
tests := []struct {
name string
body []byte
}{
{
name: "top-level automation_update",
body: []byte(`{"tools":[{"type":"function","name":"automation_update","strict":true,"parameters":{"type":"object","properties":{"cron":{"type":"string"}},"required":["cron"],"additionalProperties":false}}]}`),
},
{
name: "automation_update in another namespace",
body: []byte(`{"tools":[{"type":"namespace","name":"calendar","tools":[{"type":"function","name":"automation_update","strict":true,"parameters":{"type":"object","properties":{"cron":{"type":"string"}},"required":["cron"],"additionalProperties":false}}]}]}`),
},
{
name: "custom automation_update in codex_app",
body: []byte(`{"tools":[{"type":"namespace","name":"codex_app","tools":[{"type":"custom","name":"automation_update","strict":true,"parameters":{"type":"object","properties":{"cron":{"type":"string"}},"required":["cron"],"additionalProperties":false}}]}]}`),
},
{
name: "large schema on another codex_app function",
body: []byte(`{"tools":[{"type":"namespace","name":"codex_app","tools":[{"type":"function","name":"exec_command","strict":true,"parameters":` + largeParams + `}]}]}`),
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
out := normalizeXAITools(tt.body)
tool := gjson.GetBytes(out, "tools.0")
if tool.Get("strict").Type != gjson.True {
t.Fatalf("strict changed for unrelated tool: %s", string(out))
}
params := tool.Get("parameters")
if tt.name == "large schema on another codex_app function" {
if !params.Get("oneOf").Exists() || !params.Get("$defs").Exists() {
t.Fatalf("large schema was simplified: %s", string(out))
}
return
}
if got := params.Get("properties.cron.type").String(); got != "string" {
t.Fatalf("schema was simplified, cron type = %q: %s", got, string(out))
}
if params.Get("additionalProperties").Type != gjson.False {
t.Fatalf("additionalProperties changed: %s", string(out))
}
})
}
}
func TestXAIFunctionParametersNeedSimplification(t *testing.T) { func TestXAIFunctionParametersNeedSimplification(t *testing.T) {
auto := gjson.Parse(`{"type":"function","name":"automation_update","parameters":{"type":"object"}}`) auto := gjson.Parse(`{"type":"function","name":"automation_update","parameters":{"type":"object"}}`)
if !xaiFunctionParametersNeedSimplification(auto) { if !xaiFunctionParametersNeedSimplification(auto, "codex_app") {
t.Fatal("automation_update should always need simplification") t.Fatal("codex_app.automation_update should need simplification")
}
if xaiFunctionParametersNeedSimplification(auto, "calendar") {
t.Fatal("automation_update outside codex_app should not need simplification")
}
if xaiFunctionParametersNeedSimplification(auto, "") {
t.Fatal("top-level automation_update should not need simplification")
}
custom := gjson.Parse(`{"type":"custom","name":"automation_update","parameters":{"type":"object"}}`)
if xaiFunctionParametersNeedSimplification(custom, "codex_app") {
t.Fatal("custom codex_app.automation_update should not need simplification")
} }
safe := gjson.Parse(`{"type":"function","name":"exec_command","parameters":{"type":"object","properties":{"cmd":{"type":"string"}}}}`) safe := gjson.Parse(`{"type":"function","name":"exec_command","parameters":{"type":"object","properties":{"cmd":{"type":"string"}}}}`)
if xaiFunctionParametersNeedSimplification(safe) { if xaiFunctionParametersNeedSimplification(safe, "codex_app") {
t.Fatal("simple schema should not need simplification") t.Fatal("unrelated codex_app function should not need simplification")
} }
} }

View File

@@ -471,7 +471,7 @@ func (e *XAIWebsocketsExecutor) ExecuteStream(ctx context.Context, auth *cliprox
if sess != nil { if sess != nil {
sess.reqMu.Unlock() sess.reqMu.Unlock()
} }
return nil, statusErr{code: respHS.StatusCode, msg: string(bodyErr)} return nil, xaiStatusErr(respHS.StatusCode, bodyErr)
} }
helps.RecordAPIWebsocketError(ctx, e.cfg, "dial", errDial) helps.RecordAPIWebsocketError(ctx, e.cfg, "dial", errDial)
if sess != nil { if sess != nil {
@@ -498,10 +498,14 @@ func (e *XAIWebsocketsExecutor) ExecuteStream(ctx context.Context, auth *cliprox
e.invalidateUpstreamConn(sess, conn, "send_error", errSend) e.invalidateUpstreamConn(sess, conn, "send_error", errSend)
connRetry, respHSRetry, errDialRetry := e.ensureUpstreamConn(ctx, auth, sess, authID, wsURL, wsHeaders) connRetry, respHSRetry, errDialRetry := e.ensureUpstreamConn(ctx, auth, sess, authID, wsURL, wsHeaders)
if errDialRetry != nil || connRetry == nil { if errDialRetry != nil || connRetry == nil {
bodyErrRetry := websocketHandshakeBody(respHSRetry)
closeHTTPResponseBody(respHSRetry, "xai websockets executor: close handshake response body error") closeHTTPResponseBody(respHSRetry, "xai websockets executor: close handshake response body error")
helps.RecordAPIWebsocketError(ctx, e.cfg, "dial_retry", errDialRetry) helps.RecordAPIWebsocketError(ctx, e.cfg, "dial_retry", errDialRetry)
sess.clearActive(readCh) sess.clearActive(readCh)
sess.reqMu.Unlock() sess.reqMu.Unlock()
if respHSRetry != nil && respHSRetry.StatusCode > 0 {
return nil, xaiStatusErr(respHSRetry.StatusCode, bodyErrRetry)
}
return nil, errDialRetry return nil, errDialRetry
} }
wsReqBodyRetry := buildXAIWebsocketRequestBody(prepared.body) wsReqBodyRetry := buildXAIWebsocketRequestBody(prepared.body)
@@ -815,6 +819,13 @@ func buildXAIWebsocketWarmupCompletedPayload(createdPayload []byte) []byte {
func parseXAIWebsocketError(payload []byte) (error, bool) { func parseXAIWebsocketError(payload []byte) (error, bool) {
if wsErr, ok := parseCodexWebsocketError(payload); ok { if wsErr, ok := parseCodexWebsocketError(payload); ok {
if statusError, okStatus := wsErr.(statusErrWithHeaders); okStatus {
xaiError := xaiStatusErr(statusError.code, payload)
if xaiError.retryAfter != nil {
statusError.retryAfter = xaiError.retryAfter
}
return statusError, true
}
return wsErr, true return wsErr, true
} }
if len(payload) == 0 || !gjson.GetBytes(payload, "error").Exists() { if len(payload) == 0 || !gjson.GetBytes(payload, "error").Exists() {
@@ -833,7 +844,7 @@ func parseXAIWebsocketError(payload []byte) (error, bool) {
if errNode := gjson.GetBytes(payload, "error"); errNode.Exists() { if errNode := gjson.GetBytes(payload, "error"); errNode.Exists() {
out, _ = sjson.SetRawBytes(out, "error", []byte(errNode.Raw)) out, _ = sjson.SetRawBytes(out, "error", []byte(errNode.Raw))
} }
return statusErr{code: status, msg: string(out)}, true return xaiStatusErr(status, out), true
} }
func xaiBareWebsocketErrorStatus(payload []byte) int { func xaiBareWebsocketErrorStatus(payload []byte) int {

View File

@@ -712,6 +712,102 @@ func TestXAIWebsocketsExecuteStreamCompletesGenerateFalseWarmup(t *testing.T) {
} }
} }
func TestXAIWebsocketsExecuteStreamHandshakeFreeUsageExhaustedSetsRetryAfter(t *testing.T) {
body := []byte(`{"code":"subscription:free-usage-exhausted","error":"You've used all the included free usage for now."}`)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusTooManyRequests)
if _, errWrite := w.Write(body); errWrite != nil {
t.Errorf("write handshake rejection: %v", errWrite)
}
}))
defer server.Close()
exec := NewXAIWebsocketsExecutor(&config.Config{})
auth := &cliproxyauth.Auth{
ID: "xai-auth-free-usage",
Provider: "xai",
Attributes: map[string]string{
"base_url": server.URL,
"websockets": "true",
},
Metadata: map[string]any{"access_token": "xai-token"},
}
req := cliproxyexecutor.Request{
Model: "grok-4.3",
Payload: []byte(`{"model":"grok-4.3","input":"hello"}`),
}
opts := cliproxyexecutor.Options{
SourceFormat: sdktranslator.FormatOpenAIResponse,
ResponseFormat: sdktranslator.FormatOpenAIResponse,
}
_, err := exec.ExecuteStream(context.Background(), auth, req, opts)
if err == nil {
t.Fatal("ExecuteStream() error = nil, want handshake rejection")
}
status, ok := err.(interface{ StatusCode() int })
if !ok || status.StatusCode() != http.StatusTooManyRequests {
t.Fatalf("status = %#v, want 429", err)
}
retryable, ok := err.(interface{ RetryAfter() *time.Duration })
if !ok || retryable.RetryAfter() == nil {
t.Fatalf("expected RetryAfter for free-usage-exhausted handshake error: %#v", err)
}
if got := *retryable.RetryAfter(); got != 24*time.Hour {
t.Fatalf("RetryAfter = %v, want 24h", got)
}
if got := err.Error(); got != string(body) {
t.Fatalf("error payload = %q, want %q", got, body)
}
}
func TestParseXAIWebsocketErrorFreeUsageExhaustedSetsRetryAfter(t *testing.T) {
payload := []byte(`{"type":"error","status":429,"error":{"code":"subscription:free-usage-exhausted","message":"You've used all the included free usage for now."}}`)
err, ok := parseXAIWebsocketError(payload)
if !ok {
t.Fatal("expected xAI websocket error")
}
retryable, ok := err.(interface{ RetryAfter() *time.Duration })
if !ok || retryable.RetryAfter() == nil {
t.Fatalf("expected RetryAfter for free-usage-exhausted websocket event: %#v", err)
}
if got := *retryable.RetryAfter(); got != 24*time.Hour {
t.Fatalf("RetryAfter = %v, want 24h", got)
}
parsed := gjson.Parse(err.Error())
if got := parsed.Get("status").Int(); got != http.StatusTooManyRequests {
t.Fatalf("error status = %d, want 429; payload=%s", got, err)
}
if got := parsed.Get("error.code").String(); got != "subscription:free-usage-exhausted" {
t.Fatalf("error code = %q, want free-usage-exhausted; payload=%s", got, err)
}
}
func TestParseXAIWebsocketBareErrorFreeUsageExhaustedSetsRetryAfter(t *testing.T) {
payload := []byte(`{"status":429,"error":{"code":"subscription:free-usage-exhausted","message":"You've used all the included free usage for now."}}`)
err, ok := parseXAIWebsocketError(payload)
if !ok {
t.Fatal("expected bare xAI websocket error")
}
retryable, ok := err.(interface{ RetryAfter() *time.Duration })
if !ok || retryable.RetryAfter() == nil {
t.Fatalf("expected RetryAfter for bare free-usage-exhausted websocket event: %#v", err)
}
if got := *retryable.RetryAfter(); got != 24*time.Hour {
t.Fatalf("RetryAfter = %v, want 24h", got)
}
parsed := gjson.Parse(err.Error())
if got := parsed.Get("type").String(); got != "error" {
t.Fatalf("error type = %q, want error; payload=%s", got, err)
}
if got := parsed.Get("error.code").String(); got != "subscription:free-usage-exhausted" {
t.Fatalf("error code = %q, want free-usage-exhausted; payload=%s", got, err)
}
}
func TestXAIWebsocketsExecuteStreamStopsOnBareErrorPayload(t *testing.T) { func TestXAIWebsocketsExecuteStreamStopsOnBareErrorPayload(t *testing.T) {
upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }} upgrader := websocket.Upgrader{CheckOrigin: func(*http.Request) bool { return true }}
releaseServer := make(chan struct{}) releaseServer := make(chan struct{})