Skip to content

Commit a511353

Browse files
authored
feat(providers/openai): allow reasoning-model overrides and recognize gpt-6+ (#57)
Cherry-picks upstream charmbracelet/fantasy cfb0530 (fix: gpt series 6+, charmbracelet#354) unchanged, extends the gpt-5+ generation match into getResponsesModelConfig so the Responses client sends reasoning parameters (and drops temperature/top_p) for gpt-6 and later, and adds openai.WithReasoningModelFunc / WithLanguageModelReasoningModelFunc so callers can override the classification for both the Responses and Chat Completions clients. > Xum acted on behalf of @ibetitsmike for this merge.
1 parent bb47dc7 commit a511353

8 files changed

Lines changed: 323 additions & 18 deletions

‎providers/openai/language_model.go‎

Lines changed: 16 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,7 @@ type languageModel struct {
2525
modelID string
2626
client openai.Client
2727
objectMode fantasy.ObjectMode
28+
reasoningModelFunc func(modelID string) bool
2829
prepareCallFunc LanguageModelPrepareCallFunc
2930
mapFinishReasonFunc LanguageModelMapFinishReasonFunc
3031
extraContentFunc LanguageModelExtraContentFunc
@@ -87,6 +88,13 @@ func WithLanguageModelToPromptFunc(fn LanguageModelToPromptFunc) LanguageModelOp
8788
}
8889
}
8990

91+
// WithLanguageModelReasoningModelFunc overrides reasoning-model detection for Chat Completions.
92+
func WithLanguageModelReasoningModelFunc(fn func(modelID string) bool) LanguageModelOption {
93+
return func(l *languageModel) {
94+
l.reasoningModelFunc = fn
95+
}
96+
}
97+
9098
// WithLanguageModelObjectMode sets the object generation mode.
9199
func WithLanguageModelObjectMode(om fantasy.ObjectMode) LanguageModelOption {
92100
return func(l *languageModel) {
@@ -161,7 +169,7 @@ func (o languageModel) prepareParams(call fantasy.Call) (*openai.ChatCompletionN
161169
params.PresencePenalty = param.NewOpt(*call.PresencePenalty)
162170
}
163171

164-
if isReasoningModel(o.modelID) {
172+
if o.isReasoningModel() {
165173
// remove unsupported settings for reasoning models
166174
// see https://platform.openai.com/docs/guides/reasoning#limitations
167175
if call.Temperature != nil {
@@ -589,6 +597,13 @@ func (o languageModel) Stream(ctx context.Context, call fantasy.Call) (fantasy.S
589597
}, nil
590598
}
591599

600+
func (o languageModel) isReasoningModel() bool {
601+
if o.reasoningModelFunc != nil {
602+
return o.reasoningModelFunc(o.modelID)
603+
}
604+
return isReasoningModel(o.modelID)
605+
}
606+
592607
func isReasoningModel(modelID string) bool {
593608
return strings.HasPrefix(modelID, "o1") || strings.Contains(modelID, "-o1") ||
594609
strings.HasPrefix(modelID, "o3") || strings.Contains(modelID, "-o3") ||

‎providers/openai/language_model_hooks.go‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -130,7 +130,11 @@ func DefaultPrepareCallFunc(model fantasy.LanguageModel, params *openai.ChatComp
130130
}
131131
}
132132

133-
if isReasoningModel(model.Model()) {
133+
reasoning := isReasoningModel(model.Model())
134+
if lm, ok := model.(interface{ isReasoningModel() bool }); ok {
135+
reasoning = lm.isReasoningModel()
136+
}
137+
if reasoning {
134138
if providerOptions.LogitBias != nil {
135139
params.LogitBias = nil
136140
warnings = append(warnings, fantasy.CallWarning{
Lines changed: 88 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,88 @@
1+
package openai
2+
3+
import (
4+
"context"
5+
"testing"
6+
7+
"charm.land/fantasy"
8+
"github.com/stretchr/testify/require"
9+
)
10+
11+
func TestPrepareParams_ChatReasoningModelOverride(t *testing.T) {
12+
t.Parallel()
13+
14+
tests := []struct {
15+
name string
16+
modelID string
17+
opts []Option
18+
wantReasoning bool
19+
}{
20+
{name: "unknown default", modelID: "totally-new-model"},
21+
{name: "gpt-5 default", modelID: "gpt-5", wantReasoning: true},
22+
{
23+
name: "force reasoning", modelID: "totally-new-model", wantReasoning: true,
24+
opts: []Option{WithReasoningModelFunc(func(modelID string) bool { return modelID == "totally-new-model" })},
25+
},
26+
{
27+
name: "force non-reasoning", modelID: "gpt-5",
28+
opts: []Option{WithReasoningModelFunc(func(modelID string) bool { return modelID != "gpt-5" })},
29+
},
30+
{
31+
name: "language model option", modelID: "totally-new-model", wantReasoning: true,
32+
opts: []Option{WithLanguageModelOptions(WithLanguageModelReasoningModelFunc(func(string) bool { return true }))},
33+
},
34+
{
35+
name: "provider option takes precedence", modelID: "gpt-5",
36+
opts: []Option{
37+
WithLanguageModelOptions(WithLanguageModelReasoningModelFunc(func(string) bool { return true })),
38+
WithReasoningModelFunc(func(string) bool { return false }),
39+
},
40+
},
41+
}
42+
43+
for _, tt := range tests {
44+
t.Run(tt.name, func(t *testing.T) {
45+
t.Parallel()
46+
47+
provider, err := New(tt.opts...)
48+
require.NoError(t, err)
49+
model, err := provider.LanguageModel(context.Background(), tt.modelID)
50+
require.NoError(t, err)
51+
lm, ok := model.(languageModel)
52+
require.True(t, ok)
53+
54+
params, warnings, err := lm.prepareParams(fantasy.Call{
55+
Prompt: fantasy.Prompt{testTextMessage(fantasy.MessageRoleUser, "hello")},
56+
Temperature: new(0.7),
57+
MaxOutputTokens: new(int64(512)),
58+
ProviderOptions: fantasy.ProviderOptions{
59+
Name: &ProviderOptions{LogProbs: new(true)},
60+
},
61+
})
62+
require.NoError(t, err)
63+
64+
if tt.wantReasoning {
65+
require.False(t, params.Temperature.Valid())
66+
require.False(t, params.MaxTokens.Valid())
67+
require.True(t, params.MaxCompletionTokens.Valid())
68+
require.Equal(t, int64(512), params.MaxCompletionTokens.Value)
69+
require.False(t, params.Logprobs.Valid())
70+
var unsupported []string
71+
for _, warning := range warnings {
72+
require.Equal(t, fantasy.CallWarningTypeUnsupportedSetting, warning.Type)
73+
unsupported = append(unsupported, warning.Setting)
74+
}
75+
require.ElementsMatch(t, []string{"temperature", "Logprobs"}, unsupported)
76+
} else {
77+
require.True(t, params.Temperature.Valid())
78+
require.Equal(t, 0.7, params.Temperature.Value)
79+
require.True(t, params.MaxTokens.Valid())
80+
require.Equal(t, int64(512), params.MaxTokens.Value)
81+
require.False(t, params.MaxCompletionTokens.Valid())
82+
require.True(t, params.Logprobs.Valid())
83+
require.True(t, params.Logprobs.Value)
84+
require.Empty(t, warnings)
85+
}
86+
})
87+
}
88+
}

‎providers/openai/openai.go‎

Lines changed: 14 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@ type options struct {
3131
name string
3232
useResponsesAPI bool
3333
responsesAPIFunc func(modelID string) bool
34+
reasoningModelFunc func(modelID string) bool
3435
headers map[string]string
3536
userAgent string
3637
client option.HTTPClient
@@ -143,6 +144,15 @@ func WithResponsesAPIFunc(fn func(modelID string) bool) Option {
143144
}
144145
}
145146

147+
// WithReasoningModelFunc sets a custom classifier for which models are reasoning models.
148+
// When set, it replaces the built-in model-name heuristics for both the Responses
149+
// and Chat Completions clients.
150+
func WithReasoningModelFunc(fn func(modelID string) bool) Option {
151+
return func(o *options) {
152+
o.reasoningModelFunc = fn
153+
}
154+
}
155+
146156
// WithUserAgent sets an explicit User-Agent header, overriding the default and any
147157
// value set via WithHeaders.
148158
func WithUserAgent(ua string) Option {
@@ -194,11 +204,14 @@ func (o *provider) LanguageModel(_ context.Context, modelID string) (fantasy.Lan
194204
if objectMode == fantasy.ObjectModeJSON {
195205
objectMode = fantasy.ObjectModeAuto
196206
}
197-
return newResponsesLanguageModel(modelID, o.options.name, client, objectMode), nil
207+
return newResponsesLanguageModel(modelID, o.options.name, client, objectMode, o.options.reasoningModelFunc), nil
198208
}
199209

200210
languageModelOptions := append([]LanguageModelOption{}, o.options.languageModelOptions...)
201211
languageModelOptions = append(languageModelOptions, WithLanguageModelObjectMode(o.options.objectMode))
212+
if o.options.reasoningModelFunc != nil {
213+
languageModelOptions = append(languageModelOptions, WithLanguageModelReasoningModelFunc(o.options.reasoningModelFunc))
214+
}
202215

203216
return newLanguageModel(
204217
modelID,

‎providers/openai/responses_language_model.go‎

Lines changed: 23 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -24,19 +24,21 @@ import (
2424
const topLogprobsMax = 20
2525

2626
type responsesLanguageModel struct {
27-
provider string
28-
modelID string
29-
client openai.Client
30-
objectMode fantasy.ObjectMode
27+
provider string
28+
modelID string
29+
client openai.Client
30+
objectMode fantasy.ObjectMode
31+
reasoningModelFunc func(modelID string) bool
3132
}
3233

3334
// newResponsesLanguageModel implements a responses api model.
34-
func newResponsesLanguageModel(modelID string, provider string, client openai.Client, objectMode fantasy.ObjectMode) responsesLanguageModel {
35+
func newResponsesLanguageModel(modelID string, provider string, client openai.Client, objectMode fantasy.ObjectMode, reasoningModelFunc func(modelID string) bool) responsesLanguageModel {
3536
return responsesLanguageModel{
36-
modelID: modelID,
37-
provider: provider,
38-
client: client,
39-
objectMode: objectMode,
37+
modelID: modelID,
38+
provider: provider,
39+
client: client,
40+
objectMode: objectMode,
41+
reasoningModelFunc: reasoningModelFunc,
4042
}
4143
}
4244

@@ -77,7 +79,8 @@ func getResponsesModelConfig(modelID string) responsesModelConfig {
7779
supportsPriorityProcessing: supportsPriorityProcessing,
7880
}
7981

80-
if strings.Contains(strings.ToLower(modelID), "gpt-5-chat") {
82+
reasoningGeneration := reasoningGenerationPattern.MatchString(strings.ToLower(modelID))
83+
if reasoningGeneration && strings.Contains(strings.ToLower(modelID), "-chat") {
8184
return responsesModelConfig{
8285
isReasoningModel: false,
8386
systemMessageMode: defaults.systemMessageMode,
@@ -91,7 +94,7 @@ func getResponsesModelConfig(modelID string) responsesModelConfig {
9194
strings.HasPrefix(modelID, "o3") || strings.Contains(modelID, "-o3") ||
9295
strings.HasPrefix(modelID, "o4") || strings.Contains(modelID, "-o4") ||
9396
strings.HasPrefix(modelID, "oss") || strings.Contains(modelID, "-oss") ||
94-
strings.Contains(strings.ToLower(modelID), "gpt-5") ||
97+
reasoningGeneration ||
9598
strings.Contains(modelID, "codex-") || strings.Contains(modelID, "computer-use") {
9699
if strings.Contains(modelID, "o1-mini") || strings.Contains(modelID, "o1-preview") {
97100
return responsesModelConfig{
@@ -131,6 +134,15 @@ func (o responsesLanguageModel) prepareParams(call fantasy.Call) (*responses.Res
131134
params := &responses.ResponseNewParams{}
132135

133136
modelConfig := getResponsesModelConfig(o.modelID)
137+
if o.reasoningModelFunc != nil {
138+
if reasoning := o.reasoningModelFunc(o.modelID); reasoning != modelConfig.isReasoningModel {
139+
modelConfig.isReasoningModel = reasoning
140+
modelConfig.systemMessageMode = "system"
141+
if reasoning {
142+
modelConfig.systemMessageMode = "developer"
143+
}
144+
}
145+
}
134146

135147
if call.TopK != nil {
136148
warnings = append(warnings, fantasy.CallWarning{

‎providers/openai/responses_options.go‎

Lines changed: 13 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@ package openai
33

44
import (
55
"encoding/json"
6+
"regexp"
67
"slices"
78
"strings"
89

@@ -296,18 +297,26 @@ func ParseResponsesOptions(data map[string]any) (*ResponsesProviderOptions, erro
296297
return &options, nil
297298
}
298299

300+
// responsesGenerationPattern matches the model generations that only
301+
// speak the Responses API: gpt-4 and gpt-5 today, the newer generations
302+
// as they ship (gpt-6, gpt-10, ...), and never the legacy gpt-3 family
303+
// that predates it.
304+
var responsesGenerationPattern = regexp.MustCompile(`gpt-(?:[4-9]|[1-9]\d)`)
305+
306+
// reasoningGenerationPattern is the subset of those generations that reason:
307+
// gpt-5 and everything after it, never gpt-4.
308+
var reasoningGenerationPattern = regexp.MustCompile(`gpt-(?:[5-9]|[1-9]\d)`)
309+
299310
// IsResponsesModel checks if a model ID is a Responses API model for OpenAI.
300311
func IsResponsesModel(modelID string) bool {
301312
return slices.Contains(responsesModelIDs, modelID) ||
302-
strings.Contains(strings.ToLower(modelID), "gpt-4") ||
303-
strings.Contains(strings.ToLower(modelID), "gpt-5")
313+
responsesGenerationPattern.MatchString(strings.ToLower(modelID))
304314
}
305315

306316
// IsResponsesReasoningModel checks if a model ID is a Responses API reasoning model for OpenAI.
307317
func IsResponsesReasoningModel(modelID string) bool {
308318
return slices.Contains(responsesReasoningModelIDs, modelID) ||
309-
strings.Contains(strings.ToLower(modelID), "gpt-4") ||
310-
strings.Contains(strings.ToLower(modelID), "gpt-5")
319+
responsesGenerationPattern.MatchString(strings.ToLower(modelID))
311320
}
312321

313322
// SearchContextSize controls how much context window space the
Lines changed: 60 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,60 @@
1+
package openai
2+
3+
import (
4+
"testing"
5+
6+
"github.com/stretchr/testify/assert"
7+
)
8+
9+
func TestIsResponsesModel(t *testing.T) {
10+
tests := []struct {
11+
modelID string
12+
want bool
13+
}{
14+
// Explicitly listed models.
15+
{"gpt-4.1", true},
16+
{"gpt-4o-mini", true},
17+
{"chatgpt-4o-latest", true},
18+
{"o3", true},
19+
{"gpt-oss-120b", true},
20+
21+
// Generations caught by the pattern, listed or not.
22+
{"gpt-5", true},
23+
{"gpt-5.1-codex", true},
24+
{"gpt-6-astra", true},
25+
{"GPT-6-ASTRA", true},
26+
{"gpt-10-turbo", true},
27+
28+
// Everything predating the Responses API stays on chat
29+
// completions.
30+
{"gpt-3.5-turbo-1106", true}, // in the explicit list
31+
{"gpt-3-turbo-instruct", false},
32+
{"babbage-002", false},
33+
{"davinci-002", false},
34+
{"some-custom-model", false},
35+
}
36+
37+
for _, tt := range tests {
38+
assert.Equal(t, tt.want, IsResponsesModel(tt.modelID), tt.modelID)
39+
}
40+
}
41+
42+
func TestIsResponsesReasoningModel(t *testing.T) {
43+
tests := []struct {
44+
modelID string
45+
want bool
46+
}{
47+
{"gpt-5.1-codex", true},
48+
{"gpt-6-astra", true},
49+
{"o4-mini", true},
50+
{"gpt-oss-120b", true},
51+
52+
{"gpt-4.1-mini", true}, // gpt-4 matches, as before
53+
{"gpt-3-turbo-instruct", false},
54+
{"some-custom-model", false},
55+
}
56+
57+
for _, tt := range tests {
58+
assert.Equal(t, tt.want, IsResponsesReasoningModel(tt.modelID), tt.modelID)
59+
}
60+
}

0 commit comments

Comments
 (0)