fix(grok): sanitize composer reasoning parameters

This commit is contained in:
Heatherm Huang
2026-07-13 10:11:33 +08:00
parent c4ff604e93
commit aeb34d2003
2 changed files with 80 additions and 0 deletions
@@ -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": {},
}
@@ -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()