mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: add configurable reasoning effort to Coder agents (#26974)
This commit is contained in:
@@ -825,6 +825,12 @@ The generation goroutine supports:
|
||||
- turn limit after a user message (the LLM shouldn't be able to spin forever in loop)
|
||||
- and other things
|
||||
|
||||
##### Reasoning effort
|
||||
|
||||
Model configs may carry a `reasoning_effort` config (`{default, max}`) inside `chat_model_configs.options`. Users select a per-turn effort when sending or editing a message; the value is stored on `chat_messages.reasoning_effort` and on `chat_queued_messages.reasoning_effort` for queued messages. Queued messages carry the value through promotion, and `chats.last_reasoning_effort` tracks the most recent message that set one, mirroring `last_model_config_id`.
|
||||
|
||||
During generation preparation, the effective effort is resolved as the chat's `last_reasoning_effort` if set, else the config's `default`; clamped to the config's `max` on the global scale `none < minimal < low < medium < high < xhigh < max`; and passed through to the provider. The provider verifies whether the configured value is valid for that model at runtime. If the model config has no `reasoning_effort`, any user-selected value is ignored. The resolved value is injected into the provider-native options with `chatprovider.ApplyReasoningEffort` after provider option conversion.
|
||||
|
||||
#### Interrupt goroutine
|
||||
|
||||
The interrupt goroutine is responsible for handling interrupts. It is spawned when the event indicates the core state machine is in `I0` or `I1` (status is `interrupting`).
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"time"
|
||||
|
||||
"charm.land/fantasy"
|
||||
fantasyopenai "charm.land/fantasy/providers/openai"
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/xerrors"
|
||||
@@ -16,6 +17,7 @@ import (
|
||||
"cdr.dev/slog/v3"
|
||||
"github.com/coder/coder/v2/coderd/aibridge"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/util/ptr"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatadvisor"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattest"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
@@ -608,4 +610,39 @@ func TestNewAdvisorRuntime(t *testing.T) {
|
||||
require.Equal(t, int64(defaultAdvisorMaxOutputTokens), rt.MaxOutputTokens(),
|
||||
"zero max output tokens must be replaced with defaultAdvisorMaxOutputTokens")
|
||||
})
|
||||
|
||||
t.Run("AppliesReasoningEffortToProviderOptions", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
store := &advisorOverrideStubStore{}
|
||||
p := newAdvisorTestServer(ctx, t, store)
|
||||
|
||||
rt := p.newAdvisorRuntimeOrFallback(
|
||||
ctx,
|
||||
database.Chat{},
|
||||
codersdk.AdvisorConfig{
|
||||
Enabled: true,
|
||||
MaxUsesPerRun: 3,
|
||||
MaxOutputTokens: 16384,
|
||||
},
|
||||
fallbackModel,
|
||||
codersdk.ChatModelCallConfig{
|
||||
ReasoningEffort: &codersdk.ChatModelReasoningEffortConfig{
|
||||
Default: ptr.Ref(codersdk.ChatModelReasoningEffortHigh),
|
||||
Max: ptr.Ref(codersdk.ChatModelReasoningEffortXHigh),
|
||||
},
|
||||
ProviderOptions: &codersdk.ChatModelProviderOptions{
|
||||
OpenAI: &codersdk.ChatModelOpenAIProviderOptions{
|
||||
User: ptr.Ref("advisor-user"),
|
||||
},
|
||||
},
|
||||
},
|
||||
modelBuildOptions{},
|
||||
logger,
|
||||
)
|
||||
require.NotNil(t, rt)
|
||||
providerOptions := rt.ProviderOptions()[fantasyopenai.Name].(*fantasyopenai.ResponsesProviderOptions)
|
||||
require.Equal(t, "advisor-user", *providerOptions.User)
|
||||
require.Equal(t, fantasyopenai.ReasoningEffortHigh, *providerOptions.ReasoningEffort)
|
||||
})
|
||||
}
|
||||
|
||||
+45
-17
@@ -398,10 +398,21 @@ func (p *Server) newAdvisorRuntime(
|
||||
}
|
||||
|
||||
advisorCallConfig.MaxOutputTokens = ptr.Ref(maxOutputTokens)
|
||||
// The advisor has no per-turn effort selection; its model config's
|
||||
// default effort applies.
|
||||
advisorReasoningEffort := chatprovider.ResolveReasoningEffort(
|
||||
nil,
|
||||
advisorCallConfig.ReasoningEffort,
|
||||
)
|
||||
providerOptions := chatprovider.ProviderOptionsFromChatModelConfig(
|
||||
advisorModel,
|
||||
advisorCallConfig.ProviderOptions,
|
||||
)
|
||||
providerOptions = chatprovider.ApplyReasoningEffort(
|
||||
advisorModel,
|
||||
providerOptions,
|
||||
advisorReasoningEffort,
|
||||
)
|
||||
|
||||
rt, err := chatadvisor.NewRuntime(chatadvisor.RuntimeConfig{
|
||||
Model: advisorModel,
|
||||
@@ -1131,6 +1142,7 @@ type CreateOptions struct {
|
||||
RootChatID uuid.NullUUID
|
||||
Title string
|
||||
ModelConfigID uuid.UUID
|
||||
ReasoningEffort *string
|
||||
ChatMode database.NullChatMode
|
||||
PlanMode database.NullChatPlanMode
|
||||
ClientType database.ChatClientType
|
||||
@@ -1157,14 +1169,15 @@ const (
|
||||
|
||||
// SendMessageOptions controls user message insertion with busy-state behavior.
|
||||
type SendMessageOptions struct {
|
||||
ChatID uuid.UUID
|
||||
CreatedBy uuid.UUID
|
||||
Content []codersdk.ChatMessagePart
|
||||
ModelConfigID uuid.UUID
|
||||
APIKeyID string
|
||||
BusyBehavior SendMessageBusyBehavior
|
||||
PlanMode *database.NullChatPlanMode
|
||||
MCPServerIDs *[]uuid.UUID
|
||||
ChatID uuid.UUID
|
||||
CreatedBy uuid.UUID
|
||||
Content []codersdk.ChatMessagePart
|
||||
ModelConfigID uuid.UUID
|
||||
ReasoningEffort *string
|
||||
APIKeyID string
|
||||
BusyBehavior SendMessageBusyBehavior
|
||||
PlanMode *database.NullChatPlanMode
|
||||
MCPServerIDs *[]uuid.UUID
|
||||
}
|
||||
|
||||
// SendMessageResult contains the outcome of user message processing.
|
||||
@@ -1185,7 +1198,8 @@ type EditMessageOptions struct {
|
||||
// ModelConfigID, when non-zero, overrides the model used for
|
||||
// the replacement user message. When set to uuid.Nil the
|
||||
// original message's model is preserved.
|
||||
ModelConfigID uuid.UUID
|
||||
ModelConfigID uuid.UUID
|
||||
ReasoningEffort *string
|
||||
}
|
||||
|
||||
// EditMessageResult contains the replacement user message and chat status.
|
||||
@@ -1299,7 +1313,7 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C
|
||||
initialMessages = append(initialMessages, systemMessage(userPromptContent, opts.ModelConfigID))
|
||||
}
|
||||
initialMessages = append(initialMessages, systemMessage(workspaceAwarenessContent, opts.ModelConfigID))
|
||||
initialMessages = append(initialMessages, userMessageWithAPIKeyID(userContent, opts.ModelConfigID, opts.OwnerID, opts.APIKeyID))
|
||||
initialMessages = append(initialMessages, userMessageWithAPIKeyID(userContent, opts.ModelConfigID, opts.OwnerID, opts.APIKeyID, opts.ReasoningEffort))
|
||||
|
||||
result, err := chatstate.CreateChat(ctx, p.db, p.pubsub, chatstate.CreateChatInput{
|
||||
OrganizationID: opts.OrganizationID,
|
||||
@@ -1446,7 +1460,7 @@ func (p *Server) SendMessage(
|
||||
// Queue capacity is enforced inside tx.SendMessage; this
|
||||
// wrapper only propagates the typed error.
|
||||
sendResult, err := tx.SendMessage(chatstate.SendMessageInput{
|
||||
Message: userMessageWithAPIKeyID(content, modelConfigID, messageCreatedBy, opts.APIKeyID),
|
||||
Message: userMessageWithAPIKeyID(content, modelConfigID, messageCreatedBy, opts.APIKeyID, opts.ReasoningEffort),
|
||||
BusyBehavior: busyBehaviorToChatState(busyBehavior),
|
||||
})
|
||||
if err != nil {
|
||||
@@ -1654,12 +1668,18 @@ func (p *Server) EditMessage(
|
||||
modelOverride = uuid.NullUUID{UUID: opts.ModelConfigID, Valid: true}
|
||||
}
|
||||
|
||||
var reasoningEffortOverride database.NullChatReasoningEffort
|
||||
if opts.ReasoningEffort != nil && *opts.ReasoningEffort != "" {
|
||||
reasoningEffortOverride = database.NullChatReasoningEffort{ChatReasoningEffort: database.ChatReasoningEffort(*opts.ReasoningEffort), Valid: true}
|
||||
}
|
||||
|
||||
editResult, err := tx.EditMessage(chatstate.EditMessageInput{
|
||||
MessageID: opts.EditedMessageID,
|
||||
CreatedBy: opts.CreatedBy,
|
||||
Content: content,
|
||||
ModelConfigIDOverride: modelOverride,
|
||||
APIKeyID: sql.NullString{String: opts.APIKeyID, Valid: opts.APIKeyID != ""},
|
||||
MessageID: opts.EditedMessageID,
|
||||
CreatedBy: opts.CreatedBy,
|
||||
Content: content,
|
||||
ModelConfigIDOverride: modelOverride,
|
||||
ReasoningEffortOverride: reasoningEffortOverride,
|
||||
APIKeyID: sql.NullString{String: opts.APIKeyID, Valid: opts.APIKeyID != ""},
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, chatstate.ErrEditedMessageNotUser) {
|
||||
@@ -2340,7 +2360,13 @@ func (p *Server) generateManualTitleCandidate(
|
||||
)
|
||||
}
|
||||
|
||||
title, usage, err := generateManualTitle(titleCtx, messages, pasteText, titleModel)
|
||||
title, usage, err := generateManualTitle(
|
||||
titleCtx,
|
||||
messages,
|
||||
pasteText,
|
||||
titleModel,
|
||||
p.titleGenerationProviderOptions(ctx, titleModel, modelConfig),
|
||||
)
|
||||
finishDebugRun(err)
|
||||
result.title = title
|
||||
result.usage = usage
|
||||
@@ -2803,6 +2829,7 @@ func recordManualTitleUsage(
|
||||
CreatedBy: []uuid.UUID{chat.OwnerID},
|
||||
APIKeyID: []string{activeAPIKeyID},
|
||||
ModelConfigID: []uuid.UUID{modelConfig.ID},
|
||||
ReasoningEffort: []string{""},
|
||||
Role: []database.ChatMessageRole{database.ChatMessageRoleAssistant},
|
||||
Content: []string{content},
|
||||
ContentVersion: []int16{chatprompt.CurrentContentVersion},
|
||||
@@ -2931,6 +2958,7 @@ func appendMessageFields(
|
||||
params.CreatedBy = append(params.CreatedBy, msg.createdBy)
|
||||
params.APIKeyID = append(params.APIKeyID, apiKeyID)
|
||||
params.ModelConfigID = append(params.ModelConfigID, msg.modelConfigID)
|
||||
params.ReasoningEffort = append(params.ReasoningEffort, "")
|
||||
params.Role = append(params.Role, msg.role)
|
||||
params.Content = append(params.Content, string(msg.content.RawMessage))
|
||||
params.ContentVersion = append(params.ContentVersion, msg.contentVersion)
|
||||
|
||||
@@ -46,6 +46,7 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/database/dbgen"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
dbpubsub "github.com/coder/coder/v2/coderd/database/pubsub"
|
||||
"github.com/coder/coder/v2/coderd/util/ptr"
|
||||
"github.com/coder/coder/v2/coderd/util/slice"
|
||||
"github.com/coder/coder/v2/coderd/workspacestats"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd"
|
||||
@@ -78,9 +79,9 @@ func testAPIKeyID(t testing.TB, db database.Store, userID uuid.UUID) string {
|
||||
}
|
||||
|
||||
func chatAIGatewayTransportFactoryPointer(factory aibridge.TransportFactory) *atomic.Pointer[aibridge.TransportFactory] {
|
||||
var ptr atomic.Pointer[aibridge.TransportFactory]
|
||||
ptr.Store(&factory)
|
||||
return &ptr
|
||||
var factoryPtr atomic.Pointer[aibridge.TransportFactory]
|
||||
factoryPtr.Store(&factory)
|
||||
return &factoryPtr
|
||||
}
|
||||
|
||||
func openAIToolName(tool chattest.OpenAITool) string {
|
||||
@@ -11663,6 +11664,65 @@ func TestEditMessagePreservesModelConfigByDefault(t *testing.T) {
|
||||
"edit without model override must not change last_model_config_id")
|
||||
}
|
||||
|
||||
func TestEditMessageReasoningEffort(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
requested *string
|
||||
want string
|
||||
}{
|
||||
{name: "PreservesByDefault", want: "low"},
|
||||
{name: "Overrides", requested: ptr.Ref("high"), want: "high"},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
replica := newTestServer(t, db, ps, uuid.New())
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
|
||||
chat, err := replica.CreateChat(ctx, chatd.CreateOptions{
|
||||
OwnerID: user.ID,
|
||||
APIKeyID: testAPIKeyID(t, db, user.ID),
|
||||
OrganizationID: org.ID,
|
||||
Title: "edit-reasoning-effort",
|
||||
ModelConfigID: model.ID,
|
||||
ReasoningEffort: ptr.Ref("low"),
|
||||
InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("original")},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
initial, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{
|
||||
ChatID: chat.ID,
|
||||
AfterID: 0,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Len(t, initial, 1)
|
||||
require.True(t, initial[0].ReasoningEffort.Valid)
|
||||
require.Equal(t, database.ChatReasoningEffortLow, initial[0].ReasoningEffort.ChatReasoningEffort)
|
||||
|
||||
result, err := replica.EditMessage(ctx, chatd.EditMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
APIKeyID: testAPIKeyID(t, db, user.ID),
|
||||
EditedMessageID: initial[0].ID,
|
||||
Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edited")},
|
||||
ReasoningEffort: tc.requested,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.True(t, result.Message.ReasoningEffort.Valid)
|
||||
require.Equal(t, database.ChatReasoningEffort(tc.want), result.Message.ReasoningEffort.ChatReasoningEffort)
|
||||
|
||||
storedChat, err := db.GetChatByID(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
require.True(t, storedChat.LastReasoningEffort.Valid)
|
||||
require.Equal(t, database.ChatReasoningEffort(tc.want), storedChat.LastReasoningEffort.ChatReasoningEffort)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestEditMessageRejectsUnknownModelConfig verifies the edit handler
|
||||
// returns ErrInvalidModelConfigID when the requested model does not
|
||||
// exist, mirroring SendMessage's validation.
|
||||
@@ -11716,6 +11776,51 @@ func TestEditMessageRejectsUnknownModelConfig(t *testing.T) {
|
||||
require.Equal(t, modelA.ID, storedChat.LastModelConfigID)
|
||||
}
|
||||
|
||||
func TestPromoteQueuedPreservesReasoningEffort(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
replica := newTestServer(t, db, ps, uuid.New())
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
user, org, model := seedChatDependencies(t, db)
|
||||
|
||||
chat := dbgen.Chat(t, db, database.Chat{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: model.ID,
|
||||
Title: "promote-reasoning-effort",
|
||||
Status: database.ChatStatusError,
|
||||
})
|
||||
content, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{codersdk.ChatMessageText("queued")})
|
||||
require.NoError(t, err)
|
||||
queued, err := db.InsertChatQueuedMessageWithCreator(ctx, database.InsertChatQueuedMessageWithCreatorParams{
|
||||
ChatID: chat.ID,
|
||||
Content: content.RawMessage,
|
||||
ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true},
|
||||
ReasoningEffort: database.NullChatReasoningEffort{ChatReasoningEffort: database.ChatReasoningEffortHigh, Valid: true},
|
||||
APIKeyID: sql.NullString{String: testAPIKeyID(t, db, user.ID), Valid: true},
|
||||
CreatedBy: user.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.True(t, queued.ReasoningEffort.Valid)
|
||||
require.Equal(t, database.ChatReasoningEffortHigh, queued.ReasoningEffort.ChatReasoningEffort)
|
||||
|
||||
result, err := replica.PromoteQueued(ctx, chatd.PromoteQueuedOptions{
|
||||
ChatID: chat.ID,
|
||||
CreatedBy: user.ID,
|
||||
QueuedMessageID: queued.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.True(t, result.PromotedMessage.ReasoningEffort.Valid)
|
||||
require.Equal(t, database.ChatReasoningEffortHigh, result.PromotedMessage.ReasoningEffort.ChatReasoningEffort)
|
||||
|
||||
storedChat, err := db.GetChatByID(ctx, chat.ID)
|
||||
require.NoError(t, err)
|
||||
require.True(t, storedChat.LastReasoningEffort.Valid)
|
||||
require.Equal(t, database.ChatReasoningEffortHigh, storedChat.LastReasoningEffort.ChatReasoningEffort)
|
||||
}
|
||||
|
||||
// TestPromoteQueuedWhileRequiresActionMixedTools guards against
|
||||
func TestAcquireChatsSkipsArchivedPendingChat(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -18,7 +18,6 @@ func ProviderOptionsFromChatConfig(
|
||||
model fantasy.LanguageModel,
|
||||
options *codersdk.ChatModelOpenAIProviderOptions,
|
||||
) fantasy.ProviderOptionsData {
|
||||
reasoningEffort := ReasoningEffortFromChat(options.ReasoningEffort)
|
||||
if UsesResponsesOptions(model) {
|
||||
include := EnsureResponseIncludes(IncludeFromChat(options.Include))
|
||||
providerOptions := &fantasyopenai.ResponsesProviderOptions{
|
||||
@@ -29,7 +28,6 @@ func ProviderOptionsFromChatConfig(
|
||||
Metadata: options.Metadata,
|
||||
ParallelToolCalls: options.ParallelToolCalls,
|
||||
PromptCacheKey: chatutil.NormalizedStringPointer(options.PromptCacheKey),
|
||||
ReasoningEffort: reasoningEffort,
|
||||
ReasoningSummary: chatutil.NormalizedStringPointer(options.ReasoningSummary),
|
||||
SafetyIdentifier: chatutil.NormalizedStringPointer(options.SafetyIdentifier),
|
||||
ServiceTier: ServiceTierFromChat(options.ServiceTier),
|
||||
@@ -47,7 +45,6 @@ func ProviderOptionsFromChatConfig(
|
||||
TopLogProbs: options.TopLogProbs,
|
||||
ParallelToolCalls: options.ParallelToolCalls,
|
||||
User: chatutil.NormalizedStringPointer(options.User),
|
||||
ReasoningEffort: reasoningEffort,
|
||||
MaxCompletionTokens: options.MaxCompletionTokens,
|
||||
TextVerbosity: chatutil.NormalizedStringPointer(options.TextVerbosity),
|
||||
Prediction: options.Prediction,
|
||||
@@ -133,33 +130,6 @@ func UsesResponsesOptions(model fantasy.LanguageModel) bool {
|
||||
}
|
||||
}
|
||||
|
||||
// ReasoningEffortFromChat normalizes chat-config reasoning effort values for
|
||||
// OpenAI and returns the canonical provider effort value.
|
||||
func ReasoningEffortFromChat(value *string) *fantasyopenai.ReasoningEffort {
|
||||
if value == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
normalized := strings.ToLower(strings.TrimSpace(*value))
|
||||
if normalized == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
effort := chatutil.NormalizedEnumValue(
|
||||
normalized,
|
||||
string(fantasyopenai.ReasoningEffortMinimal),
|
||||
string(fantasyopenai.ReasoningEffortLow),
|
||||
string(fantasyopenai.ReasoningEffortMedium),
|
||||
string(fantasyopenai.ReasoningEffortHigh),
|
||||
string(fantasyopenai.ReasoningEffortXHigh),
|
||||
)
|
||||
if effort == nil {
|
||||
return nil
|
||||
}
|
||||
valueCopy := fantasyopenai.ReasoningEffort(*effort)
|
||||
return &valueCopy
|
||||
}
|
||||
|
||||
// ServiceTierFromChat normalizes chat-config service tier values for OpenAI
|
||||
// Responses API and returns the canonical provider service tier value.
|
||||
func ServiceTierFromChat(value *string) *fantasyopenai.ServiceTier {
|
||||
|
||||
@@ -30,7 +30,6 @@ func TestProviderOptionsFromChatConfigLegacy(t *testing.T) {
|
||||
TopLogProbs: &topLogProbs,
|
||||
ParallelToolCalls: ¶llelToolCalls,
|
||||
User: ptr(" user-1 "),
|
||||
ReasoningEffort: ptr(" HIGH "),
|
||||
MaxCompletionTokens: &maxCompletionTokens,
|
||||
TextVerbosity: ptr(" High "),
|
||||
Prediction: map[string]any{
|
||||
@@ -56,7 +55,7 @@ func TestProviderOptionsFromChatConfigLegacy(t *testing.T) {
|
||||
require.Same(t, options.TopLogProbs, providerOptions.TopLogProbs)
|
||||
require.Same(t, options.ParallelToolCalls, providerOptions.ParallelToolCalls)
|
||||
require.Equal(t, "user-1", requireStringPointerValue(t, providerOptions.User))
|
||||
require.Equal(t, fantasyopenai.ReasoningEffortHigh, requireReasoningEffortPointerValue(t, providerOptions.ReasoningEffort))
|
||||
require.Nil(t, providerOptions.ReasoningEffort)
|
||||
require.Same(t, options.MaxCompletionTokens, providerOptions.MaxCompletionTokens)
|
||||
require.Equal(t, "High", requireStringPointerValue(t, providerOptions.TextVerbosity))
|
||||
require.Equal(t, options.Prediction, providerOptions.Prediction)
|
||||
@@ -88,7 +87,6 @@ func TestProviderOptionsFromChatConfigResponses(t *testing.T) {
|
||||
Metadata: map[string]any{"scope": "unit"},
|
||||
ParallelToolCalls: ¶llelToolCalls,
|
||||
PromptCacheKey: ptr(" prompt-cache "),
|
||||
ReasoningEffort: ptr(" minimal "),
|
||||
ReasoningSummary: ptr(" auto "),
|
||||
SafetyIdentifier: ptr(" safety "),
|
||||
ServiceTier: ptr(" FLEX "),
|
||||
@@ -114,7 +112,7 @@ func TestProviderOptionsFromChatConfigResponses(t *testing.T) {
|
||||
require.Equal(t, options.Metadata, providerOptions.Metadata)
|
||||
require.Same(t, options.ParallelToolCalls, providerOptions.ParallelToolCalls)
|
||||
require.Equal(t, "prompt-cache", requireStringPointerValue(t, providerOptions.PromptCacheKey))
|
||||
require.Equal(t, fantasyopenai.ReasoningEffortMinimal, requireReasoningEffortPointerValue(t, providerOptions.ReasoningEffort))
|
||||
require.Nil(t, providerOptions.ReasoningEffort)
|
||||
require.Equal(t, "auto", requireStringPointerValue(t, providerOptions.ReasoningSummary))
|
||||
require.Equal(t, "safety", requireStringPointerValue(t, providerOptions.SafetyIdentifier))
|
||||
require.Equal(t, fantasyopenai.ServiceTierFlex, requireServiceTierPointerValue(t, providerOptions.ServiceTier))
|
||||
@@ -281,40 +279,6 @@ func TestUsesResponsesOptions(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestReasoningEffortFromChat(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
value *string
|
||||
want *fantasyopenai.ReasoningEffort
|
||||
}{
|
||||
{name: "Nil"},
|
||||
{name: "Empty", value: ptr(" ")},
|
||||
{name: "Minimal", value: ptr(" minimal "), want: ptr(fantasyopenai.ReasoningEffortMinimal)},
|
||||
{name: "LowCase", value: ptr(" LOW "), want: ptr(fantasyopenai.ReasoningEffortLow)},
|
||||
{name: "Medium", value: ptr("medium"), want: ptr(fantasyopenai.ReasoningEffortMedium)},
|
||||
{name: "High", value: ptr("high"), want: ptr(fantasyopenai.ReasoningEffortHigh)},
|
||||
{name: "XHigh", value: ptr("xhigh"), want: ptr(fantasyopenai.ReasoningEffortXHigh)},
|
||||
{name: "NoneUnsupported", value: ptr("none")},
|
||||
{name: "Invalid", value: ptr("max")},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := chatopenai.ReasoningEffortFromChat(tt.value)
|
||||
if tt.want == nil {
|
||||
require.Nil(t, got)
|
||||
return
|
||||
}
|
||||
require.NotNil(t, got)
|
||||
require.Equal(t, *tt.want, *got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestServiceTierFromChat(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -438,15 +402,6 @@ func requireBoolPointerValue(t *testing.T, value *bool) bool {
|
||||
return *value
|
||||
}
|
||||
|
||||
func requireReasoningEffortPointerValue(
|
||||
t *testing.T,
|
||||
value *fantasyopenai.ReasoningEffort,
|
||||
) fantasyopenai.ReasoningEffort {
|
||||
t.Helper()
|
||||
require.NotNil(t, value)
|
||||
return *value
|
||||
}
|
||||
|
||||
func requireServiceTierPointerValue(
|
||||
t *testing.T,
|
||||
value *fantasyopenai.ServiceTier,
|
||||
|
||||
@@ -812,57 +812,6 @@ func isChatModelForProvider(provider, modelID string) bool {
|
||||
}
|
||||
}
|
||||
|
||||
// ReasoningEffortFromChat normalizes chat-config reasoning effort values for a
|
||||
// provider and returns the canonical provider effort value.
|
||||
func ReasoningEffortFromChat(provider string, value *string) *string {
|
||||
if value == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
normalized := strings.ToLower(strings.TrimSpace(*value))
|
||||
if normalized == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
switch NormalizeProvider(provider) {
|
||||
case fantasyopenai.Name:
|
||||
effort := chatopenai.ReasoningEffortFromChat(value)
|
||||
if effort == nil {
|
||||
return nil
|
||||
}
|
||||
valueCopy := string(*effort)
|
||||
return &valueCopy
|
||||
case fantasyanthropic.Name:
|
||||
return chatutil.NormalizedEnumValue(
|
||||
normalized,
|
||||
string(fantasyanthropic.EffortLow),
|
||||
string(fantasyanthropic.EffortMedium),
|
||||
string(fantasyanthropic.EffortHigh),
|
||||
string(fantasyanthropic.EffortXHigh),
|
||||
string(fantasyanthropic.EffortMax),
|
||||
)
|
||||
case fantasyopenrouter.Name:
|
||||
return chatutil.NormalizedEnumValue(
|
||||
normalized,
|
||||
string(fantasyopenrouter.ReasoningEffortLow),
|
||||
string(fantasyopenrouter.ReasoningEffortMedium),
|
||||
string(fantasyopenrouter.ReasoningEffortHigh),
|
||||
)
|
||||
case fantasyvercel.Name:
|
||||
return chatutil.NormalizedEnumValue(
|
||||
normalized,
|
||||
string(fantasyvercel.ReasoningEffortNone),
|
||||
string(fantasyvercel.ReasoningEffortMinimal),
|
||||
string(fantasyvercel.ReasoningEffortLow),
|
||||
string(fantasyvercel.ReasoningEffortMedium),
|
||||
string(fantasyvercel.ReasoningEffortHigh),
|
||||
string(fantasyvercel.ReasoningEffortXHigh),
|
||||
)
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// AnthropicThinkingDisplayFromChat normalizes chat-config thinking display
|
||||
// values for Anthropic and returns the canonical provider display value.
|
||||
func AnthropicThinkingDisplayFromChat(value *string) *fantasyanthropic.ThinkingDisplay {
|
||||
@@ -1176,7 +1125,6 @@ func anthropicProviderOptionsFromChatConfig(
|
||||
) *fantasyanthropic.ProviderOptions {
|
||||
result := &fantasyanthropic.ProviderOptions{
|
||||
SendReasoning: options.SendReasoning,
|
||||
Effort: anthropicEffortFromChat(options.Effort),
|
||||
ThinkingDisplay: AnthropicThinkingDisplayFromChat(options.ThinkingDisplay),
|
||||
DisableParallelToolUse: options.DisableParallelToolUse,
|
||||
}
|
||||
@@ -1221,8 +1169,7 @@ func openAICompatProviderOptionsFromChatConfig(
|
||||
options *codersdk.ChatModelOpenAICompatProviderOptions,
|
||||
) *fantasyopenaicompat.ProviderOptions {
|
||||
return &fantasyopenaicompat.ProviderOptions{
|
||||
User: chatutil.NormalizedStringPointer(options.User),
|
||||
ReasoningEffort: chatopenai.ReasoningEffortFromChat(options.ReasoningEffort),
|
||||
User: chatutil.NormalizedStringPointer(options.User),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1242,7 +1189,6 @@ func openRouterProviderOptionsFromChatConfig(
|
||||
Enabled: options.Reasoning.Enabled,
|
||||
Exclude: options.Reasoning.Exclude,
|
||||
MaxTokens: options.Reasoning.MaxTokens,
|
||||
Effort: openRouterReasoningEffortFromChat(options.Reasoning.Effort),
|
||||
}
|
||||
}
|
||||
if options.Provider != nil {
|
||||
@@ -1275,7 +1221,6 @@ func vercelProviderOptionsFromChatConfig(
|
||||
result.Reasoning = &fantasyvercel.ReasoningOptions{
|
||||
Enabled: options.Reasoning.Enabled,
|
||||
MaxTokens: options.Reasoning.MaxTokens,
|
||||
Effort: vercelReasoningEffortFromChat(options.Reasoning.Effort),
|
||||
Exclude: options.Reasoning.Exclude,
|
||||
}
|
||||
}
|
||||
@@ -1287,30 +1232,3 @@ func vercelProviderOptionsFromChatConfig(
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func anthropicEffortFromChat(value *string) *fantasyanthropic.Effort {
|
||||
effort := ReasoningEffortFromChat(fantasyanthropic.Name, value)
|
||||
if effort == nil {
|
||||
return nil
|
||||
}
|
||||
valueCopy := fantasyanthropic.Effort(*effort)
|
||||
return &valueCopy
|
||||
}
|
||||
|
||||
func openRouterReasoningEffortFromChat(value *string) *fantasyopenrouter.ReasoningEffort {
|
||||
effort := ReasoningEffortFromChat(fantasyopenrouter.Name, value)
|
||||
if effort == nil {
|
||||
return nil
|
||||
}
|
||||
valueCopy := fantasyopenrouter.ReasoningEffort(*effort)
|
||||
return &valueCopy
|
||||
}
|
||||
|
||||
func vercelReasoningEffortFromChat(value *string) *fantasyvercel.ReasoningEffort {
|
||||
effort := ReasoningEffortFromChat(fantasyvercel.Name, value)
|
||||
if effort == nil {
|
||||
return nil
|
||||
}
|
||||
valueCopy := fantasyvercel.ReasoningEffort(*effort)
|
||||
return &valueCopy
|
||||
}
|
||||
|
||||
@@ -349,81 +349,6 @@ func (fn roundTripperFunc) RoundTrip(req *http.Request) (*http.Response, error)
|
||||
return fn(req)
|
||||
}
|
||||
|
||||
func TestReasoningEffortFromChat(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
provider string
|
||||
input *string
|
||||
want *string
|
||||
}{
|
||||
{
|
||||
name: "OpenAICaseInsensitive",
|
||||
provider: "openai",
|
||||
input: ptr.Ref(" HIGH "),
|
||||
want: ptr.Ref(string(fantasyopenai.ReasoningEffortHigh)),
|
||||
},
|
||||
{
|
||||
name: "OpenAIXHighEffort",
|
||||
provider: "openai",
|
||||
input: ptr.Ref("xhigh"),
|
||||
want: ptr.Ref(string(fantasyopenai.ReasoningEffortXHigh)),
|
||||
},
|
||||
{
|
||||
name: "AnthropicEffort",
|
||||
provider: "anthropic",
|
||||
input: ptr.Ref("max"),
|
||||
want: ptr.Ref(string(fantasyanthropic.EffortMax)),
|
||||
},
|
||||
{
|
||||
name: "AnthropicXHighEffort",
|
||||
provider: "anthropic",
|
||||
input: ptr.Ref("xhigh"),
|
||||
want: ptr.Ref(string(fantasyanthropic.EffortXHigh)),
|
||||
},
|
||||
{
|
||||
name: "OpenRouterEffort",
|
||||
provider: "openrouter",
|
||||
input: ptr.Ref("medium"),
|
||||
want: ptr.Ref(string(fantasyopenrouter.ReasoningEffortMedium)),
|
||||
},
|
||||
{
|
||||
name: "VercelEffort",
|
||||
provider: "vercel",
|
||||
input: ptr.Ref("xhigh"),
|
||||
want: ptr.Ref(string(fantasyvercel.ReasoningEffortXHigh)),
|
||||
},
|
||||
{
|
||||
name: "InvalidEffortReturnsNil",
|
||||
provider: "openai",
|
||||
input: ptr.Ref("unknown"),
|
||||
want: nil,
|
||||
},
|
||||
{
|
||||
name: "UnsupportedProviderReturnsNil",
|
||||
provider: "bedrock",
|
||||
input: ptr.Ref("high"),
|
||||
want: nil,
|
||||
},
|
||||
{
|
||||
name: "NilInputReturnsNil",
|
||||
provider: "openai",
|
||||
input: nil,
|
||||
want: nil,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := chatprovider.ReasoningEffortFromChat(tt.provider, tt.input)
|
||||
require.Equal(t, tt.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAnthropicThinkingDisplayFromChat(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -0,0 +1,158 @@
|
||||
package chatprovider
|
||||
|
||||
import (
|
||||
"slices"
|
||||
|
||||
"charm.land/fantasy"
|
||||
fantasyanthropic "charm.land/fantasy/providers/anthropic"
|
||||
fantasyazure "charm.land/fantasy/providers/azure"
|
||||
fantasybedrock "charm.land/fantasy/providers/bedrock"
|
||||
fantasyopenai "charm.land/fantasy/providers/openai"
|
||||
fantasyopenaicompat "charm.land/fantasy/providers/openaicompat"
|
||||
fantasyopenrouter "charm.land/fantasy/providers/openrouter"
|
||||
fantasyvercel "charm.land/fantasy/providers/vercel"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatopenai"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
|
||||
func reasoningEffortRank(value string) (int, bool) {
|
||||
rank := slices.Index(codersdk.ChatModelReasoningEffortValues(), value)
|
||||
return rank, rank >= 0
|
||||
}
|
||||
|
||||
func IsValidReasoningEffort(value string) bool {
|
||||
_, ok := reasoningEffortRank(value)
|
||||
return ok
|
||||
}
|
||||
|
||||
// ReasoningEffortLessOrEqual reports whether a is lower than or equal
|
||||
// to b on the global effort scale. Unknown values return false.
|
||||
func ReasoningEffortLessOrEqual(a, b string) bool {
|
||||
aRank, aOK := reasoningEffortRank(a)
|
||||
bRank, bOK := reasoningEffortRank(b)
|
||||
return aOK && bOK && aRank <= bRank
|
||||
}
|
||||
|
||||
// ResolveReasoningEffort computes the effective reasoning effort for a
|
||||
// generation. The requested per-turn value wins over the config's default,
|
||||
// and the result is clamped to the config's max on the global scale. Returns
|
||||
// nil when the model config has no reasoning effort configured, no usable
|
||||
// value remains, or the max is unknown.
|
||||
func ResolveReasoningEffort(
|
||||
requested *string,
|
||||
config *codersdk.ChatModelReasoningEffortConfig,
|
||||
) *string {
|
||||
if config == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
effective := requested
|
||||
var rank int
|
||||
var ok bool
|
||||
if effective != nil {
|
||||
rank, ok = reasoningEffortRank(*effective)
|
||||
}
|
||||
if !ok {
|
||||
effective = config.Default
|
||||
if effective != nil {
|
||||
rank, ok = reasoningEffortRank(*effective)
|
||||
}
|
||||
}
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
if config.Max != nil {
|
||||
maxRank, ok := reasoningEffortRank(*config.Max)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
if rank > maxRank {
|
||||
return config.Max
|
||||
}
|
||||
}
|
||||
return effective
|
||||
}
|
||||
|
||||
func SelectableReasoningEfforts(
|
||||
config *codersdk.ChatModelReasoningEffortConfig,
|
||||
) []string {
|
||||
if config == nil || config.Max == nil {
|
||||
return nil
|
||||
}
|
||||
maxRank, ok := reasoningEffortRank(*config.Max)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
values := codersdk.ChatModelReasoningEffortValues()
|
||||
return values[:maxRank+1]
|
||||
}
|
||||
|
||||
func ApplyReasoningEffort(
|
||||
model fantasy.LanguageModel,
|
||||
options fantasy.ProviderOptions,
|
||||
effort *string,
|
||||
) fantasy.ProviderOptions {
|
||||
if effort == nil || model == nil {
|
||||
return options
|
||||
}
|
||||
if options == nil {
|
||||
options = fantasy.ProviderOptions{}
|
||||
}
|
||||
|
||||
switch NormalizeProvider(model.Provider()) {
|
||||
case fantasyopenai.Name, fantasyazure.Name:
|
||||
providerEffort := fantasyopenai.ReasoningEffort(*effort)
|
||||
switch opts := options[fantasyopenai.Name].(type) {
|
||||
case *fantasyopenai.ResponsesProviderOptions:
|
||||
opts.ReasoningEffort = &providerEffort
|
||||
case *fantasyopenai.ProviderOptions:
|
||||
opts.ReasoningEffort = &providerEffort
|
||||
default:
|
||||
if chatopenai.UsesResponsesOptions(model) {
|
||||
options[fantasyopenai.Name] = &fantasyopenai.ResponsesProviderOptions{
|
||||
ReasoningEffort: &providerEffort,
|
||||
}
|
||||
return options
|
||||
}
|
||||
options[fantasyopenai.Name] = &fantasyopenai.ProviderOptions{
|
||||
ReasoningEffort: &providerEffort,
|
||||
}
|
||||
}
|
||||
case fantasyanthropic.Name, fantasybedrock.Name:
|
||||
providerEffort := fantasyanthropic.Effort(*effort)
|
||||
providerOptions := ensureProviderOptions[fantasyanthropic.ProviderOptions](options, fantasyanthropic.Name)
|
||||
providerOptions.Effort = &providerEffort
|
||||
case fantasyopenaicompat.Name:
|
||||
providerEffort := fantasyopenai.ReasoningEffort(*effort)
|
||||
providerOptions := ensureProviderOptions[fantasyopenaicompat.ProviderOptions](options, fantasyopenaicompat.Name)
|
||||
providerOptions.ReasoningEffort = &providerEffort
|
||||
case fantasyopenrouter.Name:
|
||||
providerEffort := fantasyopenrouter.ReasoningEffort(*effort)
|
||||
providerOptions := ensureProviderOptions[fantasyopenrouter.ProviderOptions](options, fantasyopenrouter.Name)
|
||||
if providerOptions.Reasoning == nil {
|
||||
providerOptions.Reasoning = &fantasyopenrouter.ReasoningOptions{}
|
||||
}
|
||||
providerOptions.Reasoning.Effort = &providerEffort
|
||||
case fantasyvercel.Name:
|
||||
providerEffort := fantasyvercel.ReasoningEffort(*effort)
|
||||
providerOptions := ensureProviderOptions[fantasyvercel.ProviderOptions](options, fantasyvercel.Name)
|
||||
if providerOptions.Reasoning == nil {
|
||||
providerOptions.Reasoning = &fantasyvercel.ReasoningOptions{}
|
||||
}
|
||||
providerOptions.Reasoning.Effort = &providerEffort
|
||||
}
|
||||
return options
|
||||
}
|
||||
|
||||
func ensureProviderOptions[T any, PT interface {
|
||||
*T
|
||||
fantasy.ProviderOptionsData
|
||||
}](options fantasy.ProviderOptions, name string) PT {
|
||||
providerOptions, _ := options[name].(PT)
|
||||
if providerOptions == nil {
|
||||
providerOptions = PT(new(T))
|
||||
options[name] = providerOptions
|
||||
}
|
||||
return providerOptions
|
||||
}
|
||||
@@ -0,0 +1,243 @@
|
||||
package chatprovider_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"charm.land/fantasy"
|
||||
fantasyanthropic "charm.land/fantasy/providers/anthropic"
|
||||
fantasyopenai "charm.land/fantasy/providers/openai"
|
||||
fantasyopenaicompat "charm.land/fantasy/providers/openaicompat"
|
||||
fantasyopenrouter "charm.land/fantasy/providers/openrouter"
|
||||
fantasyvercel "charm.land/fantasy/providers/vercel"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/util/ptr"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprovider"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattest"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
|
||||
func TestResolveReasoningEffort(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
requested *string
|
||||
config *codersdk.ChatModelReasoningEffortConfig
|
||||
want *string
|
||||
}{
|
||||
{name: "NilConfigIgnoresRequested", requested: new(codersdk.ChatModelReasoningEffortHigh)},
|
||||
{name: "DefaultUsedWhenNoRequested", config: effortConfig("medium", "high"), want: new(codersdk.ChatModelReasoningEffortMedium)},
|
||||
{name: "RequestedWinsOverDefault", requested: new(codersdk.ChatModelReasoningEffortHigh), config: effortConfig("medium", "high"), want: new(codersdk.ChatModelReasoningEffortHigh)},
|
||||
{name: "RequestedWinsWithoutMax", requested: new(codersdk.ChatModelReasoningEffortHigh), config: effortConfig("medium", ""), want: new(codersdk.ChatModelReasoningEffortHigh)},
|
||||
{name: "RequestedClampedToMax", requested: new(codersdk.ChatModelReasoningEffortXHigh), config: effortConfig("low", "medium"), want: new(codersdk.ChatModelReasoningEffortMedium)},
|
||||
{name: "DefaultClampedToMax", config: effortConfig("xhigh", "medium"), want: new(codersdk.ChatModelReasoningEffortMedium)},
|
||||
{name: "InvalidRequestedFallsBackToDefault", requested: ptr.Ref(" HIGH "), config: effortConfig("low", "high"), want: new(codersdk.ChatModelReasoningEffortLow)},
|
||||
{name: "InvalidMaxReturnsNil", requested: new(codersdk.ChatModelReasoningEffortMedium), config: effortConfig("low", " HIGH ")},
|
||||
{name: "EmptyConfigReturnsNil", config: &codersdk.ChatModelReasoningEffortConfig{}},
|
||||
{name: "MaxSupported", requested: new(codersdk.ChatModelReasoningEffortMax), config: effortConfig("medium", "max"), want: new(codersdk.ChatModelReasoningEffortMax)},
|
||||
{name: "NoneSupported", requested: new(codersdk.ChatModelReasoningEffortNone), config: effortConfig("medium", "xhigh"), want: new(codersdk.ChatModelReasoningEffortNone)},
|
||||
{name: "MaxOnlyConfigClampsRequested", requested: new(codersdk.ChatModelReasoningEffortXHigh), config: effortConfig("", "medium"), want: new(codersdk.ChatModelReasoningEffortMedium)},
|
||||
{name: "MaxOnlyConfigWithoutRequestedReturnsNil", config: effortConfig("", "medium")},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := chatprovider.ResolveReasoningEffort(tt.requested, tt.config)
|
||||
if tt.want == nil {
|
||||
require.Nil(t, got)
|
||||
return
|
||||
}
|
||||
require.NotNil(t, got)
|
||||
require.Equal(t, *tt.want, *got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectableReasoningEfforts(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
config *codersdk.ChatModelReasoningEffortConfig
|
||||
want []string
|
||||
}{
|
||||
{name: "NilConfig"},
|
||||
{name: "NoMax", config: effortConfig("medium", "")},
|
||||
{name: "UnknownMax", config: effortConfig("medium", " HIGH ")},
|
||||
{name: "ThroughMedium", config: effortConfig("low", "medium"), want: []string{"none", "minimal", "low", "medium"}},
|
||||
{name: "ThroughMax", config: effortConfig("medium", "max"), want: []string{"none", "minimal", "low", "medium", "high", "xhigh", "max"}},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
require.Equal(t, tt.want, chatprovider.SelectableReasoningEfforts(tt.config))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyReasoningEffort(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("CreatesOpenAIResponsesEntry", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := chatprovider.ApplyReasoningEffort(&chattest.FakeModel{ProviderName: fantasyopenai.Name, ModelName: "gpt-5"}, nil, new(codersdk.ChatModelReasoningEffortHigh))
|
||||
providerOptions, ok := got[fantasyopenai.Name].(*fantasyopenai.ResponsesProviderOptions)
|
||||
require.True(t, ok, "%T", got[fantasyopenai.Name])
|
||||
require.NotNil(t, providerOptions.ReasoningEffort)
|
||||
require.Equal(t, fantasyopenai.ReasoningEffortHigh, *providerOptions.ReasoningEffort)
|
||||
})
|
||||
|
||||
t.Run("PreservesOpenAIResponsesEntry", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
options := fantasy.ProviderOptions{
|
||||
fantasyopenai.Name: &fantasyopenai.ResponsesProviderOptions{
|
||||
Instructions: ptr.Ref("answer briefly"),
|
||||
Store: ptr.Ref(true),
|
||||
},
|
||||
}
|
||||
got := chatprovider.ApplyReasoningEffort(&chattest.FakeModel{ProviderName: fantasyopenai.Name, ModelName: "gpt-5"}, options, new(codersdk.ChatModelReasoningEffortHigh))
|
||||
providerOptions, ok := got[fantasyopenai.Name].(*fantasyopenai.ResponsesProviderOptions)
|
||||
require.True(t, ok, "%T", got[fantasyopenai.Name])
|
||||
require.Same(t, options[fantasyopenai.Name], providerOptions)
|
||||
require.Equal(t, "answer briefly", *providerOptions.Instructions)
|
||||
require.True(t, *providerOptions.Store)
|
||||
require.Equal(t, fantasyopenai.ReasoningEffortHigh, *providerOptions.ReasoningEffort)
|
||||
})
|
||||
|
||||
t.Run("PreservesOpenAILegacyEntry", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
options := fantasy.ProviderOptions{
|
||||
fantasyopenai.Name: &fantasyopenai.ProviderOptions{
|
||||
User: ptr.Ref("user"),
|
||||
ParallelToolCalls: ptr.Ref(true),
|
||||
},
|
||||
}
|
||||
got := chatprovider.ApplyReasoningEffort(&chattest.FakeModel{ProviderName: fantasyopenai.Name, ModelName: "gpt-4"}, options, new(codersdk.ChatModelReasoningEffortHigh))
|
||||
providerOptions, ok := got[fantasyopenai.Name].(*fantasyopenai.ProviderOptions)
|
||||
require.True(t, ok, "%T", got[fantasyopenai.Name])
|
||||
require.Same(t, options[fantasyopenai.Name], providerOptions)
|
||||
require.Equal(t, "user", *providerOptions.User)
|
||||
require.True(t, *providerOptions.ParallelToolCalls)
|
||||
require.Equal(t, fantasyopenai.ReasoningEffortHigh, *providerOptions.ReasoningEffort)
|
||||
})
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
provider string
|
||||
options fantasy.ProviderOptions
|
||||
assert func(*testing.T, fantasy.ProviderOptions)
|
||||
}{
|
||||
{
|
||||
name: "CreatesAnthropicEntry",
|
||||
provider: fantasyanthropic.Name,
|
||||
assert: func(t *testing.T, got fantasy.ProviderOptions) {
|
||||
providerOptions, ok := got[fantasyanthropic.Name].(*fantasyanthropic.ProviderOptions)
|
||||
require.True(t, ok, "%T", got[fantasyanthropic.Name])
|
||||
require.NotNil(t, providerOptions.Effort)
|
||||
require.Equal(t, fantasyanthropic.EffortHigh, *providerOptions.Effort)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "PreservesAnthropicEntry",
|
||||
provider: fantasyanthropic.Name,
|
||||
options: fantasy.ProviderOptions{fantasyanthropic.Name: &fantasyanthropic.ProviderOptions{SendReasoning: ptr.Ref(true)}},
|
||||
assert: func(t *testing.T, got fantasy.ProviderOptions) {
|
||||
providerOptions := got[fantasyanthropic.Name].(*fantasyanthropic.ProviderOptions)
|
||||
require.True(t, *providerOptions.SendReasoning)
|
||||
require.Equal(t, fantasyanthropic.EffortHigh, *providerOptions.Effort)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "CreatesOpenAICompatEntry",
|
||||
provider: fantasyopenaicompat.Name,
|
||||
assert: func(t *testing.T, got fantasy.ProviderOptions) {
|
||||
providerOptions, ok := got[fantasyopenaicompat.Name].(*fantasyopenaicompat.ProviderOptions)
|
||||
require.True(t, ok, "%T", got[fantasyopenaicompat.Name])
|
||||
require.NotNil(t, providerOptions.ReasoningEffort)
|
||||
require.Equal(t, fantasyopenai.ReasoningEffortHigh, *providerOptions.ReasoningEffort)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "PreservesOpenAICompatEntry",
|
||||
provider: fantasyopenaicompat.Name,
|
||||
options: fantasy.ProviderOptions{fantasyopenaicompat.Name: &fantasyopenaicompat.ProviderOptions{User: ptr.Ref("user")}},
|
||||
assert: func(t *testing.T, got fantasy.ProviderOptions) {
|
||||
providerOptions := got[fantasyopenaicompat.Name].(*fantasyopenaicompat.ProviderOptions)
|
||||
require.Equal(t, "user", *providerOptions.User)
|
||||
require.Equal(t, fantasyopenai.ReasoningEffortHigh, *providerOptions.ReasoningEffort)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "CreatesVercelEntry",
|
||||
provider: fantasyvercel.Name,
|
||||
assert: func(t *testing.T, got fantasy.ProviderOptions) {
|
||||
providerOptions, ok := got[fantasyvercel.Name].(*fantasyvercel.ProviderOptions)
|
||||
require.True(t, ok, "%T", got[fantasyvercel.Name])
|
||||
require.NotNil(t, providerOptions.Reasoning)
|
||||
require.NotNil(t, providerOptions.Reasoning.Effort)
|
||||
require.Equal(t, fantasyvercel.ReasoningEffortHigh, *providerOptions.Reasoning.Effort)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "PreservesVercelNestedEntry",
|
||||
provider: fantasyvercel.Name,
|
||||
options: fantasy.ProviderOptions{fantasyvercel.Name: &fantasyvercel.ProviderOptions{Reasoning: &fantasyvercel.ReasoningOptions{Enabled: ptr.Ref(true), MaxTokens: ptr.Ref(int64(1024))}}},
|
||||
assert: func(t *testing.T, got fantasy.ProviderOptions) {
|
||||
providerOptions := got[fantasyvercel.Name].(*fantasyvercel.ProviderOptions)
|
||||
require.True(t, *providerOptions.Reasoning.Enabled)
|
||||
require.Equal(t, int64(1024), *providerOptions.Reasoning.MaxTokens)
|
||||
require.Equal(t, fantasyvercel.ReasoningEffortHigh, *providerOptions.Reasoning.Effort)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "CreatesOpenRouterEntry",
|
||||
provider: fantasyopenrouter.Name,
|
||||
assert: func(t *testing.T, got fantasy.ProviderOptions) {
|
||||
providerOptions, ok := got[fantasyopenrouter.Name].(*fantasyopenrouter.ProviderOptions)
|
||||
require.True(t, ok, "%T", got[fantasyopenrouter.Name])
|
||||
require.NotNil(t, providerOptions.Reasoning)
|
||||
require.NotNil(t, providerOptions.Reasoning.Effort)
|
||||
require.Equal(t, fantasyopenrouter.ReasoningEffortHigh, *providerOptions.Reasoning.Effort)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "PreservesOpenRouterNestedEntry",
|
||||
provider: fantasyopenrouter.Name,
|
||||
options: fantasy.ProviderOptions{fantasyopenrouter.Name: &fantasyopenrouter.ProviderOptions{Reasoning: &fantasyopenrouter.ReasoningOptions{Enabled: ptr.Ref(true), MaxTokens: ptr.Ref(int64(1024))}}},
|
||||
assert: func(t *testing.T, got fantasy.ProviderOptions) {
|
||||
providerOptions, ok := got[fantasyopenrouter.Name].(*fantasyopenrouter.ProviderOptions)
|
||||
require.True(t, ok, "%T", got[fantasyopenrouter.Name])
|
||||
require.True(t, *providerOptions.Reasoning.Enabled)
|
||||
require.Equal(t, int64(1024), *providerOptions.Reasoning.MaxTokens)
|
||||
require.NotNil(t, providerOptions.Reasoning.Effort)
|
||||
require.Equal(t, fantasyopenrouter.ReasoningEffortHigh, *providerOptions.Reasoning.Effort)
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
got := chatprovider.ApplyReasoningEffort(&chattest.FakeModel{ProviderName: tt.provider}, tt.options, new(codersdk.ChatModelReasoningEffortHigh))
|
||||
tt.assert(t, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func effortConfig(defaultEffort, maxEffort string) *codersdk.ChatModelReasoningEffortConfig {
|
||||
cfg := &codersdk.ChatModelReasoningEffortConfig{}
|
||||
if defaultEffort != "" {
|
||||
cfg.Default = ptr.Ref(defaultEffort)
|
||||
}
|
||||
if maxEffort != "" {
|
||||
cfg.Max = ptr.Ref(maxEffort)
|
||||
}
|
||||
return cfg
|
||||
}
|
||||
@@ -22,6 +22,7 @@ type Message struct {
|
||||
Content pqtype.NullRawMessage
|
||||
Visibility database.ChatMessageVisibility
|
||||
ModelConfigID uuid.NullUUID
|
||||
ReasoningEffort database.NullChatReasoningEffort
|
||||
CreatedBy uuid.NullUUID
|
||||
ContentVersion int16
|
||||
Compressed bool
|
||||
@@ -49,6 +50,7 @@ func toInsertParams(chatID uuid.UUID, messages []Message) database.InsertChatMes
|
||||
ChatID: chatID,
|
||||
CreatedBy: make([]uuid.UUID, n),
|
||||
ModelConfigID: make([]uuid.UUID, n),
|
||||
ReasoningEffort: make([]string, n),
|
||||
APIKeyID: make([]string, n),
|
||||
Role: make([]database.ChatMessageRole, n),
|
||||
Content: make([]string, n),
|
||||
@@ -68,6 +70,9 @@ func toInsertParams(chatID uuid.UUID, messages []Message) database.InsertChatMes
|
||||
for i, m := range messages {
|
||||
params.CreatedBy[i] = nullUUIDOrNil(m.CreatedBy)
|
||||
params.ModelConfigID[i] = nullUUIDOrNil(m.ModelConfigID)
|
||||
if m.ReasoningEffort.Valid {
|
||||
params.ReasoningEffort[i] = string(m.ReasoningEffort.ChatReasoningEffort)
|
||||
}
|
||||
if m.APIKeyID.Valid {
|
||||
params.APIKeyID[i] = m.APIKeyID.String
|
||||
}
|
||||
|
||||
@@ -231,11 +231,12 @@ func (tx *Tx) insertQueuedMessage(ownerFallback uuid.UUID, m Message) (database.
|
||||
return database.ChatQueuedMessage{}, err
|
||||
}
|
||||
return tx.store.InsertChatQueuedMessageWithCreator(tx.ctx, database.InsertChatQueuedMessageWithCreatorParams{
|
||||
ChatID: tx.chatID,
|
||||
Content: rawContent,
|
||||
ModelConfigID: m.ModelConfigID,
|
||||
CreatedBy: createdBy,
|
||||
APIKeyID: m.APIKeyID,
|
||||
ChatID: tx.chatID,
|
||||
Content: rawContent,
|
||||
ModelConfigID: m.ModelConfigID,
|
||||
ReasoningEffort: m.ReasoningEffort,
|
||||
CreatedBy: createdBy,
|
||||
APIKeyID: m.APIKeyID,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -243,13 +244,14 @@ func (tx *Tx) insertQueuedMessage(ownerFallback uuid.UUID, m Message) (database.
|
||||
// suitable for promoting into active history.
|
||||
func messageFromQueuedRow(q database.ChatQueuedMessage) Message {
|
||||
return Message{
|
||||
Role: database.ChatMessageRoleUser,
|
||||
Content: pqtype.NullRawMessage{RawMessage: q.Content, Valid: q.Content != nil},
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
ModelConfigID: q.ModelConfigID,
|
||||
CreatedBy: uuid.NullUUID{UUID: q.CreatedBy, Valid: true},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
APIKeyID: q.APIKeyID,
|
||||
Role: database.ChatMessageRoleUser,
|
||||
Content: pqtype.NullRawMessage{RawMessage: q.Content, Valid: q.Content != nil},
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
ModelConfigID: q.ModelConfigID,
|
||||
ReasoningEffort: q.ReasoningEffort,
|
||||
CreatedBy: uuid.NullUUID{UUID: q.CreatedBy, Valid: true},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
APIKeyID: q.APIKeyID,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -484,11 +486,12 @@ func (tx *Tx) sendMessageInterruptRequiresAction(chat database.Chat, m Message)
|
||||
|
||||
// EditMessageInput configures [Tx.EditMessage].
|
||||
type EditMessageInput struct {
|
||||
MessageID int64
|
||||
CreatedBy uuid.UUID
|
||||
Content pqtype.NullRawMessage
|
||||
ModelConfigIDOverride uuid.NullUUID
|
||||
APIKeyID sql.NullString
|
||||
MessageID int64
|
||||
CreatedBy uuid.UUID
|
||||
Content pqtype.NullRawMessage
|
||||
ModelConfigIDOverride uuid.NullUUID
|
||||
ReasoningEffortOverride database.NullChatReasoningEffort
|
||||
APIKeyID sql.NullString
|
||||
}
|
||||
|
||||
// EditMessageResult is returned by [Tx.EditMessage].
|
||||
@@ -564,18 +567,23 @@ func (tx *Tx) EditMessage(input EditMessageInput) (EditMessageResult, error) {
|
||||
if input.ModelConfigIDOverride.Valid {
|
||||
modelConfig = input.ModelConfigIDOverride
|
||||
}
|
||||
reasoningEffort := target.ReasoningEffort
|
||||
if input.ReasoningEffortOverride.Valid {
|
||||
reasoningEffort = input.ReasoningEffortOverride
|
||||
}
|
||||
apiKeyID := input.APIKeyID
|
||||
if !apiKeyID.Valid {
|
||||
return EditMessageResult{}, xerrors.Errorf("api_key_id is required")
|
||||
}
|
||||
replacement := Message{
|
||||
Role: database.ChatMessageRoleUser,
|
||||
Content: input.Content,
|
||||
Visibility: target.Visibility,
|
||||
ModelConfigID: modelConfig,
|
||||
CreatedBy: uuid.NullUUID{UUID: input.CreatedBy, Valid: true},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
APIKeyID: apiKeyID,
|
||||
Role: database.ChatMessageRoleUser,
|
||||
Content: input.Content,
|
||||
Visibility: target.Visibility,
|
||||
ModelConfigID: modelConfig,
|
||||
ReasoningEffort: reasoningEffort,
|
||||
CreatedBy: uuid.NullUUID{UUID: input.CreatedBy, Valid: true},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
APIKeyID: apiKeyID,
|
||||
}
|
||||
insertedReplacement, err := tx.insertMessages([]Message{replacement})
|
||||
if err != nil {
|
||||
|
||||
@@ -29,15 +29,20 @@ func systemMessage(rawContent pqtype.NullRawMessage, modelConfigID uuid.UUID) ch
|
||||
}
|
||||
}
|
||||
|
||||
func userMessageWithAPIKeyID(rawContent pqtype.NullRawMessage, modelConfigID, createdBy uuid.UUID, apiKeyID string) chatstate.Message {
|
||||
func userMessageWithAPIKeyID(rawContent pqtype.NullRawMessage, modelConfigID, createdBy uuid.UUID, apiKeyID string, reasoningEffort *string) chatstate.Message {
|
||||
var effort database.NullChatReasoningEffort
|
||||
if reasoningEffort != nil && *reasoningEffort != "" {
|
||||
effort = database.NullChatReasoningEffort{ChatReasoningEffort: database.ChatReasoningEffort(*reasoningEffort), Valid: true}
|
||||
}
|
||||
return chatstate.Message{
|
||||
Role: database.ChatMessageRoleUser,
|
||||
Content: rawContent,
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: modelConfigID != uuid.Nil},
|
||||
CreatedBy: uuid.NullUUID{UUID: createdBy, Valid: createdBy != uuid.Nil},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
APIKeyID: sql.NullString{String: apiKeyID, Valid: apiKeyID != ""},
|
||||
Role: database.ChatMessageRoleUser,
|
||||
Content: rawContent,
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: modelConfigID != uuid.Nil},
|
||||
ReasoningEffort: effort,
|
||||
CreatedBy: uuid.NullUUID{UUID: createdBy, Valid: createdBy != uuid.Nil},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
APIKeyID: sql.NullString{String: apiKeyID, Valid: apiKeyID != ""},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -532,7 +532,16 @@ func (server *Server) prepareGeneration(
|
||||
}
|
||||
}
|
||||
|
||||
var requestedEffort *string
|
||||
if chat.LastReasoningEffort.Valid {
|
||||
requestedEffort = new(string(chat.LastReasoningEffort.ChatReasoningEffort))
|
||||
}
|
||||
reasoningEffort := chatprovider.ResolveReasoningEffort(
|
||||
requestedEffort,
|
||||
callConfig.ReasoningEffort,
|
||||
)
|
||||
providerOptions := chatprovider.ProviderOptionsFromChatModelConfig(model, callConfig.ProviderOptions)
|
||||
providerOptions = chatprovider.ApplyReasoningEffort(model, providerOptions, reasoningEffort)
|
||||
|
||||
activeToolNames := activeToolNamesForTurn(tools, currentPlanMode, chat.ParentChatID, approvedPlanMCPConfigIDs)
|
||||
if isExploreSubagent {
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
fantasyopenai "charm.land/fantasy/providers/openai"
|
||||
"github.com/google/uuid"
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -13,6 +14,7 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbgen"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/v2/coderd/util/ptr"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprovider"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatstate"
|
||||
@@ -85,6 +87,80 @@ func TestLatestAssistantText(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestPrepareGenerationClampsRequestedReasoningEffortToMax(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := chatdTestContext(t)
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID})
|
||||
org := dbgen.Organization(t, db, database.Organization{})
|
||||
dbgen.OrganizationMember(t, db, database.OrganizationMember{
|
||||
UserID: user.ID,
|
||||
OrganizationID: org.ID,
|
||||
})
|
||||
provider := dbgen.AIProviderWithOptionalKey(t, db, database.AIProvider{
|
||||
Type: database.AIProviderTypeOpenai,
|
||||
}, "test-key")
|
||||
modelConfigRaw, err := json.Marshal(codersdk.ChatModelCallConfig{
|
||||
ReasoningEffort: &codersdk.ChatModelReasoningEffortConfig{
|
||||
Default: ptr.Ref(codersdk.ChatModelReasoningEffortLow),
|
||||
Max: ptr.Ref(codersdk.ChatModelReasoningEffortMedium),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
modelConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
|
||||
Model: "gpt-4o-mini",
|
||||
Options: modelConfigRaw,
|
||||
AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true},
|
||||
}, func(p *database.InsertChatModelConfigParams) {
|
||||
p.Enabled = true
|
||||
})
|
||||
|
||||
created, err := chatstate.CreateChat(ctx, db, ps, chatstate.CreateChatInput{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
LastModelConfigID: modelConfig.ID,
|
||||
Title: "clamp reasoning effort",
|
||||
ClientType: database.ChatClientTypeApi,
|
||||
InitialMessages: []chatstate.Message{
|
||||
{
|
||||
Role: database.ChatMessageRoleUser,
|
||||
Content: mustMarshalText(t, "hello"),
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true},
|
||||
ReasoningEffort: database.NullChatReasoningEffort{
|
||||
ChatReasoningEffort: database.ChatReasoningEffortHigh,
|
||||
Valid: true,
|
||||
},
|
||||
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
APIKeyID: sql.NullString{String: apiKey.ID, Valid: true},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
server := newInternalTestServer(
|
||||
t,
|
||||
db,
|
||||
ps,
|
||||
chatprovider.ProviderAPIKeys{},
|
||||
withInternalTestServerTransportFactory(&aibridgeTestFactory{}),
|
||||
)
|
||||
prepared, err := server.prepareGeneration(ctx, generationPrepareInput{
|
||||
Chat: created.Chat,
|
||||
Messages: created.InitialMessages,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(prepared.Cleanup)
|
||||
|
||||
providerOptions, ok := prepared.ProviderOptions[fantasyopenai.Name].(*fantasyopenai.ResponsesProviderOptions)
|
||||
require.True(t, ok, "%T", prepared.ProviderOptions[fantasyopenai.Name])
|
||||
require.NotNil(t, providerOptions.ReasoningEffort)
|
||||
require.Equal(t, fantasyopenai.ReasoningEffortMedium, *providerOptions.ReasoningEffort)
|
||||
}
|
||||
|
||||
// TestDeriveFinalTurnRunResult exercises the re-derivation path that replaces
|
||||
// the old in-memory generationSideEffects stash. The server here never ran
|
||||
// prepareGeneration, so a passing test proves the finish-turn inputs are
|
||||
|
||||
@@ -26,9 +26,10 @@ func ChatPersonalModelOverrideKey(
|
||||
// When Malformed is true, Mode is the provided default and ModelConfigID is
|
||||
// uuid.Nil.
|
||||
type ParsedChatPersonalModelOverride struct {
|
||||
Mode codersdk.ChatPersonalModelOverrideMode
|
||||
ModelConfigID uuid.UUID
|
||||
Malformed bool
|
||||
Mode codersdk.ChatPersonalModelOverrideMode
|
||||
ModelConfigID uuid.UUID
|
||||
ReasoningEffort *string
|
||||
Malformed bool
|
||||
}
|
||||
|
||||
// ParseChatPersonalModelOverride parses a stored personal model override.
|
||||
@@ -61,15 +62,20 @@ func ParseChatPersonalModelOverride(
|
||||
Malformed: true,
|
||||
}
|
||||
}
|
||||
modelConfigID, err := uuid.Parse(rawModelConfigID)
|
||||
if err != nil {
|
||||
rawID, rawEffort, hasEffort := strings.Cut(rawModelConfigID, ":")
|
||||
modelConfigID, err := uuid.Parse(rawID)
|
||||
if err != nil || (hasEffort && rawEffort == "") {
|
||||
return ParsedChatPersonalModelOverride{
|
||||
Mode: defaultMode,
|
||||
Malformed: true,
|
||||
}
|
||||
}
|
||||
return ParsedChatPersonalModelOverride{
|
||||
parsed := ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeModel,
|
||||
ModelConfigID: modelConfigID,
|
||||
}
|
||||
if hasEffort {
|
||||
parsed.ReasoningEffort = &rawEffort
|
||||
}
|
||||
return parsed
|
||||
}
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/util/ptr"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
@@ -63,6 +64,25 @@ func TestParseChatPersonalModelOverride(t *testing.T) {
|
||||
ModelConfigID: modelConfigID,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "ModelWithReasoningEffort",
|
||||
raw: "model:" + modelConfigID.String() + ":high",
|
||||
defaultMode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
want: chatd.ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeModel,
|
||||
ModelConfigID: modelConfigID,
|
||||
ReasoningEffort: ptr.Ref("high"),
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "ModelWithEmptyReasoningEffort",
|
||||
raw: "model:" + modelConfigID.String() + ":",
|
||||
defaultMode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
want: chatd.ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
Malformed: true,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "InvalidModelUUID",
|
||||
raw: "model:not-a-uuid",
|
||||
|
||||
+49
-16
@@ -2,6 +2,7 @@ package chatd
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"slices"
|
||||
@@ -86,10 +87,11 @@ var preferredTitleModels = []struct {
|
||||
}
|
||||
|
||||
type shortTextCandidate struct {
|
||||
provider string
|
||||
model string
|
||||
route aiGatewayModelRoute
|
||||
lm fantasy.LanguageModel
|
||||
provider string
|
||||
model string
|
||||
route aiGatewayModelRoute
|
||||
lm fantasy.LanguageModel
|
||||
providerOptions fantasy.ProviderOptions
|
||||
}
|
||||
|
||||
func selectPreferredConfiguredShortTextModelConfig(
|
||||
@@ -185,7 +187,7 @@ func (p *Server) GenerateChatTitleAsync(ctx context.Context, chat database.Chat)
|
||||
messages,
|
||||
pasteText,
|
||||
string(route.Provider.Type),
|
||||
modelConfig.Model,
|
||||
modelConfig,
|
||||
model,
|
||||
route,
|
||||
modelOpts,
|
||||
@@ -216,7 +218,7 @@ func (p *Server) maybeGenerateChatTitle(
|
||||
messages []database.ChatMessage,
|
||||
pasteText map[uuid.UUID]string,
|
||||
fallbackProvider string,
|
||||
fallbackModelName string,
|
||||
fallbackConfig database.ChatModelConfig,
|
||||
fallbackModel fantasy.LanguageModel,
|
||||
fallbackRoute aiGatewayModelRoute,
|
||||
modelOpts modelBuildOptions,
|
||||
@@ -257,17 +259,19 @@ func (p *Server) maybeGenerateChatTitle(
|
||||
var candidate shortTextCandidate
|
||||
if overrideSet {
|
||||
candidate = shortTextCandidate{
|
||||
provider: string(overrideRoute.Provider.Type),
|
||||
model: overrideConfig.Model,
|
||||
route: overrideRoute,
|
||||
lm: overrideModel,
|
||||
provider: string(overrideRoute.Provider.Type),
|
||||
model: overrideConfig.Model,
|
||||
route: overrideRoute,
|
||||
lm: overrideModel,
|
||||
providerOptions: p.titleGenerationProviderOptions(ctx, overrideModel, overrideConfig),
|
||||
}
|
||||
} else {
|
||||
candidate = shortTextCandidate{
|
||||
provider: fallbackProvider,
|
||||
model: fallbackModelName,
|
||||
route: fallbackRoute,
|
||||
lm: fallbackModel,
|
||||
provider: fallbackProvider,
|
||||
model: fallbackConfig.Model,
|
||||
route: fallbackRoute,
|
||||
lm: fallbackModel,
|
||||
providerOptions: p.titleGenerationProviderOptions(ctx, fallbackModel, fallbackConfig),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -309,7 +313,7 @@ func (p *Server) maybeGenerateChatTitle(
|
||||
)
|
||||
}
|
||||
|
||||
title, err := generateTitle(candidateCtx, candidateModel, input)
|
||||
title, err := generateTitle(candidateCtx, candidateModel, candidate.providerOptions, input)
|
||||
finishDebugRun(err)
|
||||
if err != nil {
|
||||
if overrideSet {
|
||||
@@ -348,6 +352,28 @@ func (p *Server) maybeGenerateChatTitle(
|
||||
p.publishChatPubsubEvent(chat, codersdk.ChatWatchEventKindTitleChange, nil)
|
||||
}
|
||||
|
||||
func (p *Server) titleGenerationProviderOptions(
|
||||
ctx context.Context,
|
||||
model fantasy.LanguageModel,
|
||||
config database.ChatModelConfig,
|
||||
) fantasy.ProviderOptions {
|
||||
callConfig := codersdk.ChatModelCallConfig{}
|
||||
if len(config.Options) > 0 {
|
||||
if err := json.Unmarshal(config.Options, &callConfig); err != nil {
|
||||
p.logger.Debug(ctx, "failed to parse title generation model call config",
|
||||
slog.F("model_config_id", config.ID),
|
||||
slog.Error(err),
|
||||
)
|
||||
}
|
||||
}
|
||||
providerOptions := chatprovider.ProviderOptionsFromChatModelConfig(model, callConfig.ProviderOptions)
|
||||
return chatprovider.ApplyReasoningEffort(
|
||||
model,
|
||||
providerOptions,
|
||||
chatprovider.ResolveReasoningEffort(nil, callConfig.ReasoningEffort),
|
||||
)
|
||||
}
|
||||
|
||||
func (p *Server) newQuickgenDebugModel(
|
||||
ctx context.Context,
|
||||
chat database.Chat,
|
||||
@@ -471,9 +497,10 @@ func (p *Server) prepareQuickgenDebugCandidate(
|
||||
func generateTitle(
|
||||
ctx context.Context,
|
||||
model fantasy.LanguageModel,
|
||||
providerOptions fantasy.ProviderOptions,
|
||||
input string,
|
||||
) (string, error) {
|
||||
title, err := generateStructuredTitle(ctx, model, titleGenerationPrompt, input)
|
||||
title, err := generateStructuredTitle(ctx, model, providerOptions, titleGenerationPrompt, input)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
@@ -483,12 +510,14 @@ func generateTitle(
|
||||
func generateStructuredTitle(
|
||||
ctx context.Context,
|
||||
model fantasy.LanguageModel,
|
||||
providerOptions fantasy.ProviderOptions,
|
||||
systemPrompt string,
|
||||
userInput string,
|
||||
) (string, error) {
|
||||
title, _, err := generateStructuredTitleWithUsage(
|
||||
ctx,
|
||||
model,
|
||||
providerOptions,
|
||||
systemPrompt,
|
||||
userInput,
|
||||
)
|
||||
@@ -501,6 +530,7 @@ func generateStructuredTitle(
|
||||
func generateStructuredTitleWithUsage(
|
||||
ctx context.Context,
|
||||
model fantasy.LanguageModel,
|
||||
providerOptions fantasy.ProviderOptions,
|
||||
systemPrompt string,
|
||||
userInput string,
|
||||
) (string, fantasy.Usage, error) {
|
||||
@@ -534,6 +564,7 @@ func generateStructuredTitleWithUsage(
|
||||
SchemaDescription: "Propose a short chat title.",
|
||||
MaxOutputTokens: &maxOutputTokens,
|
||||
Temperature: ptr.Ref(quickgenTemperature),
|
||||
ProviderOptions: providerOptions,
|
||||
})
|
||||
return genErr
|
||||
}, nil)
|
||||
@@ -844,6 +875,7 @@ func generateManualTitle(
|
||||
messages []database.ChatMessage,
|
||||
pasteText map[uuid.UUID]string,
|
||||
fallbackModel fantasy.LanguageModel,
|
||||
providerOptions fantasy.ProviderOptions,
|
||||
) (string, fantasy.Usage, error) {
|
||||
turns := extractManualTitleTurns(messages, pasteText)
|
||||
selected := selectManualTitleTurnIndexes(turns)
|
||||
@@ -874,6 +906,7 @@ func generateManualTitle(
|
||||
title, usage, err := generateStructuredTitleWithUsage(
|
||||
titleCtx,
|
||||
fallbackModel,
|
||||
providerOptions,
|
||||
systemPrompt,
|
||||
userInput,
|
||||
)
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"time"
|
||||
|
||||
"charm.land/fantasy"
|
||||
fantasyopenai "charm.land/fantasy/providers/openai"
|
||||
fantasyopenaicompat "charm.land/fantasy/providers/openaicompat"
|
||||
"github.com/google/uuid"
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
@@ -585,7 +586,7 @@ func TestMaybeGenerateChatTitlePreservesUpdatedAt(t *testing.T) {
|
||||
[]database.ChatMessage{message},
|
||||
nil,
|
||||
"openai",
|
||||
"test-model",
|
||||
database.ChatModelConfig{Model: "test-model"},
|
||||
model,
|
||||
aiGatewayModelRoute{},
|
||||
modelBuildOptions{},
|
||||
@@ -608,6 +609,62 @@ func TestMaybeGenerateChatTitlePreservesUpdatedAt(t *testing.T) {
|
||||
require.Equal(t, wantTitle, gotTitle)
|
||||
}
|
||||
|
||||
func TestMaybeGenerateChatTitleAppliesModelConfigReasoningEffort(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
chat, messages := titleOverrideTestChatAndMessages(t)
|
||||
reasoningEffort := "high"
|
||||
maxReasoningEffort := "max"
|
||||
modelConfigRaw, err := json.Marshal(codersdk.ChatModelCallConfig{
|
||||
ReasoningEffort: &codersdk.ChatModelReasoningEffortConfig{
|
||||
Default: &reasoningEffort,
|
||||
Max: &maxReasoningEffort,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
model := &chattest.FakeModel{
|
||||
ProviderName: fantasyopenai.Name,
|
||||
ModelName: "gpt-4o-mini",
|
||||
GenerateObjectFn: func(_ context.Context, call fantasy.ObjectCall) (*fantasy.ObjectResponse, error) {
|
||||
require.NotNil(t, call.MaxOutputTokens)
|
||||
require.Equal(t, int64(256), *call.MaxOutputTokens)
|
||||
providerOptions, ok := call.ProviderOptions[fantasyopenai.Name].(*fantasyopenai.ResponsesProviderOptions)
|
||||
require.True(t, ok, "%T", call.ProviderOptions[fantasyopenai.Name])
|
||||
require.NotNil(t, providerOptions.ReasoningEffort)
|
||||
require.Equal(t, fantasyopenai.ReasoningEffortHigh, *providerOptions.ReasoningEffort)
|
||||
return &fantasy.ObjectResponse{
|
||||
Object: map[string]any{"title": "Reasoning title"},
|
||||
}, nil
|
||||
},
|
||||
}
|
||||
|
||||
db := dbmock.NewMockStore(gomock.NewController(t))
|
||||
db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return("", nil)
|
||||
db.EXPECT().UpdateChatTitleByID(gomock.Any(), database.UpdateChatTitleByIDParams{
|
||||
ID: chat.ID,
|
||||
Title: "Reasoning title",
|
||||
}).Return(chatWithTitle(chat, "Reasoning title"), nil)
|
||||
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
server := titleOverrideTestServer(db, logger)
|
||||
server.maybeGenerateChatTitle(
|
||||
ctx,
|
||||
chat,
|
||||
messages,
|
||||
nil,
|
||||
fantasyopenai.Name,
|
||||
database.ChatModelConfig{Model: "gpt-4o-mini", Options: modelConfigRaw},
|
||||
model,
|
||||
aiGatewayModelRoute{},
|
||||
modelBuildOptions{},
|
||||
&generatedChatTitle{},
|
||||
logger,
|
||||
nil,
|
||||
)
|
||||
}
|
||||
|
||||
func Test_titleGenerationPrompt_UsesSlimRules(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -652,6 +709,7 @@ func Test_generateManualTitle_UsesTimeout(t *testing.T) {
|
||||
messages,
|
||||
nil,
|
||||
model,
|
||||
nil,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "Refresh title", title)
|
||||
@@ -689,6 +747,7 @@ func Test_generateManualTitle_TruncatesFirstUserInput(t *testing.T) {
|
||||
messages,
|
||||
nil,
|
||||
model,
|
||||
nil,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
@@ -723,6 +782,7 @@ func Test_generateManualTitle_ReturnsUsageForEmptyNormalizedTitle(t *testing.T)
|
||||
messages,
|
||||
nil,
|
||||
model,
|
||||
nil,
|
||||
)
|
||||
require.ErrorContains(t, err, "generated title was empty")
|
||||
require.Equal(t, int64(11), usage.InputTokens)
|
||||
@@ -828,6 +888,7 @@ func TestGenerateStructuredTitleWithUsage_OpenAICompatibleRequiredToolChoice(t *
|
||||
title, _, err := generateStructuredTitleWithUsage(
|
||||
t.Context(),
|
||||
model,
|
||||
nil,
|
||||
titleGenerationPrompt,
|
||||
"summarize failed workspace build logs",
|
||||
)
|
||||
|
||||
+85
-54
@@ -192,48 +192,46 @@ func (p *Server) resolveConfiguredModelOverride(
|
||||
resolveModelConfig modelOverrideConfigResolver,
|
||||
resolveProviderKeys modelOverrideProviderKeysResolver,
|
||||
failureMode modelOverrideFailureMode,
|
||||
) (database.ChatModelConfig, bool, error) {
|
||||
trimmed := strings.TrimSpace(raw)
|
||||
if trimmed == "" {
|
||||
return database.ChatModelConfig{}, false, nil
|
||||
}
|
||||
configuredModelConfigID, err := uuid.Parse(trimmed)
|
||||
if err != nil {
|
||||
) (database.ChatModelConfig, *string, bool, error) {
|
||||
parsed, ok := parseModelOverride(raw)
|
||||
if !ok {
|
||||
p.logger.Info(ctx,
|
||||
"invalid model override, ignoring",
|
||||
slog.F("override_context", overrideContext),
|
||||
slog.F("raw_model_config_id", trimmed),
|
||||
slog.Error(err),
|
||||
slog.F("raw_model_config_id", strings.TrimSpace(raw)),
|
||||
)
|
||||
return database.ChatModelConfig{}, false, nil
|
||||
return database.ChatModelConfig{}, nil, false, nil
|
||||
}
|
||||
if parsed.modelConfigID == uuid.Nil {
|
||||
return database.ChatModelConfig{}, nil, false, nil
|
||||
}
|
||||
|
||||
modelConfig, providerName, err := resolveModelConfig(
|
||||
ctx,
|
||||
configuredModelConfigID,
|
||||
parsed.modelConfigID,
|
||||
)
|
||||
if err != nil {
|
||||
if failureMode == modelOverrideFailureModeHard {
|
||||
label := modelOverrideErrorLabel(overrideContext)
|
||||
switch {
|
||||
case errors.Is(err, sql.ErrNoRows):
|
||||
return database.ChatModelConfig{}, true, xerrors.Errorf(
|
||||
return database.ChatModelConfig{}, parsed.reasoningEffort, true, xerrors.Errorf(
|
||||
"%s model override is unavailable: %s",
|
||||
label,
|
||||
configuredModelConfigID,
|
||||
parsed.modelConfigID,
|
||||
)
|
||||
case errors.Is(err, errInvalidModelOverrideMetadata):
|
||||
return database.ChatModelConfig{}, true, xerrors.Errorf(
|
||||
return database.ChatModelConfig{}, parsed.reasoningEffort, true, xerrors.Errorf(
|
||||
"%s model override metadata is invalid for %s: %w",
|
||||
label,
|
||||
configuredModelConfigID,
|
||||
parsed.modelConfigID,
|
||||
err,
|
||||
)
|
||||
default:
|
||||
return database.ChatModelConfig{}, true, xerrors.Errorf(
|
||||
return database.ChatModelConfig{}, parsed.reasoningEffort, true, xerrors.Errorf(
|
||||
"resolve %s model override %s: %w",
|
||||
label,
|
||||
configuredModelConfigID,
|
||||
parsed.modelConfigID,
|
||||
err,
|
||||
)
|
||||
}
|
||||
@@ -244,36 +242,36 @@ func (p *Server) resolveConfiguredModelOverride(
|
||||
p.logger.Info(ctx,
|
||||
"model override is unavailable, ignoring",
|
||||
slog.F("override_context", overrideContext),
|
||||
slog.F("model_config_id", configuredModelConfigID),
|
||||
slog.F("model_config_id", parsed.modelConfigID),
|
||||
)
|
||||
case errors.Is(err, errInvalidModelOverrideMetadata):
|
||||
p.logger.Info(ctx,
|
||||
"model override metadata is invalid, ignoring",
|
||||
slog.F("override_context", overrideContext),
|
||||
slog.F("model_config_id", configuredModelConfigID),
|
||||
slog.F("model_config_id", parsed.modelConfigID),
|
||||
slog.Error(err),
|
||||
)
|
||||
default:
|
||||
p.logger.Warn(ctx,
|
||||
"failed to resolve model override, ignoring",
|
||||
slog.F("override_context", overrideContext),
|
||||
slog.F("model_config_id", configuredModelConfigID),
|
||||
slog.F("model_config_id", parsed.modelConfigID),
|
||||
slog.Error(err),
|
||||
)
|
||||
}
|
||||
return database.ChatModelConfig{}, false, nil
|
||||
return database.ChatModelConfig{}, nil, false, nil
|
||||
}
|
||||
|
||||
providerKeys, err := resolveProviderKeys(ctx, ownerID, modelConfigAIProviderID(modelConfig))
|
||||
if err != nil {
|
||||
return database.ChatModelConfig{}, false, xerrors.Errorf(
|
||||
return database.ChatModelConfig{}, nil, false, xerrors.Errorf(
|
||||
"resolve provider API keys: %w",
|
||||
err,
|
||||
)
|
||||
}
|
||||
if !userCanUseProviderKeys(providerKeys, providerName) {
|
||||
if failureMode == modelOverrideFailureModeHard {
|
||||
return database.ChatModelConfig{}, true, xerrors.Errorf(
|
||||
return database.ChatModelConfig{}, parsed.reasoningEffort, true, xerrors.Errorf(
|
||||
"%s model override credentials are unavailable for provider %q",
|
||||
modelOverrideErrorLabel(overrideContext),
|
||||
providerName,
|
||||
@@ -283,22 +281,22 @@ func (p *Server) resolveConfiguredModelOverride(
|
||||
p.logger.Info(ctx,
|
||||
"model override credentials are unavailable, ignoring",
|
||||
slog.F("override_context", overrideContext),
|
||||
slog.F("model_config_id", configuredModelConfigID),
|
||||
slog.F("model_config_id", parsed.modelConfigID),
|
||||
slog.F("provider", providerName),
|
||||
)
|
||||
return database.ChatModelConfig{}, false, nil
|
||||
return database.ChatModelConfig{}, nil, false, nil
|
||||
}
|
||||
return modelConfig, true, nil
|
||||
return modelConfig, parsed.reasoningEffort, true, nil
|
||||
}
|
||||
|
||||
func (p *Server) resolvePersonalSubagentModelConfigID(
|
||||
ctx context.Context,
|
||||
ownerID uuid.UUID,
|
||||
overrideContext codersdk.ChatModelOverrideContext,
|
||||
) (uuid.UUID, bool, error) {
|
||||
) (uuid.UUID, *string, bool, error) {
|
||||
personalContext, err := personalModelOverrideContextForSubagent(overrideContext)
|
||||
if err != nil {
|
||||
return uuid.Nil, false, err
|
||||
return uuid.Nil, nil, false, err
|
||||
}
|
||||
raw, err := p.db.GetUserChatPersonalModelOverride(
|
||||
ctx,
|
||||
@@ -309,7 +307,7 @@ func (p *Server) resolvePersonalSubagentModelConfigID(
|
||||
)
|
||||
if err != nil {
|
||||
if !xerrors.Is(err, sql.ErrNoRows) {
|
||||
return uuid.Nil, false, xerrors.Errorf(
|
||||
return uuid.Nil, nil, false, xerrors.Errorf(
|
||||
"get %s personal model override: %w",
|
||||
subagentModelOverrideLogLabel(overrideContext),
|
||||
err,
|
||||
@@ -332,7 +330,7 @@ func (p *Server) resolvePersonalSubagentModelConfigID(
|
||||
}
|
||||
switch parsed.Mode {
|
||||
case codersdk.ChatPersonalModelOverrideModeChatDefault:
|
||||
return uuid.Nil, true, nil
|
||||
return uuid.Nil, nil, true, nil
|
||||
case codersdk.ChatPersonalModelOverrideModeDeploymentDefault:
|
||||
case codersdk.ChatPersonalModelOverrideModeModel:
|
||||
modelConfig, ok, err := p.resolvePersonalModelOverride(
|
||||
@@ -342,10 +340,10 @@ func (p *Server) resolvePersonalSubagentModelConfigID(
|
||||
parsed.ModelConfigID,
|
||||
)
|
||||
if err != nil {
|
||||
return uuid.Nil, false, err
|
||||
return uuid.Nil, nil, false, err
|
||||
}
|
||||
if ok {
|
||||
return modelConfig.ID, true, nil
|
||||
return modelConfig.ID, parsed.ReasoningEffort, true, nil
|
||||
}
|
||||
default:
|
||||
p.logger.Warn(ctx,
|
||||
@@ -356,7 +354,7 @@ func (p *Server) resolvePersonalSubagentModelConfigID(
|
||||
)
|
||||
}
|
||||
|
||||
return uuid.Nil, false, nil
|
||||
return uuid.Nil, nil, false, nil
|
||||
}
|
||||
|
||||
func (p *Server) resolvePersonalModelOverride(
|
||||
@@ -417,43 +415,75 @@ func (p *Server) resolvePersonalModelOverride(
|
||||
return modelConfig, true, nil
|
||||
}
|
||||
|
||||
func withResolvedReasoningEffort(
|
||||
modelConfig database.ChatModelConfig,
|
||||
reasoningEffort *string,
|
||||
) database.ChatModelConfig {
|
||||
if reasoningEffort == nil {
|
||||
return modelConfig
|
||||
}
|
||||
callConfig := codersdk.ChatModelCallConfig{}
|
||||
if len(modelConfig.Options) > 0 {
|
||||
if err := json.Unmarshal(modelConfig.Options, &callConfig); err != nil {
|
||||
return modelConfig
|
||||
}
|
||||
}
|
||||
resolvedEffort := chatprovider.ResolveReasoningEffort(
|
||||
reasoningEffort,
|
||||
callConfig.ReasoningEffort,
|
||||
)
|
||||
if resolvedEffort == nil {
|
||||
return modelConfig
|
||||
}
|
||||
callConfig.ReasoningEffort = &codersdk.ChatModelReasoningEffortConfig{
|
||||
Default: resolvedEffort,
|
||||
Max: resolvedEffort,
|
||||
}
|
||||
options, err := json.Marshal(callConfig)
|
||||
if err != nil {
|
||||
return modelConfig
|
||||
}
|
||||
modelConfig.Options = options
|
||||
return modelConfig
|
||||
}
|
||||
|
||||
func (p *Server) resolveSubagentModelConfigID(
|
||||
ctx context.Context,
|
||||
ownerID uuid.UUID,
|
||||
overrideContext codersdk.ChatModelOverrideContext,
|
||||
) (uuid.UUID, error) {
|
||||
) (uuid.UUID, *string, error) {
|
||||
//nolint:gocritic // Chatd needs its scoped config and user-data access here.
|
||||
chatdCtx := dbauthz.AsChatd(ctx)
|
||||
personalOverridesEnabled, err := p.db.GetChatPersonalModelOverridesEnabled(chatdCtx)
|
||||
if err != nil {
|
||||
return uuid.Nil, xerrors.Errorf(
|
||||
return uuid.Nil, nil, xerrors.Errorf(
|
||||
"get chat personal model overrides enabled: %w",
|
||||
err,
|
||||
)
|
||||
}
|
||||
if personalOverridesEnabled {
|
||||
modelConfigID, resolved, err := p.resolvePersonalSubagentModelConfigID(
|
||||
modelConfigID, reasoningEffort, resolved, err := p.resolvePersonalSubagentModelConfigID(
|
||||
chatdCtx,
|
||||
ownerID,
|
||||
overrideContext,
|
||||
)
|
||||
if err != nil {
|
||||
return uuid.Nil, err
|
||||
return uuid.Nil, nil, err
|
||||
}
|
||||
if resolved {
|
||||
return modelConfigID, nil
|
||||
return modelConfigID, reasoningEffort, nil
|
||||
}
|
||||
}
|
||||
|
||||
raw, err := readSubagentModelOverride(chatdCtx, p.db, overrideContext)
|
||||
if err != nil {
|
||||
return uuid.Nil, xerrors.Errorf(
|
||||
return uuid.Nil, nil, xerrors.Errorf(
|
||||
"get %s model override: %w",
|
||||
subagentModelOverrideLogLabel(overrideContext),
|
||||
err,
|
||||
)
|
||||
}
|
||||
modelConfig, ok, err := p.resolveConfiguredModelOverride(
|
||||
modelConfig, reasoningEffort, ok, err := p.resolveConfiguredModelOverride(
|
||||
chatdCtx,
|
||||
string(overrideContext),
|
||||
raw,
|
||||
@@ -463,12 +493,12 @@ func (p *Server) resolveSubagentModelConfigID(
|
||||
modelOverrideFailureModeSoft,
|
||||
)
|
||||
if err != nil {
|
||||
return uuid.Nil, err
|
||||
return uuid.Nil, nil, err
|
||||
}
|
||||
if !ok {
|
||||
return uuid.Nil, nil
|
||||
return uuid.Nil, nil, nil
|
||||
}
|
||||
return modelConfig.ID, nil
|
||||
return modelConfig.ID, reasoningEffort, nil
|
||||
}
|
||||
|
||||
func modelConfigAIProviderID(modelConfig database.ChatModelConfig) uuid.UUID {
|
||||
@@ -932,17 +962,18 @@ func parseSubagentToolChatID(raw string) (uuid.UUID, error) {
|
||||
}
|
||||
|
||||
// childSubagentChatOptions carries per-child overrides for subagent chat
|
||||
// creation. modelConfigIDOverride and planModeOverride apply to any
|
||||
// subagent. inheritedMCPServerIDs is an Explore-only snapshot of the
|
||||
// spawning parent turn's effective external MCP entitlement.
|
||||
// resolveExploreToolSnapshot computes and persists it on the child chat.
|
||||
// Non-Explore children ignore this field.
|
||||
// creation. modelConfigIDOverride, reasoningEffortOverride, and
|
||||
// planModeOverride apply to any subagent. inheritedMCPServerIDs is an
|
||||
// Explore-only snapshot of the spawning parent turn's effective external MCP
|
||||
// entitlement. resolveExploreToolSnapshot computes and persists it on the
|
||||
// child chat. Non-Explore children ignore this field.
|
||||
type childSubagentChatOptions struct {
|
||||
chatMode database.NullChatMode
|
||||
systemPrompt string
|
||||
modelConfigIDOverride *uuid.UUID
|
||||
planModeOverride *database.NullChatPlanMode
|
||||
inheritedMCPServerIDs []uuid.UUID
|
||||
chatMode database.NullChatMode
|
||||
systemPrompt string
|
||||
modelConfigIDOverride *uuid.UUID
|
||||
reasoningEffortOverride *string
|
||||
planModeOverride *database.NullChatPlanMode
|
||||
inheritedMCPServerIDs []uuid.UUID
|
||||
}
|
||||
|
||||
// resolveExploreToolSnapshot computes the child chat's inherited MCP
|
||||
@@ -1118,7 +1149,7 @@ func (p *Server) createChildSubagentChatWithOptions(
|
||||
// workspace context the same way a top-level chat does: pinned from the
|
||||
// agent's latest snapshot (see hydrateChatContextOnCreate below). The
|
||||
// parent's context is not copied into child history.
|
||||
initialMessages = append(initialMessages, userMessageWithAPIKeyID(userContent, modelConfigID, parent.OwnerID, childAPIKeyID))
|
||||
initialMessages = append(initialMessages, userMessageWithAPIKeyID(userContent, modelConfigID, parent.OwnerID, childAPIKeyID, opts.reasoningEffortOverride))
|
||||
|
||||
publisher := p.pubsub
|
||||
if publisher == nil {
|
||||
|
||||
@@ -56,7 +56,7 @@ func allSubagentDefinitions() []subagentDefinition {
|
||||
id: subagentTypeGeneral,
|
||||
description: "substantial delegated research, analysis, reasoning, review, planning support, and implementation",
|
||||
buildOptions: func(ctx context.Context, p *Server, parent database.Chat, _ database.Chat, _ uuid.UUID, _ string) (childSubagentChatOptions, error) {
|
||||
modelConfigID, err := p.resolveSubagentModelConfigID(
|
||||
modelConfigID, reasoningEffort, err := p.resolveSubagentModelConfigID(
|
||||
ctx,
|
||||
parent.OwnerID,
|
||||
codersdk.ChatModelOverrideContextGeneral,
|
||||
@@ -67,6 +67,7 @@ func allSubagentDefinitions() []subagentDefinition {
|
||||
options := childSubagentChatOptions{}
|
||||
if modelConfigID != uuid.Nil {
|
||||
options.modelConfigIDOverride = &modelConfigID
|
||||
options.reasoningEffortOverride = reasoningEffort
|
||||
}
|
||||
return options, nil
|
||||
},
|
||||
@@ -75,7 +76,7 @@ func allSubagentDefinitions() []subagentDefinition {
|
||||
id: subagentTypeExplore,
|
||||
description: "narrow repository-local read-only code discovery and code tracing",
|
||||
buildOptions: func(ctx context.Context, p *Server, _ database.Chat, turnParent database.Chat, currentModelConfigID uuid.UUID, _ string) (childSubagentChatOptions, error) {
|
||||
modelConfigID, err := p.resolveSubagentModelConfigID(
|
||||
modelConfigID, reasoningEffort, err := p.resolveSubagentModelConfigID(
|
||||
ctx,
|
||||
turnParent.OwnerID,
|
||||
codersdk.ChatModelOverrideContextExplore,
|
||||
@@ -101,9 +102,10 @@ func allSubagentDefinitions() []subagentDefinition {
|
||||
ChatMode: database.ChatModeExplore,
|
||||
Valid: true,
|
||||
},
|
||||
modelConfigIDOverride: &modelConfigID,
|
||||
planModeOverride: &clearPlanMode,
|
||||
inheritedMCPServerIDs: inheritedMCPServerIDs,
|
||||
modelConfigIDOverride: &modelConfigID,
|
||||
reasoningEffortOverride: reasoningEffort,
|
||||
planModeOverride: &clearPlanMode,
|
||||
inheritedMCPServerIDs: inheritedMCPServerIDs,
|
||||
}, nil
|
||||
},
|
||||
},
|
||||
|
||||
@@ -29,6 +29,7 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/v2/coderd/database/pubsub"
|
||||
coderdpubsub "github.com/coder/coder/v2/coderd/pubsub"
|
||||
"github.com/coder/coder/v2/coderd/util/ptr"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatloop"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprovider"
|
||||
@@ -1489,7 +1490,7 @@ func TestResolveConfiguredModelOverride_AcceptsAmbientCredentialsProvider(
|
||||
Enabled: true,
|
||||
}
|
||||
|
||||
resolvedModelConfig, ok, err := server.resolveConfiguredModelOverride(
|
||||
resolvedModelConfig, reasoningEffort, ok, err := server.resolveConfiguredModelOverride(
|
||||
ctx,
|
||||
"plan",
|
||||
modelConfig.ID.String(),
|
||||
@@ -1514,6 +1515,7 @@ func TestResolveConfiguredModelOverride_AcceptsAmbientCredentialsProvider(
|
||||
modelOverrideFailureModeSoft,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, reasoningEffort)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, modelConfig, resolvedModelConfig)
|
||||
require.Empty(t, logSink.entriesAtLevelWithMessage(
|
||||
@@ -1522,6 +1524,76 @@ func TestResolveConfiguredModelOverride_AcceptsAmbientCredentialsProvider(
|
||||
))
|
||||
}
|
||||
|
||||
func TestWithResolvedReasoningEffort(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
baseOptions, err := json.Marshal(codersdk.ChatModelCallConfig{
|
||||
MaxOutputTokens: ptr.Ref(int64(123)),
|
||||
ReasoningEffort: &codersdk.ChatModelReasoningEffortConfig{
|
||||
Default: ptr.Ref(codersdk.ChatModelReasoningEffortLow),
|
||||
Max: ptr.Ref(codersdk.ChatModelReasoningEffortHigh),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
t.Run("NilEffortReturnsOriginalConfig", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
modelConfig := database.ChatModelConfig{Options: baseOptions}
|
||||
require.Equal(t, modelConfig, withResolvedReasoningEffort(modelConfig, nil))
|
||||
})
|
||||
|
||||
t.Run("InvalidJSONReturnsOriginalConfig", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
modelConfig := database.ChatModelConfig{Options: []byte(`{`)}
|
||||
require.Equal(t, modelConfig, withResolvedReasoningEffort(modelConfig, ptr.Ref(codersdk.ChatModelReasoningEffortHigh)))
|
||||
})
|
||||
|
||||
t.Run("NoReasoningConfigReturnsOriginalConfig", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
modelConfig := database.ChatModelConfig{Options: []byte(`{"max_output_tokens":123}`)}
|
||||
require.Equal(t, modelConfig, withResolvedReasoningEffort(modelConfig, ptr.Ref(codersdk.ChatModelReasoningEffortHigh)))
|
||||
})
|
||||
|
||||
t.Run("ClampsRequestedEffortAndPreservesOtherOptions", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
modelConfig := database.ChatModelConfig{Options: baseOptions}
|
||||
|
||||
got := withResolvedReasoningEffort(modelConfig, ptr.Ref(codersdk.ChatModelReasoningEffortXHigh))
|
||||
|
||||
var callConfig codersdk.ChatModelCallConfig
|
||||
require.NoError(t, json.Unmarshal(got.Options, &callConfig))
|
||||
require.Equal(t, ptr.Ref(int64(123)), callConfig.MaxOutputTokens)
|
||||
require.Equal(t, ptr.Ref(codersdk.ChatModelReasoningEffortHigh), callConfig.ReasoningEffort.Default)
|
||||
require.Equal(t, ptr.Ref(codersdk.ChatModelReasoningEffortHigh), callConfig.ReasoningEffort.Max)
|
||||
})
|
||||
}
|
||||
|
||||
func TestCreateChildSubagentChat_StoresReasoningEffortOverride(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{})
|
||||
|
||||
ctx := chatdTestContext(t)
|
||||
user, org, model := seedInternalChatDeps(t, db)
|
||||
parentChat := createInternalParentChat(
|
||||
ctx, t, server, db, org.ID, user.ID, model.ID, "parent-effort-override",
|
||||
)
|
||||
ctx = aibridge.WithDelegatedAPIKeyID(ctx, testAPIKeyID(t, server.db, parentChat.OwnerID))
|
||||
child, err := server.createChildSubagentChatWithOptions(
|
||||
ctx,
|
||||
parentChat,
|
||||
"delegate work",
|
||||
"",
|
||||
childSubagentChatOptions{reasoningEffortOverride: ptr.Ref("high")},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
childChat, err := db.GetChatByID(ctx, child.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, database.NullChatReasoningEffort{ChatReasoningEffort: database.ChatReasoningEffortHigh, Valid: true}, childChat.LastReasoningEffort)
|
||||
}
|
||||
|
||||
func TestCreateChildSubagentChat_OverrideWorksWhenParentHasNoModel(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@ package chatd
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
|
||||
"charm.land/fantasy"
|
||||
"github.com/google/uuid"
|
||||
@@ -14,6 +15,28 @@ import (
|
||||
|
||||
const titleGenerationOverrideContext = "title_generation"
|
||||
|
||||
type parsedModelOverride struct {
|
||||
modelConfigID uuid.UUID
|
||||
reasoningEffort *string
|
||||
}
|
||||
|
||||
func parseModelOverride(raw string) (parsedModelOverride, bool) {
|
||||
trimmed := strings.TrimSpace(raw)
|
||||
if trimmed == "" {
|
||||
return parsedModelOverride{}, true
|
||||
}
|
||||
rawID, rawEffort, hasEffort := strings.Cut(trimmed, ":")
|
||||
modelConfigID, err := uuid.Parse(rawID)
|
||||
if err != nil || (hasEffort && rawEffort == "") {
|
||||
return parsedModelOverride{}, false
|
||||
}
|
||||
parsed := parsedModelOverride{modelConfigID: modelConfigID}
|
||||
if hasEffort {
|
||||
parsed.reasoningEffort = &rawEffort
|
||||
}
|
||||
return parsed, true
|
||||
}
|
||||
|
||||
func readTitleGenerationModelOverride(
|
||||
ctx context.Context,
|
||||
db database.Store,
|
||||
@@ -47,7 +70,7 @@ func (p *Server) resolveTitleGenerationModelOverride(
|
||||
)
|
||||
}
|
||||
|
||||
modelConfig, overrideSet, err := p.resolveConfiguredModelOverride(
|
||||
modelConfig, overrideEffort, overrideSet, err := p.resolveConfiguredModelOverride(
|
||||
ctx,
|
||||
titleGenerationOverrideContext,
|
||||
raw,
|
||||
@@ -64,6 +87,7 @@ func (p *Server) resolveTitleGenerationModelOverride(
|
||||
if !overrideSet {
|
||||
return database.ChatModelConfig{}, nil, aiGatewayModelRoute{}, false, nil
|
||||
}
|
||||
modelConfig = withResolvedReasoningEffort(modelConfig, overrideEffort)
|
||||
|
||||
//nolint:gocritic // Title overrides need chatd-scoped provider reads for user-owned chats.
|
||||
route, err := p.resolveModelRouteForConfig(dbauthz.AsChatd(ctx), chat.OwnerID, modelConfig)
|
||||
|
||||
@@ -3,6 +3,7 @@ package chatd
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"strconv"
|
||||
@@ -11,6 +12,7 @@ import (
|
||||
"testing"
|
||||
|
||||
"charm.land/fantasy"
|
||||
fantasyopenai "charm.land/fantasy/providers/openai"
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
@@ -21,6 +23,7 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/aibridge"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbmock"
|
||||
"github.com/coder/coder/v2/coderd/util/ptr"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattest"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
@@ -65,7 +68,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideUnset(t *testing.T) {
|
||||
messages,
|
||||
nil,
|
||||
"openai",
|
||||
"fallback-chat-model",
|
||||
database.ChatModelConfig{Model: "fallback-chat-model"},
|
||||
fallbackModel,
|
||||
aiGatewayModelRoute{},
|
||||
modelBuildOptions{},
|
||||
@@ -115,7 +118,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideReadDBError(t *testing.T)
|
||||
messages,
|
||||
nil,
|
||||
"openai",
|
||||
"fallback-chat-model",
|
||||
database.ChatModelConfig{Model: "fallback-chat-model"},
|
||||
fallbackModel,
|
||||
aiGatewayModelRoute{},
|
||||
modelBuildOptions{},
|
||||
@@ -164,7 +167,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideMalformedFallsThrough(t *
|
||||
messages,
|
||||
nil,
|
||||
"openai",
|
||||
"fallback-chat-model",
|
||||
database.ChatModelConfig{Model: "fallback-chat-model"},
|
||||
fallbackModel,
|
||||
aiGatewayModelRoute{},
|
||||
modelBuildOptions{},
|
||||
@@ -187,16 +190,29 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideSetUsable(t *testing.T) {
|
||||
db := dbmock.NewMockStore(ctrl)
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
|
||||
chat, messages := titleOverrideTestChatAndMessages(t)
|
||||
overrideConfig := titleOverrideModelConfig("gpt-4.1", true)
|
||||
overrideConfig := titleOverrideModelConfig("gpt-5", true)
|
||||
providerID := uuid.New()
|
||||
overrideConfig.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true}
|
||||
options, err := json.Marshal(codersdk.ChatModelCallConfig{
|
||||
ReasoningEffort: &codersdk.ChatModelReasoningEffortConfig{
|
||||
Default: ptr.Ref(codersdk.ChatModelReasoningEffortLow),
|
||||
Max: ptr.Ref(codersdk.ChatModelReasoningEffortHigh),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
overrideConfig.Options = options
|
||||
wantTitle := "Override title"
|
||||
|
||||
var requestCount atomic.Int32
|
||||
factory := &aibridgeTestFactory{rt: roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
requestCount.Add(1)
|
||||
bodyBytes, err := io.ReadAll(req.Body)
|
||||
require.NoError(t, err)
|
||||
var raw map[string]any
|
||||
require.NoError(t, json.Unmarshal(bodyBytes, &raw))
|
||||
require.Equal(t, string(fantasyopenai.ReasoningEffortHigh), raw["reasoning"].(map[string]any)["effort"])
|
||||
text := strconv.Quote(`{"title":"` + wantTitle + `"}`)
|
||||
body := `{"id":"resp_test","object":"response","created_at":0,"status":"completed","model":"gpt-4.1","output":[{"id":"msg_test","type":"message","role":"assistant","content":[{"type":"output_text","text":` + text + `}]}],"usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}`
|
||||
body := `{"id":"resp_test","object":"response","created_at":0,"status":"completed","model":"gpt-5","output":[{"id":"msg_test","type":"message","role":"assistant","content":[{"type":"output_text","text":` + text + `}]}],"usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}`
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"application/json"}},
|
||||
@@ -217,7 +233,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideSetUsable(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil)
|
||||
db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return(overrideConfig.ID.String()+":xhigh", nil)
|
||||
db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil)
|
||||
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes()
|
||||
db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{
|
||||
@@ -238,7 +254,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideSetUsable(t *testing.T) {
|
||||
messages,
|
||||
nil,
|
||||
"openai",
|
||||
"fallback-chat-model",
|
||||
database.ChatModelConfig{Model: "fallback-chat-model"},
|
||||
fallbackModel,
|
||||
aiGatewayModelRoute{},
|
||||
modelBuildOptions{ActiveAPIKeyID: uuid.NewString()},
|
||||
@@ -280,7 +296,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideSetUnusableSkips(t *testi
|
||||
messages,
|
||||
nil,
|
||||
"openai",
|
||||
"fallback-chat-model",
|
||||
database.ChatModelConfig{Model: "fallback-chat-model"},
|
||||
fallbackModel,
|
||||
aiGatewayModelRoute{},
|
||||
modelBuildOptions{},
|
||||
@@ -334,7 +350,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideCallFailureSkipsFallback(
|
||||
messages,
|
||||
nil,
|
||||
"openai",
|
||||
"fallback-chat-model",
|
||||
database.ChatModelConfig{Model: "fallback-chat-model"},
|
||||
fallbackModel,
|
||||
aiGatewayModelRoute{},
|
||||
modelBuildOptions{ActiveAPIKeyID: uuid.NewString()},
|
||||
@@ -667,6 +683,41 @@ func TestResolveManualTitleModel_TitleGenerationOverrideSetUnusable(t *testing.T
|
||||
require.Equal(t, database.ChatModelConfig{}, gotConfig)
|
||||
}
|
||||
|
||||
func TestParseModelOverride(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
modelConfigID := uuid.New()
|
||||
tests := []struct {
|
||||
name string
|
||||
raw string
|
||||
wantID uuid.UUID
|
||||
wantEffort *string
|
||||
wantOK bool
|
||||
}{
|
||||
{name: "Empty", raw: "", wantOK: true},
|
||||
{name: "Whitespace", raw: " \t\n ", wantOK: true},
|
||||
{name: "IDOnly", raw: modelConfigID.String(), wantID: modelConfigID, wantOK: true},
|
||||
{name: "IDWithEffort", raw: modelConfigID.String() + ":high", wantID: modelConfigID, wantEffort: ptr.Ref("high"), wantOK: true},
|
||||
{name: "IDEmptyEffort", raw: modelConfigID.String() + ":", wantOK: false},
|
||||
{name: "OuterWhitespace", raw: " \t" + modelConfigID.String() + ":high\n ", wantID: modelConfigID, wantEffort: ptr.Ref("high"), wantOK: true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got, ok := parseModelOverride(tt.raw)
|
||||
|
||||
require.Equal(t, tt.wantOK, ok)
|
||||
if !tt.wantOK {
|
||||
return
|
||||
}
|
||||
require.Equal(t, tt.wantID, got.modelConfigID)
|
||||
require.Equal(t, tt.wantEffort, got.reasoningEffort)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func titleOverrideTestChatAndMessages(t *testing.T) (database.Chat, []database.ChatMessage) {
|
||||
t.Helper()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user