Skip to content
Merged
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
17 changes: 16 additions & 1 deletion providers/openai/language_model.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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) {
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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") ||
Expand Down
6 changes: 5 additions & 1 deletion providers/openai/language_model_hooks.go
Original file line number Diff line number Diff line change
Expand Up @@ -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{
Expand Down
88 changes: 88 additions & 0 deletions providers/openai/language_model_params_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
})
}
}
15 changes: 14 additions & 1 deletion providers/openai/openai.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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,
Expand Down
34 changes: 23 additions & 11 deletions providers/openai/responses_language_model.go
Original file line number Diff line number Diff line change
Expand Up @@ -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,
}
}

Expand Down Expand Up @@ -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,
Expand All @@ -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{
Expand Down Expand Up @@ -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{
Expand Down
17 changes: 13 additions & 4 deletions providers/openai/responses_options.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package openai

import (
"encoding/json"
"regexp"
"slices"
"strings"

Expand Down Expand Up @@ -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
Expand Down
60 changes: 60 additions & 0 deletions providers/openai/responses_options_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
}
Loading
Loading