diff --git a/providers/openai/language_model.go b/providers/openai/language_model.go index b46651c8b..bf5e43060 100644 --- a/providers/openai/language_model.go +++ b/providers/openai/language_model.go @@ -25,6 +25,7 @@ type languageModel struct { modelID string client openai.Client objectMode fantasy.ObjectMode + reasoningModelFunc func(modelID string) bool prepareCallFunc LanguageModelPrepareCallFunc mapFinishReasonFunc LanguageModelMapFinishReasonFunc extraContentFunc LanguageModelExtraContentFunc @@ -87,6 +88,13 @@ func WithLanguageModelToPromptFunc(fn LanguageModelToPromptFunc) LanguageModelOp } } +// WithLanguageModelReasoningModelFunc overrides reasoning-model detection for Chat Completions. +func WithLanguageModelReasoningModelFunc(fn func(modelID string) bool) LanguageModelOption { + return func(l *languageModel) { + l.reasoningModelFunc = fn + } +} + // WithLanguageModelObjectMode sets the object generation mode. func WithLanguageModelObjectMode(om fantasy.ObjectMode) LanguageModelOption { return func(l *languageModel) { @@ -161,7 +169,7 @@ func (o languageModel) prepareParams(call fantasy.Call) (*openai.ChatCompletionN params.PresencePenalty = param.NewOpt(*call.PresencePenalty) } - if isReasoningModel(o.modelID) { + if o.isReasoningModel() { // remove unsupported settings for reasoning models // see https://platform.openai.com/docs/guides/reasoning#limitations if call.Temperature != nil { @@ -589,6 +597,13 @@ func (o languageModel) Stream(ctx context.Context, call fantasy.Call) (fantasy.S }, nil } +func (o languageModel) isReasoningModel() bool { + if o.reasoningModelFunc != nil { + return o.reasoningModelFunc(o.modelID) + } + return isReasoningModel(o.modelID) +} + func isReasoningModel(modelID string) bool { return strings.HasPrefix(modelID, "o1") || strings.Contains(modelID, "-o1") || strings.HasPrefix(modelID, "o3") || strings.Contains(modelID, "-o3") || diff --git a/providers/openai/language_model_hooks.go b/providers/openai/language_model_hooks.go index f1c851570..c4ee8efaf 100644 --- a/providers/openai/language_model_hooks.go +++ b/providers/openai/language_model_hooks.go @@ -130,7 +130,11 @@ func DefaultPrepareCallFunc(model fantasy.LanguageModel, params *openai.ChatComp } } - if isReasoningModel(model.Model()) { + reasoning := isReasoningModel(model.Model()) + if lm, ok := model.(interface{ isReasoningModel() bool }); ok { + reasoning = lm.isReasoningModel() + } + if reasoning { if providerOptions.LogitBias != nil { params.LogitBias = nil warnings = append(warnings, fantasy.CallWarning{ diff --git a/providers/openai/language_model_params_test.go b/providers/openai/language_model_params_test.go new file mode 100644 index 000000000..f40942860 --- /dev/null +++ b/providers/openai/language_model_params_test.go @@ -0,0 +1,88 @@ +package openai + +import ( + "context" + "testing" + + "charm.land/fantasy" + "github.com/stretchr/testify/require" +) + +func TestPrepareParams_ChatReasoningModelOverride(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + modelID string + opts []Option + wantReasoning bool + }{ + {name: "unknown default", modelID: "totally-new-model"}, + {name: "gpt-5 default", modelID: "gpt-5", wantReasoning: true}, + { + name: "force reasoning", modelID: "totally-new-model", wantReasoning: true, + opts: []Option{WithReasoningModelFunc(func(modelID string) bool { return modelID == "totally-new-model" })}, + }, + { + name: "force non-reasoning", modelID: "gpt-5", + opts: []Option{WithReasoningModelFunc(func(modelID string) bool { return modelID != "gpt-5" })}, + }, + { + name: "language model option", modelID: "totally-new-model", wantReasoning: true, + opts: []Option{WithLanguageModelOptions(WithLanguageModelReasoningModelFunc(func(string) bool { return true }))}, + }, + { + name: "provider option takes precedence", modelID: "gpt-5", + opts: []Option{ + WithLanguageModelOptions(WithLanguageModelReasoningModelFunc(func(string) bool { return true })), + WithReasoningModelFunc(func(string) bool { return false }), + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + provider, err := New(tt.opts...) + require.NoError(t, err) + model, err := provider.LanguageModel(context.Background(), tt.modelID) + require.NoError(t, err) + lm, ok := model.(languageModel) + require.True(t, ok) + + params, warnings, err := lm.prepareParams(fantasy.Call{ + Prompt: fantasy.Prompt{testTextMessage(fantasy.MessageRoleUser, "hello")}, + Temperature: new(0.7), + MaxOutputTokens: new(int64(512)), + ProviderOptions: fantasy.ProviderOptions{ + Name: &ProviderOptions{LogProbs: new(true)}, + }, + }) + require.NoError(t, err) + + if tt.wantReasoning { + require.False(t, params.Temperature.Valid()) + require.False(t, params.MaxTokens.Valid()) + require.True(t, params.MaxCompletionTokens.Valid()) + require.Equal(t, int64(512), params.MaxCompletionTokens.Value) + require.False(t, params.Logprobs.Valid()) + var unsupported []string + for _, warning := range warnings { + require.Equal(t, fantasy.CallWarningTypeUnsupportedSetting, warning.Type) + unsupported = append(unsupported, warning.Setting) + } + require.ElementsMatch(t, []string{"temperature", "Logprobs"}, unsupported) + } else { + require.True(t, params.Temperature.Valid()) + require.Equal(t, 0.7, params.Temperature.Value) + require.True(t, params.MaxTokens.Valid()) + require.Equal(t, int64(512), params.MaxTokens.Value) + require.False(t, params.MaxCompletionTokens.Valid()) + require.True(t, params.Logprobs.Valid()) + require.True(t, params.Logprobs.Value) + require.Empty(t, warnings) + } + }) + } +} diff --git a/providers/openai/openai.go b/providers/openai/openai.go index b4878966a..35847a68a 100644 --- a/providers/openai/openai.go +++ b/providers/openai/openai.go @@ -31,6 +31,7 @@ type options struct { name string useResponsesAPI bool responsesAPIFunc func(modelID string) bool + reasoningModelFunc func(modelID string) bool headers map[string]string userAgent string client option.HTTPClient @@ -143,6 +144,15 @@ func WithResponsesAPIFunc(fn func(modelID string) bool) Option { } } +// WithReasoningModelFunc sets a custom classifier for which models are reasoning models. +// When set, it replaces the built-in model-name heuristics for both the Responses +// and Chat Completions clients. +func WithReasoningModelFunc(fn func(modelID string) bool) Option { + return func(o *options) { + o.reasoningModelFunc = fn + } +} + // WithUserAgent sets an explicit User-Agent header, overriding the default and any // value set via WithHeaders. func WithUserAgent(ua string) Option { @@ -194,11 +204,14 @@ func (o *provider) LanguageModel(_ context.Context, modelID string) (fantasy.Lan if objectMode == fantasy.ObjectModeJSON { objectMode = fantasy.ObjectModeAuto } - return newResponsesLanguageModel(modelID, o.options.name, client, objectMode), nil + return newResponsesLanguageModel(modelID, o.options.name, client, objectMode, o.options.reasoningModelFunc), nil } languageModelOptions := append([]LanguageModelOption{}, o.options.languageModelOptions...) languageModelOptions = append(languageModelOptions, WithLanguageModelObjectMode(o.options.objectMode)) + if o.options.reasoningModelFunc != nil { + languageModelOptions = append(languageModelOptions, WithLanguageModelReasoningModelFunc(o.options.reasoningModelFunc)) + } return newLanguageModel( modelID, diff --git a/providers/openai/responses_language_model.go b/providers/openai/responses_language_model.go index a54a21981..84eb2f1cd 100644 --- a/providers/openai/responses_language_model.go +++ b/providers/openai/responses_language_model.go @@ -24,19 +24,21 @@ import ( const topLogprobsMax = 20 type responsesLanguageModel struct { - provider string - modelID string - client openai.Client - objectMode fantasy.ObjectMode + provider string + modelID string + client openai.Client + objectMode fantasy.ObjectMode + reasoningModelFunc func(modelID string) bool } // newResponsesLanguageModel implements a responses api model. -func newResponsesLanguageModel(modelID string, provider string, client openai.Client, objectMode fantasy.ObjectMode) responsesLanguageModel { +func newResponsesLanguageModel(modelID string, provider string, client openai.Client, objectMode fantasy.ObjectMode, reasoningModelFunc func(modelID string) bool) responsesLanguageModel { return responsesLanguageModel{ - modelID: modelID, - provider: provider, - client: client, - objectMode: objectMode, + modelID: modelID, + provider: provider, + client: client, + objectMode: objectMode, + reasoningModelFunc: reasoningModelFunc, } } @@ -77,7 +79,8 @@ func getResponsesModelConfig(modelID string) responsesModelConfig { supportsPriorityProcessing: supportsPriorityProcessing, } - if strings.Contains(strings.ToLower(modelID), "gpt-5-chat") { + reasoningGeneration := reasoningGenerationPattern.MatchString(strings.ToLower(modelID)) + if reasoningGeneration && strings.Contains(strings.ToLower(modelID), "-chat") { return responsesModelConfig{ isReasoningModel: false, systemMessageMode: defaults.systemMessageMode, @@ -91,7 +94,7 @@ func getResponsesModelConfig(modelID string) responsesModelConfig { strings.HasPrefix(modelID, "o3") || strings.Contains(modelID, "-o3") || strings.HasPrefix(modelID, "o4") || strings.Contains(modelID, "-o4") || strings.HasPrefix(modelID, "oss") || strings.Contains(modelID, "-oss") || - strings.Contains(strings.ToLower(modelID), "gpt-5") || + reasoningGeneration || strings.Contains(modelID, "codex-") || strings.Contains(modelID, "computer-use") { if strings.Contains(modelID, "o1-mini") || strings.Contains(modelID, "o1-preview") { return responsesModelConfig{ @@ -131,6 +134,15 @@ func (o responsesLanguageModel) prepareParams(call fantasy.Call) (*responses.Res params := &responses.ResponseNewParams{} modelConfig := getResponsesModelConfig(o.modelID) + if o.reasoningModelFunc != nil { + if reasoning := o.reasoningModelFunc(o.modelID); reasoning != modelConfig.isReasoningModel { + modelConfig.isReasoningModel = reasoning + modelConfig.systemMessageMode = "system" + if reasoning { + modelConfig.systemMessageMode = "developer" + } + } + } if call.TopK != nil { warnings = append(warnings, fantasy.CallWarning{ diff --git a/providers/openai/responses_options.go b/providers/openai/responses_options.go index 362393a39..5cd878237 100644 --- a/providers/openai/responses_options.go +++ b/providers/openai/responses_options.go @@ -3,6 +3,7 @@ package openai import ( "encoding/json" + "regexp" "slices" "strings" @@ -296,18 +297,26 @@ func ParseResponsesOptions(data map[string]any) (*ResponsesProviderOptions, erro return &options, nil } +// responsesGenerationPattern matches the model generations that only +// speak the Responses API: gpt-4 and gpt-5 today, the newer generations +// as they ship (gpt-6, gpt-10, ...), and never the legacy gpt-3 family +// that predates it. +var responsesGenerationPattern = regexp.MustCompile(`gpt-(?:[4-9]|[1-9]\d)`) + +// reasoningGenerationPattern is the subset of those generations that reason: +// gpt-5 and everything after it, never gpt-4. +var reasoningGenerationPattern = regexp.MustCompile(`gpt-(?:[5-9]|[1-9]\d)`) + // IsResponsesModel checks if a model ID is a Responses API model for OpenAI. func IsResponsesModel(modelID string) bool { return slices.Contains(responsesModelIDs, modelID) || - strings.Contains(strings.ToLower(modelID), "gpt-4") || - strings.Contains(strings.ToLower(modelID), "gpt-5") + responsesGenerationPattern.MatchString(strings.ToLower(modelID)) } // IsResponsesReasoningModel checks if a model ID is a Responses API reasoning model for OpenAI. func IsResponsesReasoningModel(modelID string) bool { return slices.Contains(responsesReasoningModelIDs, modelID) || - strings.Contains(strings.ToLower(modelID), "gpt-4") || - strings.Contains(strings.ToLower(modelID), "gpt-5") + responsesGenerationPattern.MatchString(strings.ToLower(modelID)) } // SearchContextSize controls how much context window space the diff --git a/providers/openai/responses_options_test.go b/providers/openai/responses_options_test.go new file mode 100644 index 000000000..c187a661c --- /dev/null +++ b/providers/openai/responses_options_test.go @@ -0,0 +1,60 @@ +package openai + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestIsResponsesModel(t *testing.T) { + tests := []struct { + modelID string + want bool + }{ + // Explicitly listed models. + {"gpt-4.1", true}, + {"gpt-4o-mini", true}, + {"chatgpt-4o-latest", true}, + {"o3", true}, + {"gpt-oss-120b", true}, + + // Generations caught by the pattern, listed or not. + {"gpt-5", true}, + {"gpt-5.1-codex", true}, + {"gpt-6-astra", true}, + {"GPT-6-ASTRA", true}, + {"gpt-10-turbo", true}, + + // Everything predating the Responses API stays on chat + // completions. + {"gpt-3.5-turbo-1106", true}, // in the explicit list + {"gpt-3-turbo-instruct", false}, + {"babbage-002", false}, + {"davinci-002", false}, + {"some-custom-model", false}, + } + + for _, tt := range tests { + assert.Equal(t, tt.want, IsResponsesModel(tt.modelID), tt.modelID) + } +} + +func TestIsResponsesReasoningModel(t *testing.T) { + tests := []struct { + modelID string + want bool + }{ + {"gpt-5.1-codex", true}, + {"gpt-6-astra", true}, + {"o4-mini", true}, + {"gpt-oss-120b", true}, + + {"gpt-4.1-mini", true}, // gpt-4 matches, as before + {"gpt-3-turbo-instruct", false}, + {"some-custom-model", false}, + } + + for _, tt := range tests { + assert.Equal(t, tt.want, IsResponsesReasoningModel(tt.modelID), tt.modelID) + } +} diff --git a/providers/openai/responses_params_test.go b/providers/openai/responses_params_test.go index f87603b41..e1fa05e69 100644 --- a/providers/openai/responses_params_test.go +++ b/providers/openai/responses_params_test.go @@ -1,6 +1,7 @@ package openai import ( + "context" "encoding/json" "testing" @@ -9,6 +10,109 @@ import ( "github.com/stretchr/testify/require" ) +func TestPrepareParams_ReasoningModelClassification(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + modelID string + force *bool + wantReasoning bool + wantRole string + }{ + {name: "gpt-5", modelID: "gpt-5", wantReasoning: true, wantRole: "developer"}, + {name: "gpt-6", modelID: "gpt-6-astra", wantReasoning: true, wantRole: "developer"}, + {name: "uppercase", modelID: "GPT-6-ASTRA", wantReasoning: true, wantRole: "developer"}, + {name: "gpt-10", modelID: "gpt-10-turbo", wantReasoning: true, wantRole: "developer"}, + {name: "gpt-40", modelID: "gpt-40-turbo", wantReasoning: true, wantRole: "developer"}, + {name: "gpt-100", modelID: "gpt-100-turbo", wantReasoning: true, wantRole: "developer"}, + {name: "gpt-5-chat", modelID: "gpt-5-chat-latest", wantRole: "system"}, + {name: "gpt-6-chat", modelID: "gpt-6-chat-latest", wantRole: "system"}, + {name: "gpt-4o", modelID: "gpt-4o", wantRole: "system"}, + {name: "unknown default", modelID: "totally-new-model", wantRole: "system"}, + {name: "force reasoning", modelID: "totally-new-model", force: new(true), wantReasoning: true, wantRole: "developer"}, + {name: "force non-reasoning", modelID: "gpt-5", force: new(false), wantRole: "system"}, + {name: "preserve remove mode", modelID: "o1-mini", force: new(true), wantReasoning: true}, + {name: "override remove mode", modelID: "o1-mini", force: new(false), wantRole: "system"}, + {name: "agree non-reasoning", modelID: "gpt-4o", force: new(false), wantRole: "system"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + opts := []Option{WithUseResponsesAPI()} + if tt.modelID == "totally-new-model" || tt.modelID == "o1-mini" { + opts = append(opts, WithResponsesAPIFunc(func(string) bool { return true })) + } + if tt.force != nil { + opts = append(opts, WithReasoningModelFunc(func(modelID string) bool { + require.Equal(t, tt.modelID, modelID) + return *tt.force + })) + } + provider, err := New(opts...) + require.NoError(t, err) + model, err := provider.LanguageModel(context.Background(), tt.modelID) + require.NoError(t, err) + lm, ok := model.(responsesLanguageModel) + require.True(t, ok) + + call := testCall(fantasy.Prompt{ + testTextMessage(fantasy.MessageRoleSystem, "Be helpful."), + testTextMessage(fantasy.MessageRoleUser, "hello"), + }, &ResponsesProviderOptions{ + ReasoningEffort: new(ReasoningEffortHigh), + ReasoningSummary: new("detailed"), + }) + call.Temperature = new(0.7) + call.TopP = new(0.9) + params, warnings, err := lm.prepareParams(call) + require.NoError(t, err) + + var unsupported []string + for _, warning := range warnings { + if warning.Type == fantasy.CallWarningTypeUnsupportedSetting { + unsupported = append(unsupported, warning.Setting) + } + } + if tt.wantReasoning { + require.Equal(t, "high", string(params.Reasoning.Effort)) + require.Equal(t, "detailed", string(params.Reasoning.Summary)) + require.False(t, params.Temperature.Valid()) + require.False(t, params.TopP.Valid()) + require.ElementsMatch(t, []string{"temperature", "topP"}, unsupported) + } else { + require.Empty(t, params.Reasoning.Effort) + require.Empty(t, params.Reasoning.Summary) + require.True(t, params.Temperature.Valid()) + require.Equal(t, 0.7, params.Temperature.Value) + require.True(t, params.TopP.Valid()) + require.Equal(t, 0.9, params.TopP.Value) + require.ElementsMatch(t, []string{"reasoningEffort", "reasoningSummary"}, unsupported) + } + + encoded, err := json.Marshal(params) + require.NoError(t, err) + var body map[string]json.RawMessage + require.NoError(t, json.Unmarshal(encoded, &body)) + _, hasReasoning := body["reasoning"] + require.Equal(t, tt.wantReasoning, hasReasoning) + var input []struct { + Role string `json:"role"` + } + require.NoError(t, json.Unmarshal(body["input"], &input)) + if tt.wantRole == "" { + require.Len(t, input, 1) + require.Equal(t, "user", input[0].Role) + } else { + require.Len(t, input, 2) + require.Equal(t, tt.wantRole, input[0].Role) + } + }) + } +} + func TestPrepareParams_Store(t *testing.T) { t.Parallel()