diff --git a/backend/internal/service/openai_gateway_grok.go b/backend/internal/service/openai_gateway_grok.go index 502f99f707..379b586136 100644 --- a/backend/internal/service/openai_gateway_grok.go +++ b/backend/internal/service/openai_gateway_grok.go @@ -154,6 +154,10 @@ func patchGrokResponsesBody(body []byte, upstreamModel string) ([]byte, error) { if err != nil { return nil, err } + out, err = sanitizeGrokResponsesModelCapabilities(out, upstreamModel) + if err != nil { + return nil, err + } for _, unsupportedField := range []string{"prompt_cache_retention", "safety_identifier"} { if gjson.GetBytes(out, unsupportedField).Exists() { out, err = sjson.DeleteBytes(out, unsupportedField) @@ -187,6 +191,38 @@ func patchGrokResponsesBody(body []byte, upstreamModel string) ([]byte, error) { return out, nil } +func sanitizeGrokResponsesModelCapabilities(body []byte, upstreamModel string) ([]byte, error) { + if !grokModelRejectsReasoningEffort(upstreamModel) { + return body, nil + } + + out := body + for _, field := range []string{"reasoning", "reasoning_effort", "reasoningEffort"} { + if !gjson.GetBytes(out, field).Exists() { + continue + } + var err error + out, err = sjson.DeleteBytes(out, field) + if err != nil { + return nil, fmt.Errorf("remove unsupported Grok Composer %s: %w", field, err) + } + } + return out, nil +} + +func grokModelRejectsReasoningEffort(model string) bool { + model = strings.TrimSpace(strings.ToLower(model)) + if slash := strings.LastIndex(model, "/"); slash >= 0 { + model = strings.TrimSpace(model[slash+1:]) + } + switch model { + case "grok-composer", "grok-composer-2.5-fast", "composer-2.5": + return true + default: + return false + } +} + var grokResponsesUnsupportedRecursiveFields = map[string]struct{}{ "external_web_access": {}, } diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index 6a0942eb72..13020b1c28 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -42,6 +42,50 @@ func TestPatchGrokResponsesBodySetsMappedModelAndDropsUnsupportedFields(t *testi require.Equal(t, "high", gjson.GetBytes(patched, "reasoning.effort").String()) } +func TestPatchGrokResponsesBodySanitizesComposerReasoningParameters(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + upstreamModel string + wantReasoning bool + }{ + {name: "composer fast", upstreamModel: "grok-composer-2.5-fast"}, + {name: "composer shorthand", upstreamModel: "grok-composer"}, + {name: "composer legacy alias", upstreamModel: "composer-2.5"}, + {name: "provider-prefixed composer", upstreamModel: "xai/grok-composer-2.5-fast"}, + {name: "grok 4.5", upstreamModel: "grok-4.5", wantReasoning: true}, + } + + body := []byte(`{ + "model": "grok", + "input": "hello", + "reasoning": {"effort": "medium", "summary": "auto"}, + "reasoning_effort": "medium", + "reasoningEffort": "medium" + }`) + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + patched, err := patchGrokResponsesBody(body, tt.upstreamModel) + require.NoError(t, err) + require.True(t, json.Valid(patched)) + require.Equal(t, tt.upstreamModel, gjson.GetBytes(patched, "model").String()) + + if tt.wantReasoning { + require.Equal(t, "medium", gjson.GetBytes(patched, "reasoning.effort").String()) + require.Equal(t, "medium", gjson.GetBytes(patched, "reasoning_effort").String()) + require.Equal(t, "medium", gjson.GetBytes(patched, "reasoningEffort").String()) + return + } + + require.False(t, gjson.GetBytes(patched, "reasoning").Exists()) + require.False(t, gjson.GetBytes(patched, "reasoning_effort").Exists()) + require.False(t, gjson.GetBytes(patched, "reasoningEffort").Exists()) + }) + } +} + func TestExtractGrokResponsesReasoningEffortSupportsOpenAICompatibleField(t *testing.T) { t.Parallel()