diff --git a/coderd/x/chatd/ARCHITECTURE.md b/coderd/x/chatd/ARCHITECTURE.md index bef7b57116bb2..409ca071a1853 100644 --- a/coderd/x/chatd/ARCHITECTURE.md +++ b/coderd/x/chatd/ARCHITECTURE.md @@ -876,7 +876,7 @@ Request preparation reads the transport from the model instead of recomputing it The first two happen together in `chatprovider.ProviderOptionsForCall`, the only entry point in `chatprovider` that builds provider options for a call; it delegates transport-aware OpenAI conversion to `chatopenai.ProviderOptionsFromChatConfig`. Config conversion and effort injection cannot pick different option types because one function owns both. -Paths that build their own clients get a `Model` from the same constructor, including the compaction override, quick generation (used by turn status labels and debug models), and the advisor runtime. Within quick generation, only title generation converts the model config through `ProviderOptionsForCall`; the turn status label and chat summary paths deliberately send no provider options, because they are short structured calls that set their own output bounds. Debug recording replaces the wrapped client and preserves the resolved transport. Computer-use turns substitute a hardcoded default model that has no config of its own; it carries its own transport, so the chat model's `openai_config` does not follow it. +Debug recording replaces the wrapped client and preserves the resolved transport. Computer-use turns substitute a hardcoded default model that has no config of its own; it carries its own transport, so the chat model's `openai_config` does not follow it. Azure is deliberately exempt: its provider always enables the Responses API for known models and exposes no equivalent per-model hook, so the transport keeps following the known-model list for Azure. Ignoring the override there is what keeps the decisions above in agreement with the Azure client. The exemption is narrower than it appears, because chatd never builds an azure-typed provider as a fantasy azure client: `fantasyConfigForAIBridge` folds every provider type other than anthropic, bedrock, and openai into openai-compat, which always speaks Chat Completions. diff --git a/coderd/x/chatd/advisor_internal_test.go b/coderd/x/chatd/advisor_internal_test.go index 8c1e979eb8520..20b07d8a40aa0 100644 --- a/coderd/x/chatd/advisor_internal_test.go +++ b/coderd/x/chatd/advisor_internal_test.go @@ -113,12 +113,11 @@ func (p *Server) resolveAdvisorModelOverrideOrFallback( modelOpts modelBuildOptions, logger slog.Logger, ) (chatprovider.Model, codersdk.ChatModelCallConfig) { - model, cfg, err := p.resolveAdvisorModelOverride( + resolved, err := p.resolveAdvisorModelOverride( ctx, chat, advisorCfg, - fallbackModel, - fallbackCallConfig, + resolvedModelCall{model: fallbackModel, callConfig: fallbackCallConfig}, modelOpts, logger, ) @@ -126,7 +125,7 @@ func (p *Server) resolveAdvisorModelOverrideOrFallback( logger.Warn(ctx, "failed to resolve advisor model override, continuing with chat model", slog.Error(err)) return fallbackModel, fallbackCallConfig } - return model, cfg + return resolved.model, resolved.callConfig } func (p *Server) newAdvisorRuntimeOrFallback( @@ -142,8 +141,7 @@ func (p *Server) newAdvisorRuntimeOrFallback( ctx, chat, advisorCfg, - fallbackModel, - fallbackCallConfig, + resolvedModelCall{model: fallbackModel, callConfig: fallbackCallConfig}, modelOpts, logger, ) @@ -268,6 +266,37 @@ func TestResolveAdvisorModelOverride(t *testing.T) { require.Equal(t, fallbackCallConfig, gotCfg) }) + t.Run("InvalidOptionsJSONWithLinkedProviderReturnsFallback", func(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitShort) + configID := uuid.New() + store := &advisorOverrideStubStore{ + getEnabledChatModelConfigByID: func(context.Context, uuid.UUID) (database.ChatModelConfig, error) { + return database.ChatModelConfig{ + ID: configID, + Model: "gpt-5.2", + Enabled: true, + Options: []byte("not valid json"), + DisplayName: "gpt-5.2", + AIProviderID: uuid.NullUUID{UUID: uuid.New(), Valid: true}, + }, nil + }, + } + p := newAdvisorTestServer(ctx, t, store) + + resolved, err := p.resolveAdvisorModelOverride( + ctx, + database.Chat{}, + codersdk.AdvisorConfig{ModelConfigID: configID}, + resolvedModelCall{model: fallbackModel, callConfig: fallbackCallConfig}, + modelBuildOptions{}, + logger, + ) + require.NoError(t, err) + require.Equal(t, fallbackModel, resolved.model) + require.Equal(t, fallbackCallConfig, resolved.callConfig) + }) + t.Run("MissingProviderKeyReturnsFallback", func(t *testing.T) { t.Parallel() ctx := testutil.Context(t, testutil.WaitShort) @@ -416,7 +445,9 @@ func TestResolveAdvisorModelOverride(t *testing.T) { require.True(t, gotModel.Valid()) require.Equal(t, "openai", gotModel.Provider()) require.Equal(t, "gpt-5.2", gotModel.ModelID()) - require.Equal(t, fallbackCallConfig, gotCfg) + require.Equal(t, codersdk.ChatModelCallConfig{ + MaxOutputTokens: ptr.Ref(defaultChatMaxOutputTokens), + }, gotCfg) }) } @@ -446,17 +477,16 @@ func TestResolveAdvisorModelOverridePromotesAIBridgeErrors(t *testing.T) { p := newAdvisorTestServer(ctx, t, store) ctx = aibridge.WithDelegatedAPIKeyID(ctx, uuid.NewString()) - model, _, err := p.resolveAdvisorModelOverride( + resolved, err := p.resolveAdvisorModelOverride( ctx, database.Chat{ID: uuid.New(), OwnerID: uuid.New()}, codersdk.AdvisorConfig{ModelConfigID: configID}, - chatprovider.NewModel(&chattest.FakeModel{ProviderName: "stub", ModelName: "stub"}, nil), - codersdk.ChatModelCallConfig{}, + resolvedModelCall{model: chatprovider.NewModel(&chattest.FakeModel{ProviderName: "stub", ModelName: "stub"}, nil)}, modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, slog.Make(), ) require.ErrorContains(t, err, "AI Gateway transport factory") - require.False(t, model.Valid()) + require.False(t, resolved.model.Valid()) } // TestStripAdvisorGuidanceBlock exercises the filter that keeps the advisor diff --git a/coderd/x/chatd/chatadvisor/runner.go b/coderd/x/chatd/chatadvisor/runner.go index da847ce1b0f90..0f3e9c1f0c365 100644 --- a/coderd/x/chatd/chatadvisor/runner.go +++ b/coderd/x/chatd/chatadvisor/runner.go @@ -47,14 +47,14 @@ func (rt *Runtime) RunAdvisor( // resetProviderOptionsForNestedCall mutates its argument; give it a // clone so the Runtime's stored options stay unchanged across calls. - nestedProviderOptions := cloneProviderOptions(rt.cfg.ProviderOptions) - resetProviderOptionsForNestedCall(nestedProviderOptions) + nestedCall := rt.cfg.CallTemplate + nestedCall.ProviderOptions = cloneProviderOptions(rt.cfg.CallTemplate.ProviderOptions) + resetProviderOptionsForNestedCall(nestedCall.ProviderOptions) assistantOpts := chatloop.GenerateAssistantOptions{ - Model: rt.cfg.Model, - Messages: BuildAdvisorMessages(question, conversationSnapshot), - ModelConfig: rt.cfg.ModelConfig, - ProviderOptions: nestedProviderOptions, + Model: rt.cfg.Model, + Messages: BuildAdvisorMessages(question, conversationSnapshot), + CallTemplate: nestedCall, } if opts != nil && opts.OnAdviceDelta != nil { assistantOpts.PublishMessagePart = func(role codersdk.ChatMessageRole, part codersdk.ChatMessagePart) { diff --git a/coderd/x/chatd/chatadvisor/runner_test.go b/coderd/x/chatd/chatadvisor/runner_test.go index c3830ec4fee66..84b4144a03e96 100644 --- a/coderd/x/chatd/chatadvisor/runner_test.go +++ b/coderd/x/chatd/chatadvisor/runner_test.go @@ -14,7 +14,6 @@ import ( "github.com/coder/coder/v2/coderd/x/chatd/chatadvisor" "github.com/coder/coder/v2/coderd/x/chatd/chattest" - "github.com/coder/coder/v2/codersdk" ) func TestAdvisorRunAdvice(t *testing.T) { @@ -425,12 +424,12 @@ func TestNewRuntimeValidation(t *testing.T) { errText: "advisor max output tokens must be positive", }, { - name: "MismatchedModelConfigMaxOutputTokens", + name: "MismatchedCallTemplateMaxOutputTokens", cfg: chatadvisor.RuntimeConfig{ Model: model, MaxUsesPerRun: 1, MaxOutputTokens: matchingTokens, - ModelConfig: codersdk.ChatModelCallConfig{ + CallTemplate: fantasy.Call{ MaxOutputTokens: &mismatchedTokens, }, }, @@ -473,7 +472,7 @@ func TestNewRuntimeDeepClonesOpenAIResponsesProviderOptions(t *testing.T) { }), nil }, }, - ProviderOptions: parentProviderOpts, + CallTemplate: fantasy.Call{ProviderOptions: parentProviderOpts}, MaxUsesPerRun: 1, MaxOutputTokens: 64, }) @@ -532,7 +531,7 @@ func TestAdvisorRunDisablesStoreAndIsConsistentAcrossCalls(t *testing.T) { }), nil }, }, - ProviderOptions: parentProviderOpts, + CallTemplate: fantasy.Call{ProviderOptions: parentProviderOpts}, MaxUsesPerRun: 2, MaxOutputTokens: 64, }) diff --git a/coderd/x/chatd/chatadvisor/runtime.go b/coderd/x/chatd/chatadvisor/runtime.go index d7282e9706dda..a15c0eaa92c09 100644 --- a/coderd/x/chatd/chatadvisor/runtime.go +++ b/coderd/x/chatd/chatadvisor/runtime.go @@ -6,15 +6,13 @@ import ( "charm.land/fantasy" fantasyopenai "charm.land/fantasy/providers/openai" "golang.org/x/xerrors" - - "github.com/coder/coder/v2/codersdk" ) // RuntimeConfig configures a single advisor runtime instance. type RuntimeConfig struct { - Model fantasy.LanguageModel - ModelConfig codersdk.ChatModelCallConfig - ProviderOptions fantasy.ProviderOptions + Model fantasy.LanguageModel + // CallTemplate's provider options are cloned for each nested call. + CallTemplate fantasy.Call MaxUsesPerRun int MaxOutputTokens int64 } @@ -44,19 +42,19 @@ func NewRuntime(cfg RuntimeConfig) (*Runtime, error) { if cfg.MaxOutputTokens <= 0 { return nil, xerrors.New("advisor max output tokens must be positive") } - if cfg.ModelConfig.MaxOutputTokens != nil && - *cfg.ModelConfig.MaxOutputTokens != cfg.MaxOutputTokens { + if cfg.CallTemplate.MaxOutputTokens != nil && + *cfg.CallTemplate.MaxOutputTokens != cfg.MaxOutputTokens { return nil, xerrors.Errorf( - "advisor model_config.max_output_tokens (%d) must match runtime max output tokens (%d)", - *cfg.ModelConfig.MaxOutputTokens, + "advisor call template max output tokens (%d) must match runtime max output tokens (%d)", + *cfg.CallTemplate.MaxOutputTokens, cfg.MaxOutputTokens, ) } normalized := cfg - normalized.ProviderOptions = cloneProviderOptions(cfg.ProviderOptions) + normalized.CallTemplate.ProviderOptions = cloneProviderOptions(cfg.CallTemplate.ProviderOptions) maxOutputTokens := cfg.MaxOutputTokens - normalized.ModelConfig.MaxOutputTokens = &maxOutputTokens + normalized.CallTemplate.MaxOutputTokens = &maxOutputTokens return &Runtime{cfg: normalized}, nil } @@ -134,7 +132,7 @@ func (rt *Runtime) ProviderOptions() fantasy.ProviderOptions { if rt == nil { return nil } - return rt.cfg.ProviderOptions + return rt.cfg.CallTemplate.ProviderOptions } func (rt *Runtime) tryAcquire() bool { diff --git a/coderd/x/chatd/chatd.go b/coderd/x/chatd/chatd.go index 792e72bcd6ac3..0e6b938028e09 100644 --- a/coderd/x/chatd/chatd.go +++ b/coderd/x/chatd/chatd.go @@ -246,13 +246,12 @@ func (p *Server) resolveAdvisorModelOverride( ctx context.Context, chat database.Chat, advisorCfg codersdk.AdvisorConfig, - fallbackModel chatprovider.Model, - fallbackCallConfig codersdk.ChatModelCallConfig, + fallback resolvedModelCall, modelOpts modelBuildOptions, logger slog.Logger, -) (chatprovider.Model, codersdk.ChatModelCallConfig, error) { +) (resolvedModelCall, error) { if advisorCfg.ModelConfigID == uuid.Nil { - return fallbackModel, fallbackCallConfig, nil + return fallback, nil } // Re-read the override instead of using the cache so disabled models @@ -268,7 +267,7 @@ func (p *Server) resolveAdvisorModelOverride( "advisor model config is disabled or unavailable, continuing with chat model", slog.F("model_config_id", advisorCfg.ModelConfigID), ) - return fallbackModel, fallbackCallConfig, nil + return fallback, nil } logger.Warn( ctx, @@ -276,90 +275,60 @@ func (p *Server) resolveAdvisorModelOverride( slog.F("model_config_id", advisorCfg.ModelConfigID), slog.Error(err), ) - return fallbackModel, fallbackCallConfig, nil + return fallback, nil } - overrideCallConfig := codersdk.ChatModelCallConfig{} - if len(overrideConfig.Options) > 0 { - if err := json.Unmarshal(overrideConfig.Options, &overrideCallConfig); err != nil { - logger.Warn( - ctx, - "failed to parse advisor model config, continuing with chat model", - slog.F("model_config_id", advisorCfg.ModelConfigID), - slog.Error(err), - ) - return fallbackModel, fallbackCallConfig, nil - } - } - - route, err := p.resolveModelRouteForConfig( - ctx, - chat.OwnerID, - overrideConfig, - ) - if err != nil { - if overrideConfig.AIProviderID.Valid { - return chatprovider.Model{}, codersdk.ChatModelCallConfig{}, xerrors.Errorf("resolve advisor override route: %w", err) - } - logger.Warn( - ctx, - "failed to resolve advisor override route, continuing with chat model", - slog.F("model_config_id", advisorCfg.ModelConfigID), - slog.Error(err), - ) - return fallbackModel, fallbackCallConfig, nil - } - overrideModel, err := p.newModel(ctx, modelClientRequest{ - Chat: chat, - ModelName: overrideConfig.Model, - UserAgent: chatprovider.UserAgent(), - ExtraHeaders: chatprovider.CoderHeaders(chat), - ConfigOptions: overrideConfig.Options, - }, route, modelOpts) + resolved, err := p.resolveModelCall(ctx, modelCallSpec{ + purpose: "advisor", + chat: chat, + explicitConfig: &overrideConfig, + buildOptions: modelOpts, + }) if err != nil { - if overrideConfig.AIProviderID.Valid { - return chatprovider.Model{}, codersdk.ChatModelCallConfig{}, xerrors.Errorf("create advisor override model: %w", err) + // Malformed options always fall back; route and client errors are + // hard failures only when the config has a linked provider. + var parseErr modelCallConfigParseError + if overrideConfig.AIProviderID.Valid && !xerrors.As(err, &parseErr) { + return resolvedModelCall{}, xerrors.Errorf("resolve advisor override model: %w", err) } logger.Warn( ctx, - "failed to create advisor override model, continuing with chat model", + "failed to resolve advisor override model, continuing with chat model", slog.F("model_config_id", advisorCfg.ModelConfigID), slog.Error(err), ) - return fallbackModel, fallbackCallConfig, nil + return fallback, nil } if advisorCfg.ReasoningEffort != nil { resolvedEffort := chatprovider.ResolveReasoningEffort( advisorCfg.ReasoningEffort, - overrideCallConfig.ReasoningEffort, + resolved.callConfig.ReasoningEffort, ) if resolvedEffort != nil { - overrideCallConfig.ReasoningEffort = &codersdk.ChatModelReasoningEffortConfig{ + resolved.callConfig.ReasoningEffort = &codersdk.ChatModelReasoningEffortConfig{ Default: resolvedEffort, Max: resolvedEffort, } } } - return overrideModel, overrideCallConfig, nil + return resolved, nil } func (p *Server) newAdvisorRuntime( ctx context.Context, chat database.Chat, advisorCfg codersdk.AdvisorConfig, - fallbackModel chatprovider.Model, - fallbackCallConfig codersdk.ChatModelCallConfig, + fallback resolvedModelCall, modelOpts modelBuildOptions, logger slog.Logger, ) (*chatadvisor.Runtime, error) { - advisorModel, advisorCallConfig, err := p.resolveAdvisorModelOverride( + advisor, err := p.resolveAdvisorModelOverride( ctx, chat, advisorCfg, - fallbackModel, - fallbackCallConfig, + fallback, modelOpts, logger, ) @@ -389,15 +358,14 @@ func (p *Server) newAdvisorRuntime( maxOutputTokens = defaultAdvisorMaxOutputTokens } - advisorCallConfig.MaxOutputTokens = ptr.Ref(maxOutputTokens) + advisor.callConfig.MaxOutputTokens = ptr.Ref(maxOutputTokens) // The override resolver pins an explicit advisor effort into the model // config. Fallback models keep their configured default effort. - providerOptions := chatprovider.ProviderOptionsForCall(advisorModel, advisorCallConfig, nil) + advisor.providerOptions = advisor.deriveProviderOptions(advisor.callConfig, nil) rt, err := chatadvisor.NewRuntime(chatadvisor.RuntimeConfig{ - Model: advisorModel.LanguageModel(), - ModelConfig: advisorCallConfig, - ProviderOptions: providerOptions, + Model: advisor.model.LanguageModel(), + CallTemplate: advisor.newCall(), MaxUsesPerRun: maxUsesPerRun, MaxOutputTokens: maxOutputTokens, }) @@ -2529,23 +2497,20 @@ func (p *Server) generateManualTitleCandidate( } modelOpts := modelBuildOptions{ActiveAPIKeyID: apiKeyID} - model, modelConfig, err := p.resolveManualTitleModel(ctx, store, chat, modelOpts) + resolved, err := p.resolveManualTitleModel(ctx, store, chat, modelOpts) if err != nil { return "", err } titleCtx := ctx - titleModel := model finishDebugRun := func(error) {} - if debugSvc := p.debugService(); debugSvc != nil && debugSvc.IsEnabled(ctx, chat.ID, chat.OwnerID) { - titleCtx, titleModel, finishDebugRun = p.prepareManualTitleDebugRun( + if resolved.debugEnabled { + titleCtx, finishDebugRun = p.prepareManualTitleDebugRun( ctx, - debugSvc, + p.debugService(), chat, - modelConfig, - modelOpts, + resolved, messages, - model, ) } @@ -2553,8 +2518,8 @@ func (p *Server) generateManualTitleCandidate( titleCtx, messages, pasteText, - titleModel.LanguageModel(), - p.titleGenerationProviderOptions(ctx, titleModel, modelConfig), + resolved.model.LanguageModel(), + titleObjectCall(resolved), ) finishDebugRun(err) if err != nil { @@ -2602,61 +2567,12 @@ func (p *Server) prepareManualTitleDebugRun( ctx context.Context, debugSvc *chatdebug.Service, chat database.Chat, - modelConfig database.ChatModelConfig, - modelOpts modelBuildOptions, + resolved resolvedModelCall, messages []database.ChatMessage, - fallbackModel chatprovider.Model, -) (context.Context, chatprovider.Model, func(error)) { +) (context.Context, func(error)) { titleCtx := ctx - titleModel := fallbackModel finishDebugRun := func(error) {} - - route, routeErr := p.resolveModelRouteForConfig(ctx, chat.OwnerID, modelConfig) - var routeProvider string - if routeErr == nil { - routeProvider = string(route.Provider.Type) - } else if modelConfig.AIProviderID.Valid { - // Route resolution failed, but the linked provider still identifies the - // type for the debug run record. Best-effort: leave empty if disabled. - if provider, err := p.enabledAIProviderByID(ctx, modelConfig.AIProviderID.UUID); err == nil { - routeProvider = string(provider.Type) - } - } - debugOpts := modelOpts - debugOpts.RecordHTTP = true - var debugModelErr error - var debugModel chatprovider.Model - if routeErr != nil { - debugModelErr = routeErr - } else { - debugModel, debugModelErr = p.newModel(ctx, modelClientRequest{ - Chat: chat, - ModelName: modelConfig.Model, - UserAgent: chatprovider.UserAgent(), - ExtraHeaders: chatprovider.CoderHeaders(chat), - ConfigOptions: modelConfig.Options, - }, route, debugOpts) - } - switch { - case debugModelErr != nil: - p.logger.Warn(ctx, "failed to create debug-aware manual title model", - slog.F("chat_id", chat.ID), - slog.F("model", modelConfig.Model), - slog.Error(debugModelErr), - ) - case !debugModel.Valid(): - p.logger.Warn(ctx, "manual title debug model creation returned nil", - slog.F("chat_id", chat.ID), - slog.F("model", modelConfig.Model), - ) - default: - titleModel = debugModel.WithLanguageModel(chatdebug.WrapModel(debugModel.LanguageModel(), debugSvc, chatdebug.RecorderOptions{ - ChatID: chat.ID, - OwnerID: chat.OwnerID, - Provider: routeProvider, - Model: modelConfig.Model, - })) - } + modelConfig := resolved.dbConfig var historyTipMessageID int64 if len(messages) > 0 { @@ -2684,7 +2600,7 @@ func (p *Server) prepareManualTitleDebugRun( debugRun, createRunErr := debugSvc.CreateRun(createRunCtx, chatdebug.CreateRunParams{ ChatID: chat.ID, ModelConfigID: modelConfig.ID, - Provider: routeProvider, + Provider: string(resolved.route.Provider.Type), Model: modelConfig.Model, Kind: chatdebug.KindTitleGeneration, Status: chatdebug.StatusInProgress, @@ -2699,7 +2615,7 @@ func (p *Server) prepareManualTitleDebugRun( slog.F("model", modelConfig.Model), slog.Error(createRunErr), ) - return titleCtx, titleModel, finishDebugRun + return titleCtx, finishDebugRun } runContext := chatdebugRunContext(debugRun) @@ -2719,7 +2635,7 @@ func (p *Server) prepareManualTitleDebugRun( } } - return titleCtx, titleModel, finishDebugRun + return titleCtx, finishDebugRun } func chatdebugRunContext(run database.ChatDebugRun) chatdebug.RunContext { @@ -2780,15 +2696,15 @@ func (p *Server) resolveManualTitleModel( store database.Store, chat database.Chat, modelOpts modelBuildOptions, -) (chatprovider.Model, database.ChatModelConfig, error) { - overrideConfig, overrideModel, _, overrideSet, overrideErr := p.resolveTitleGenerationModelOverride( +) (resolvedModelCall, error) { + overrideResolved, overrideSet, overrideErr := p.resolveTitleGenerationModelOverride( ctx, chat, modelOpts, ) if overrideErr != nil { if overrideSet { - return chatprovider.Model{}, database.ChatModelConfig{}, xerrors.Errorf( + return resolvedModelCall{}, xerrors.Errorf( "resolve manual title generation model override: %w", overrideErr, ) @@ -2798,7 +2714,7 @@ func (p *Server) resolveManualTitleModel( slog.Error(overrideErr), ) } else if overrideSet { - return overrideModel, overrideConfig, nil + return overrideResolved, nil } configs, err := store.GetEnabledChatModelConfigs(ctx) @@ -2815,22 +2731,12 @@ func (p *Server) resolveManualTitleModel( return p.resolveFallbackManualTitleModel(ctx, chat, modelOpts) } - route, err := p.resolveModelRouteForConfig(ctx, chat.OwnerID, config) - if err != nil { - p.logger.Debug(ctx, "manual title preferred model unavailable", - slog.F("chat_id", chat.ID), - slog.F("model", config.Model), - slog.Error(err), - ) - return p.resolveFallbackManualTitleModel(ctx, chat, modelOpts) - } - model, err := p.newModel(ctx, modelClientRequest{ - Chat: chat, - ModelName: config.Model, - UserAgent: chatprovider.UserAgent(), - ExtraHeaders: chatprovider.CoderHeaders(chat), - ConfigOptions: config.Options, - }, route, modelOpts) + resolved, err := p.resolveModelCall(ctx, modelCallSpec{ + purpose: "title", + chat: chat, + explicitConfig: &config, + buildOptions: modelOpts, + }) if err != nil { p.logger.Debug(ctx, "manual title preferred model unavailable", slog.F("chat_id", chat.ID), @@ -2839,40 +2745,34 @@ func (p *Server) resolveManualTitleModel( ) return p.resolveFallbackManualTitleModel(ctx, chat, modelOpts) } - - return model, config, nil + return resolved, nil } func (p *Server) resolveFallbackManualTitleModel( ctx context.Context, chat database.Chat, modelOpts modelBuildOptions, -) (chatprovider.Model, database.ChatModelConfig, error) { +) (resolvedModelCall, error) { config, err := p.resolveModelConfig(ctx, chat) if err != nil { - return chatprovider.Model{}, database.ChatModelConfig{}, xerrors.Errorf( + return resolvedModelCall{}, xerrors.Errorf( "resolve fallback manual title model config: %w", err, ) } - route, err := p.resolveModelRouteForConfig(ctx, chat.OwnerID, config) - if err != nil { - return chatprovider.Model{}, database.ChatModelConfig{}, err - } - model, err := p.newModel(ctx, modelClientRequest{ - Chat: chat, - ModelName: config.Model, - UserAgent: chatprovider.UserAgent(), - ExtraHeaders: chatprovider.CoderHeaders(chat), - ConfigOptions: config.Options, - }, route, modelOpts) + resolved, err := p.resolveModelCall(ctx, modelCallSpec{ + purpose: "title", + chat: chat, + explicitConfig: &config, + buildOptions: modelOpts, + }) if err != nil { - return chatprovider.Model{}, database.ChatModelConfig{}, xerrors.Errorf( + return resolvedModelCall{}, xerrors.Errorf( "create fallback manual title model: %w", err, ) } - return model, config, nil + return resolved, nil } func mergeManualTitleMessages( @@ -3447,13 +3347,9 @@ func (p *Server) trackWorkspaceUsage( } type runChatResult struct { - FinalAssistantText string - StatusLabelModel chatprovider.Model - FallbackProvider string - FallbackRoute aiGatewayModelRoute - FallbackModel string - ModelBuildOptions modelBuildOptions - StatusLabelOptions json.RawMessage + FinalAssistantText string + // StatusLabelCall is nil when status-label model resolution failed. + StatusLabelCall *resolvedModelCall TriggerMessageID int64 HistoryTipMessageID int64 } @@ -4026,59 +3922,6 @@ func buildProviderTools(options *codersdk.ChatModelProviderOptions) []chatloop.P return tools } -func (p *Server) resolveChatModel( - ctx context.Context, - chat database.Chat, - modelOpts modelBuildOptions, -) ( - model chatprovider.Model, - dbConfig database.ChatModelConfig, - route aiGatewayModelRoute, - debugEnabled bool, - resolvedProvider string, - resolvedModel string, - err error, -) { - dbConfig, err = p.resolveModelConfig(ctx, chat) - if err != nil { - return chatprovider.Model{}, database.ChatModelConfig{}, aiGatewayModelRoute{}, false, "", "", xerrors.Errorf("resolve model config: %w", err) - } - - if !dbConfig.Enabled { - return chatprovider.Model{}, database.ChatModelConfig{}, aiGatewayModelRoute{}, false, "", "", xerrors.Errorf("chat model config %s is disabled", dbConfig.ID) - } - - route, err = p.resolveModelRouteForConfig(ctx, chat.OwnerID, dbConfig) - if err != nil { - return chatprovider.Model{}, database.ChatModelConfig{}, aiGatewayModelRoute{}, false, "", "", err - } - - providerHint := route.ModelProviderHint - resolvedProvider, resolvedModel, err = chatprovider.ResolveModelWithProviderHint( - dbConfig.Model, - providerHint, - ) - if err != nil { - return chatprovider.Model{}, database.ChatModelConfig{}, aiGatewayModelRoute{}, false, "", "", xerrors.Errorf( - "resolve model metadata: %w", err, - ) - } - - model, debugEnabled, err = p.newDebugAwareModel(ctx, modelClientRequest{ - Chat: chat, - ModelName: dbConfig.Model, - UserAgent: chatprovider.UserAgent(), - ExtraHeaders: chatprovider.CoderHeaders(chat), - ConfigOptions: dbConfig.Options, - }, route, modelOpts) - if err != nil { - return chatprovider.Model{}, database.ChatModelConfig{}, aiGatewayModelRoute{}, false, "", "", xerrors.Errorf( - "create model: %w", err, - ) - } - return model, dbConfig, route, debugEnabled, resolvedProvider, resolvedModel, nil -} - func (p *Server) aiProviderConfig(ctx context.Context, provider database.AIProvider) (chatprovider.ConfiguredProvider, error) { keys, err := p.db.GetAIProviderKeysByProviderID(ctx, provider.ID) if err != nil { @@ -4626,21 +4469,16 @@ func (p *Server) generateFinalTurnStatusLabel( } assistantText := strings.TrimSpace(runResult.FinalAssistantText) - if assistantText == "" || !runResult.StatusLabelModel.Valid() { + if assistantText == "" || runResult.StatusLabelCall == nil { return fallbackTurnStatusLabel(status) } - statusLabel := p.generateTurnStatusLabel( + statusLabel := generateTurnStatusLabel( ctx, chat, status, assistantText, - runResult.FallbackProvider, - runResult.FallbackModel, - runResult.StatusLabelModel, - runResult.FallbackRoute, - runResult.ModelBuildOptions, - runResult.StatusLabelOptions, + *runResult.StatusLabelCall, logger, p.existingDebugService(), runResult.TriggerMessageID, @@ -4861,14 +4699,14 @@ func (p *Server) generateAndStoreChatSummary( } modelOpts := modelBuildOptions{ActiveAPIKeyID: apiKeyID} - model, _, ok := p.resolveChatSummaryModel(ctx, logger, chat, modelOpts) + resolved, ok := p.resolveChatSummaryModel(ctx, logger, chat, modelOpts) if !ok { return } summaryCtx, cancelGen := context.WithTimeout(ctx, chatSummaryGenerateTimeout) defer cancelGen() - summary, _, genErr := generateChatSummary(summaryCtx, model, transcript) + summary, _, genErr := generateChatSummary(summaryCtx, resolved.model.LanguageModel(), summaryObjectCall(resolved), transcript) if genErr != nil { logger.Debug(ctx, "failed to generate chat summary", @@ -4884,15 +4722,18 @@ func (p *Server) resolveChatSummaryModel( logger slog.Logger, chat database.Chat, modelOpts modelBuildOptions, -) (fantasy.LanguageModel, database.ChatModelConfig, bool) { - //nolint:dogsled // resolveChatModel returns rich routing metadata; summary generation only needs the model and its config. - model, dbConfig, _, _, _, _, err := p.resolveChatModel(ctx, chat, modelOpts) +) (resolvedModelCall, bool) { + resolved, err := p.resolveModelCall(ctx, modelCallSpec{ + purpose: "chat_summary", + chat: chat, + buildOptions: modelOpts, + }) if err != nil { logger.Debug(ctx, "failed to resolve chat model for summary", slog.F("chat_id", chat.ID), slog.Error(err)) - return nil, database.ChatModelConfig{}, false + return resolvedModelCall{}, false } - return model.LanguageModel(), dbConfig, true + return resolved, true } func shouldGenerateChatSummary(chat database.Chat, messages []database.ChatMessage) bool { diff --git a/coderd/x/chatd/chatd_debug.go b/coderd/x/chatd/chatd_debug.go index 8fdf19c6a2eab..3787a3a073f7f 100644 --- a/coderd/x/chatd/chatd_debug.go +++ b/coderd/x/chatd/chatd_debug.go @@ -6,7 +6,6 @@ import ( "cdr.dev/slog/v3" "github.com/coder/coder/v2/coderd/x/chatd/chatdebug" - "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" ) const ( @@ -110,36 +109,3 @@ func (p *Server) scheduleDebugCleanup( p.logger.Error(context.WithoutCancel(ctx), "failed to schedule chat debug cleanup", logFields...) } } - -func (p *Server) newDebugAwareModel( - ctx context.Context, - req modelClientRequest, - route aiGatewayModelRoute, - opts modelBuildOptions, -) (chatprovider.Model, bool, error) { - provider, resolvedModel, err := chatprovider.ResolveModelWithProviderHint(req.ModelName, route.ModelProviderHint) - if err != nil { - return chatprovider.Model{}, false, err - } - route.ModelProviderHint = provider - req.ModelName = resolvedModel - - debugSvc := p.debugService() - debugEnabled := debugSvc != nil && debugSvc.IsEnabled(ctx, req.Chat.ID, req.Chat.OwnerID) - opts.RecordHTTP = debugEnabled - - model, err := p.newModel(ctx, req, route, opts) - if err != nil { - return chatprovider.Model{}, debugEnabled, err - } - if !debugEnabled { - return model, false, nil - } - - return model.WithLanguageModel(chatdebug.WrapModel(model.LanguageModel(), debugSvc, chatdebug.RecorderOptions{ - ChatID: req.Chat.ID, - OwnerID: req.Chat.OwnerID, - Provider: provider, - Model: resolvedModel, - })), true, nil -} diff --git a/coderd/x/chatd/chatd_internal_test.go b/coderd/x/chatd/chatd_internal_test.go index 4d16d0b5f3ab3..00b2ecde4002e 100644 --- a/coderd/x/chatd/chatd_internal_test.go +++ b/coderd/x/chatd/chatd_internal_test.go @@ -31,12 +31,10 @@ import ( coderdpubsub "github.com/coder/coder/v2/coderd/pubsub" "github.com/coder/coder/v2/coderd/rbac" "github.com/coder/coder/v2/coderd/workspacestats" - "github.com/coder/coder/v2/coderd/x/chatd/chatdebug" "github.com/coder/coder/v2/coderd/x/chatd/chatloop" openaicomputeruse "github.com/coder/coder/v2/coderd/x/chatd/chatopenai/computeruse" "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/chattest" "github.com/coder/coder/v2/coderd/x/chatd/chattool" skillspkg "github.com/coder/coder/v2/coderd/x/skills" "github.com/coder/coder/v2/codersdk" @@ -3657,74 +3655,6 @@ func TestServer_inflightContext(t *testing.T) { } } -// TestPrepareManualTitleDebugRun_RouteFailureDerivesProviderFromConfig drives -// the fallback branch in prepareManualTitleDebugRun: AI-gateway route -// resolution fails (the BYOK key lookup returns a non-ErrNoRows error) while -// the linked provider stays enabled, so the debug run records the provider -// type derived from modelConfig.AIProviderID instead of an empty string. -func TestPrepareManualTitleDebugRun_RouteFailureDerivesProviderFromConfig(t *testing.T) { - t.Parallel() - - ctx := testutil.Context(t, testutil.WaitShort) - ctrl := gomock.NewController(t) - db := dbmock.NewMockStore(ctrl) - logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) - - ownerID := uuid.New() - providerID := uuid.New() - chat := database.Chat{ID: uuid.New(), OwnerID: ownerID} - modelConfig := database.ChatModelConfig{ - ID: uuid.New(), - Model: "claude-sonnet-4", - AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true}, - } - provider := database.AIProvider{ - ID: providerID, - Type: database.AIProviderTypeAnthropic, - Name: "anthropic", - Enabled: true, - } - - // Resolved twice: once by gatewayProviderForConfig during route resolution, - // once by the fallback's own enabledAIProviderByID lookup. - db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes() - // A non-ErrNoRows BYOK error fails route resolution while the provider stays - // enabled, which is exactly the gap the fallback covers. - db.EXPECT().GetUserAIProviderKeyByProviderID(gomock.Any(), database.GetUserAIProviderKeyByProviderIDParams{ - UserID: ownerID, - AIProviderID: providerID, - }).Return(database.UserAIProviderKey{}, sql.ErrConnDone) - - var gotProvider sql.NullString - db.EXPECT().InsertChatDebugRun(gomock.Any(), gomock.Any()).DoAndReturn( - func(_ context.Context, params database.InsertChatDebugRunParams) (database.ChatDebugRun, error) { - gotProvider = params.Provider - return database.ChatDebugRun{ChatID: params.ChatID, Provider: params.Provider}, nil - }, - ) - - server := &Server{ - db: db, - logger: logger, - allowBYOK: true, - } - debugSvc := chatdebug.NewService(db, logger, nil) - fallbackModel := chatprovider.NewModel(&chattest.FakeModel{ProviderName: "stub", ModelName: "stub"}, nil) - - server.prepareManualTitleDebugRun( - ctx, - debugSvc, - chat, - modelConfig, - modelBuildOptions{}, - nil, - fallbackModel, - ) - - require.True(t, gotProvider.Valid, "debug run provider should be populated from the linked config") - require.Equal(t, "anthropic", gotProvider.String) -} - // TestResolveFallbackModelConfigID verifies that admission does not reuse // a disabled last model and rejects a disabled default. func TestResolveFallbackModelConfigID(t *testing.T) { diff --git a/coderd/x/chatd/chatloop/chatloop.go b/coderd/x/chatd/chatloop/chatloop.go index 89063ff061280..13a154b6940c7 100644 --- a/coderd/x/chatd/chatloop/chatloop.go +++ b/coderd/x/chatd/chatloop/chatloop.go @@ -215,8 +215,9 @@ type GenerateAssistantOptions struct { Clock quartz.Clock ContextLimitFallback int64 - ModelConfig codersdk.ChatModelCallConfig - ProviderOptions fantasy.ProviderOptions + // CallTemplate is copied before GenerateAssistant attaches the prompt and + // tools. + CallTemplate fantasy.Call PublishMessagePart func(codersdk.ChatMessageRole, codersdk.ChatMessagePart) // OnModelStreamStart runs immediately before the provider stream is @@ -305,9 +306,9 @@ type GenerateCompactionOptions struct { ResolvedModel string ModelConfigID uuid.UUID - // ProviderOptions carry summary-model call options such as an - // override's reasoning effort. - ProviderOptions fantasy.ProviderOptions + // SummaryCall is copied before GenerateCompaction attaches the summary + // prompt. + SummaryCall fantasy.Call PublishMessagePart func(codersdk.ChatMessageRole, codersdk.ChatMessagePart) @@ -399,17 +400,9 @@ func GenerateAssistant(ctx context.Context, opts GenerateAssistantOptions) (Assi opts.Metrics.PromptSizeBytes.WithLabelValues(provider, modelName).Observe(float64(EstimatePromptSize(prepared))) opts.Metrics.StepsTotal.WithLabelValues(provider, modelName).Inc() - call := fantasy.Call{ - Prompt: prepared, - Tools: buildToolDefinitions(opts.Tools, opts.ActiveTools, opts.ProviderTools), - MaxOutputTokens: opts.ModelConfig.MaxOutputTokens, - Temperature: opts.ModelConfig.Temperature, - TopP: opts.ModelConfig.TopP, - TopK: opts.ModelConfig.TopK, - PresencePenalty: opts.ModelConfig.PresencePenalty, - FrequencyPenalty: opts.ModelConfig.FrequencyPenalty, - ProviderOptions: opts.ProviderOptions, - } + call := opts.CallTemplate + call.Prompt = prepared + call.Tools = buildToolDefinitions(opts.Tools, opts.ActiveTools, opts.ProviderTools) stepStart := opts.Clock.Now() if opts.OnModelStreamStart != nil { diff --git a/coderd/x/chatd/chatloop/compaction.go b/coderd/x/chatd/chatloop/compaction.go index 320edca3d9237..1b8edb6106ba7 100644 --- a/coderd/x/chatd/chatloop/compaction.go +++ b/coderd/x/chatd/chatloop/compaction.go @@ -94,12 +94,10 @@ type CompactionOptions struct { ChatID uuid.UUID HistoryTipMessageID int64 - // Summary model identity and call options; see - // GenerateCompactionOptions. ResolvedProvider string ResolvedModel string ModelConfigID uuid.UUID - ProviderOptions fantasy.ProviderOptions + SummaryCall fantasy.Call // Force skips the threshold gate (including the threshold=100 // disable and the zero-usage early return). Set for manual, @@ -236,7 +234,7 @@ func normalizedCompactionGenerateConfig(opts GenerateCompactionOptions) (Compact ResolvedProvider: opts.ResolvedProvider, ResolvedModel: opts.ResolvedModel, ModelConfigID: opts.ModelConfigID, - ProviderOptions: opts.ProviderOptions, + SummaryCall: opts.SummaryCall, Force: opts.Force, Source: opts.Source, ToolCallID: opts.ToolCallID, @@ -440,7 +438,6 @@ func generateCompactionSummary( Role: fantasy.MessageRoleUser, Content: summaryParts, }) - toolChoice := fantasy.ToolChoiceNone summaryCtx, finishDebugRun := startCompactionDebugRun(ctx, options) defer func() { @@ -458,11 +455,9 @@ func generateCompactionSummary( finishDebugRun(err) }() - response, err := model.Generate(summaryCtx, fantasy.Call{ - Prompt: summaryPrompt, - ToolChoice: &toolChoice, - ProviderOptions: options.ProviderOptions, - }) + call := options.SummaryCall + call.Prompt = summaryPrompt + response, err := model.Generate(summaryCtx, call) if err != nil { return "", xerrors.Errorf("generate summary text: %w", err) } diff --git a/coderd/x/chatd/compaction_override.go b/coderd/x/chatd/compaction_override.go index fc764ab58dbb5..2db780a947da6 100644 --- a/coderd/x/chatd/compaction_override.go +++ b/coderd/x/chatd/compaction_override.go @@ -2,16 +2,13 @@ package chatd import ( "context" - "encoding/json" - "charm.land/fantasy" "github.com/google/uuid" "golang.org/x/xerrors" "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/database/dbauthz" "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" - "github.com/coder/coder/v2/codersdk" ) const compactionOverrideContext = "compaction" @@ -32,28 +29,12 @@ func readCompactionModelOverride( return raw, nil } -// compactionModelOverride carries the built compaction override model plus -// the identity metadata debug runs and prompt sanitization need. -type compactionModelOverride struct { - modelConfig database.ChatModelConfig - model chatprovider.Model - resolvedProvider string - resolvedModel string - // providerOptions include the override's reasoning effort for the - // summary call. - providerOptions fantasy.ProviderOptions -} - // resolvedCompactionOverride is the compaction override resolved at // prepare time. The provider/model identity is resolved without building // the model client so metrics recorded before the client exists // (still-over-limit) attribute to the same model as the compact action's. type resolvedCompactionOverride struct { - Config database.ChatModelConfig - // ResolvedProvider and ResolvedModel match the built client's - // identity: ResolveModelWithProviderHint normalizes its hint, so the - // normalized provider name here and the route's raw provider type in - // buildCompactionOverrideModel yield the same result. + Config database.ChatModelConfig ResolvedProvider string ResolvedModel string } @@ -106,76 +87,3 @@ func (p *Server) resolveCompactionOverrideConfig( ResolvedModel: resolvedModel, }, nil } - -// buildCompactionOverrideModel resolves the route and constructs the model -// client for a usable override config. Errors are hard failures: a usable -// override that cannot be constructed must fail the generation visibly -// instead of silently compacting with the chat model. -func (p *Server) buildCompactionOverrideModel( - ctx context.Context, - chat database.Chat, - modelConfig database.ChatModelConfig, - modelOpts modelBuildOptions, -) (compactionModelOverride, error) { - //nolint:gocritic // Compaction overrides need chatd-scoped provider reads for user-owned chats. - route, err := p.resolveModelRouteForConfig(dbauthz.AsChatd(ctx), chat.OwnerID, modelConfig) - if err != nil { - return compactionModelOverride{}, xerrors.Errorf( - "resolve compaction model override route: %w", - err, - ) - } - resolvedProvider, resolvedModel, err := chatprovider.ResolveModelWithProviderHint( - modelConfig.Model, - route.ModelProviderHint, - ) - if err != nil { - return compactionModelOverride{}, xerrors.Errorf( - "resolve compaction model override metadata: %w", - err, - ) - } - model, _, err := p.newDebugAwareModel(ctx, modelClientRequest{ - Chat: chat, - ModelName: modelConfig.Model, - UserAgent: chatprovider.UserAgent(), - ExtraHeaders: chatprovider.CoderHeaders(chat), - ConfigOptions: modelConfig.Options, - }, route, modelOpts) - if err != nil { - return compactionModelOverride{}, xerrors.Errorf( - "create compaction model override: %w", - err, - ) - } - providerOptions, err := compactionOverrideProviderOptions(model, modelConfig) - if err != nil { - return compactionModelOverride{}, err - } - return compactionModelOverride{ - modelConfig: modelConfig, - model: model, - resolvedProvider: resolvedProvider, - resolvedModel: resolvedModel, - providerOptions: providerOptions, - }, nil -} - -// compactionOverrideProviderOptions converts the override config's call -// options, including the admin-resolved reasoning effort, into provider -// options for the summary call. -func compactionOverrideProviderOptions( - model chatprovider.Model, - modelConfig database.ChatModelConfig, -) (fantasy.ProviderOptions, error) { - callConfig := codersdk.ChatModelCallConfig{} - if len(modelConfig.Options) > 0 { - if err := json.Unmarshal(modelConfig.Options, &callConfig); err != nil { - return nil, xerrors.Errorf( - "parse compaction model override call config: %w", - err, - ) - } - } - return chatprovider.ProviderOptionsForCall(model, callConfig, nil), nil -} diff --git a/coderd/x/chatd/compaction_override_internal_test.go b/coderd/x/chatd/compaction_override_internal_test.go index 166263c3d5219..50a69adb15b2c 100644 --- a/coderd/x/chatd/compaction_override_internal_test.go +++ b/coderd/x/chatd/compaction_override_internal_test.go @@ -5,7 +5,7 @@ import ( "encoding/json" "testing" - fantasyanthropic "charm.land/fantasy/providers/anthropic" + fantasyopenai "charm.land/fantasy/providers/openai" "github.com/google/uuid" "github.com/stretchr/testify/require" "go.uber.org/mock/gomock" @@ -13,49 +13,10 @@ import ( "cdr.dev/slog/v3/sloggers/slogtest" "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/database/dbmock" - "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" - "github.com/coder/coder/v2/coderd/x/chatd/chattest" "github.com/coder/coder/v2/codersdk" "github.com/coder/coder/v2/testutil" ) -func TestCompactionOverrideProviderOptions(t *testing.T) { - t.Parallel() - - model := chatprovider.NewModel(&chattest.FakeModel{ProviderName: "anthropic", ModelName: "claude-3-5-haiku"}, nil) - - t.Run("NoOptions", func(t *testing.T) { - t.Parallel() - opts, err := compactionOverrideProviderOptions(model, database.ChatModelConfig{}) - require.NoError(t, err) - require.Nil(t, opts) - }) - - t.Run("ReasoningEffort", func(t *testing.T) { - t.Parallel() - effort := "low" - options, err := json.Marshal(codersdk.ChatModelCallConfig{ - ReasoningEffort: &codersdk.ChatModelReasoningEffortConfig{ - Default: &effort, - Max: &effort, - }, - }) - require.NoError(t, err) - opts, err := compactionOverrideProviderOptions(model, database.ChatModelConfig{Options: options}) - require.NoError(t, err) - anthropicOpts, ok := opts[fantasyanthropic.Name].(*fantasyanthropic.ProviderOptions) - require.True(t, ok) - require.NotNil(t, anthropicOpts.Effort) - require.Equal(t, fantasyanthropic.Effort("low"), *anthropicOpts.Effort) - }) - - t.Run("MalformedOptions", func(t *testing.T) { - t.Parallel() - _, err := compactionOverrideProviderOptions(model, database.ChatModelConfig{Options: []byte("{")}) - require.ErrorContains(t, err, "parse compaction model override call config") - }) -} - func TestResolveCompactionOverrideConfig_Unset(t *testing.T) { t.Parallel() @@ -184,6 +145,15 @@ func TestCompactionOverride_SetUsable(t *testing.T) { overrideConfig := titleOverrideModelConfig("gpt-4.1", true) providerID := uuid.New() overrideConfig.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true} + effort := "low" + options, err := json.Marshal(codersdk.ChatModelCallConfig{ + ReasoningEffort: &codersdk.ChatModelReasoningEffortConfig{ + Default: &effort, + Max: &effort, + }, + }) + require.NoError(t, err) + overrideConfig.Options = options db.EXPECT().GetChatCompactionModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil) db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil) @@ -199,19 +169,28 @@ func TestCompactionOverride_SetUsable(t *testing.T) { require.NotNil(t, resolved) require.Equal(t, overrideConfig.ID, resolved.Config.ID) - override, err := server.buildCompactionOverrideModel( - ctx, - chat, - resolved.Config, - modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, - ) + override, err := server.resolveModelCall(ctx, modelCallSpec{ + purpose: "compaction", + chat: chat, + explicitConfig: &resolved.Config, + chatdScopedRoute: true, + buildOptions: modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, + }) require.NoError(t, err) - require.NotNil(t, override.model) - require.Equal(t, overrideConfig.ID, override.modelConfig.ID) + require.True(t, override.model.Valid()) + require.Equal(t, overrideConfig.ID, override.dbConfig.ID) require.Equal(t, "openai", override.resolvedProvider) require.Equal(t, "gpt-4.1", override.resolvedModel) // Prepare-time identity must match the built client's so // still-over-limit metrics land on the same series. require.Equal(t, override.resolvedProvider, resolved.ResolvedProvider) require.Equal(t, override.resolvedModel, resolved.ResolvedModel) + switch opts := override.providerOptions[fantasyopenai.Name].(type) { + case *fantasyopenai.ResponsesProviderOptions: + require.Equal(t, fantasyopenai.ReasoningEffort(effort), *opts.ReasoningEffort) + case *fantasyopenai.ProviderOptions: + require.Equal(t, fantasyopenai.ReasoningEffort(effort), *opts.ReasoningEffort) + default: + t.Fatalf("unexpected openai provider options type %T", opts) + } } diff --git a/coderd/x/chatd/computer_use.go b/coderd/x/chatd/computer_use.go index 249f4b3ae3388..5e4ab54e6c721 100644 --- a/coderd/x/chatd/computer_use.go +++ b/coderd/x/chatd/computer_use.go @@ -7,11 +7,9 @@ import ( "golang.org/x/xerrors" "cdr.dev/slog/v3" - "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/database/dbauthz" "github.com/coder/coder/v2/coderd/x/chatd/chatloop" openaicomputeruse "github.com/coder/coder/v2/coderd/x/chatd/chatopenai/computeruse" - "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" "github.com/coder/coder/v2/coderd/x/chatd/chattool" "github.com/coder/coder/v2/codersdk" "github.com/coder/coder/v2/codersdk/workspacesdk" @@ -56,52 +54,6 @@ func (p *Server) computerUseProviderAndModelFromConfig( return provider, modelProvider, modelName, nil } -func (p *Server) resolveComputerUseModel( - ctx context.Context, - chat database.Chat, - route aiGatewayModelRoute, - computerUseProvider codersdk.ChatComputerUseProvider, - computerUseModelProvider string, - computerUseModelName string, - modelOpts modelBuildOptions, -) ( - model chatprovider.Model, - debugEnabled bool, - resolvedProvider string, - resolvedModel string, - err error, -) { - resolvedProvider, resolvedModel, err = chatprovider.ResolveModelWithProviderHint( - computerUseModelName, - computerUseModelProvider, - ) - if err != nil { - return chatprovider.Model{}, false, "", "", xerrors.Errorf( - "resolve computer use model metadata for provider %q model %q: %w", - computerUseProvider, - computerUseModelName, - err, - ) - } - - model, debugEnabled, err = p.newDebugAwareModel(ctx, modelClientRequest{ - Chat: chat, - ModelName: computerUseModelName, - UserAgent: chatprovider.UserAgent(), - ExtraHeaders: chatprovider.CoderHeaders(chat), - }, route, modelOpts) - if err != nil { - return chatprovider.Model{}, false, "", "", xerrors.Errorf( - "resolve computer use model for provider %q model %q: %w", - computerUseProvider, - computerUseModelName, - err, - ) - } - - return model, debugEnabled, resolvedProvider, resolvedModel, nil -} - type computerUseProviderToolOptions struct { provider codersdk.ChatComputerUseProvider isPlanModeTurn bool diff --git a/coderd/x/chatd/generation.go b/coderd/x/chatd/generation.go index 5fb7f00e6fd21..b7f03424d9ae0 100644 --- a/coderd/x/chatd/generation.go +++ b/coderd/x/chatd/generation.go @@ -53,8 +53,7 @@ type generationPrepared struct { ResolvedProvider string ModelConfigID uuid.UUID - ModelConfig codersdk.ChatModelCallConfig - ProviderOptions fantasy.ProviderOptions + CallTemplate fantasy.Call ContextLimitFallback int64 DynamicToolNames map[string]bool @@ -731,8 +730,7 @@ func (s *taskStarter) generateAssistant( ActiveTools: prepared.ActiveTools, ProviderTools: prepared.ProviderTools, ContextLimitFallback: prepared.ContextLimitFallback, - ModelConfig: prepared.ModelConfig, - ProviderOptions: prepared.ProviderOptions, + CallTemplate: prepared.CallTemplate, PublishMessagePart: attempt.publish, OnModelStreamStart: attempt.startModelInvocation, Logger: s.opts.Logger, @@ -919,7 +917,14 @@ func (s *taskStarter) generateCompaction( compactionOpts := prepared.Compaction.Options metricProvider, metricModel := compactionMetricIdentity(prepared.Compaction) if override := prepared.Compaction.Override; override != nil { - overrideModel, err := s.server.buildCompactionOverrideModel(ctx, prepared.Chat, override.Config, prepared.ModelBuildOptions) + // A usable override that fails to build is a hard generation failure. + overrideModel, err := s.server.resolveModelCall(ctx, modelCallSpec{ + purpose: "compaction", + chat: prepared.Chat, + explicitConfig: &override.Config, + chatdScopedRoute: true, + buildOptions: prepared.ModelBuildOptions, + }) if err != nil { return xerrors.Errorf("build compaction model override: %w", err) } @@ -930,15 +935,15 @@ func (s *taskStarter) generateCompaction( compactionOpts.Model = overrideModel.model.LanguageModel() compactionOpts.ResolvedProvider = overrideModel.resolvedProvider compactionOpts.ResolvedModel = overrideModel.resolvedModel - compactionOpts.ModelConfigID = overrideModel.modelConfig.ID - compactionOpts.ProviderOptions = overrideModel.providerOptions + compactionOpts.ModelConfigID = overrideModel.dbConfig.ID + compactionOpts.SummaryCall = compactionSummaryCall(overrideModel) compactionOpts.Messages = sanitizeCompactionPrompt( ctx, logger, compactionOpts.Messages, overrideModel.model, prepared.Compaction.ChatModelConfig, - overrideModel.modelConfig, + overrideModel.dbConfig, ) } preResult, err := s.server.hooks.Trigger(ctx, chathooks.ChatFor(prepared.Chat, input.hookTurnID()), chathooks.Message{}, agenthooks.EventPreCompact, dispatch.CapacityClassGeneration) diff --git a/coderd/x/chatd/generation_preparer.go b/coderd/x/chatd/generation_preparer.go index 6602d11c84b2c..47f117444c213 100644 --- a/coderd/x/chatd/generation_preparer.go +++ b/coderd/x/chatd/generation_preparer.go @@ -2,7 +2,6 @@ package chatd import ( "context" - "encoding/json" "slices" "strings" "sync" @@ -83,17 +82,9 @@ func (server *Server) prepareGeneration( ) var ( - model chatprovider.Model - modelConfig database.ChatModelConfig - modelRoute aiGatewayModelRoute - modelOpts modelBuildOptions - callConfig codersdk.ChatModelCallConfig - promptRows []database.ChatMessage - mcpConfigs []database.MCPServerConfig - mcpTokens []database.MCPServerUserToken - debugEnabled bool - resolvedProvider string - debugModel string + promptRows []database.ChatMessage + mcpConfigs []database.MCPServerConfig + mcpTokens []database.MCPServerUserToken ) var g errgroup.Group @@ -118,22 +109,21 @@ func (server *Server) prepareGeneration( if err != nil { return generationPrepared{}, xerrors.Errorf("ensure synthetic API key: %w", err) } - modelOpts = modelBuildOptions{ActiveAPIKeyID: apiKeyID} + modelOpts := modelBuildOptions{ActiveAPIKeyID: apiKeyID} - model, modelConfig, modelRoute, debugEnabled, resolvedProvider, debugModel, err = server.resolveChatModel(ctx, chat, modelOpts) + requestedEffort := chatRequestedEffort(chat) + resolved, err := server.resolveModelCall(ctx, modelCallSpec{ + purpose: "standard_turn", + chat: chat, + requestedEffort: requestedEffort, + buildOptions: modelOpts, + }) if err != nil { return generationPrepared{}, err } - if len(modelConfig.Options) > 0 { - if err := json.Unmarshal(modelConfig.Options, &callConfig); err != nil { - return generationPrepared{}, xerrors.Errorf("parse model call config: %w", err) - } - } - - if callConfig.MaxOutputTokens == nil { - maxOutputTokens := int64(32_000) - callConfig.MaxOutputTokens = &maxOutputTokens - } + // The chat config keeps driving compaction, sanitization, and debug + // attribution even when computer use swaps the resolved call below. + modelConfig := resolved.dbConfig // Computer-use turns swap in a specialized model, so the substitution // must happen before anything model-sensitive runs: file-part @@ -147,28 +137,30 @@ func (server *Server) prepareGeneration( if err != nil { return generationPrepared{}, xerrors.Errorf("resolve computer use provider and model: %w", err) } - computerUseRoute, keyErr := server.resolveModelRouteForProviderType(ctx, chat.OwnerID, cuModelProvider) - if keyErr != nil { - return generationPrepared{}, xerrors.Errorf("resolve computer use provider route: %w", keyErr) - } - modelRoute = computerUseRoute - cuModel, cuDebugEnabled, cuResolvedProvider, cuResolvedModel, cuErr := server.resolveComputerUseModel( - ctx, - chat, - computerUseRoute, - computerUseProvider, - cuModelProvider, - cuModelName, - modelOpts, - ) + cuResolved, cuErr := server.resolveModelCall(ctx, modelCallSpec{ + purpose: "computer_use", + chat: chat, + fixedModel: &fixedModelCall{ + providerType: cuModelProvider, + modelName: cuModelName, + callConfig: resolved.callConfig, + }, + requestedEffort: requestedEffort, + buildOptions: modelOpts, + }) if cuErr != nil { - return generationPrepared{}, cuErr + return generationPrepared{}, xerrors.Errorf( + "resolve computer use model for provider %q model %q: %w", + computerUseProvider, + cuModelName, + cuErr, + ) } - model = cuModel - debugEnabled = cuDebugEnabled - resolvedProvider = cuResolvedProvider - debugModel = cuResolvedModel + resolved = cuResolved } + model := resolved.model + callConfig := resolved.callConfig + modelRoute := resolved.route currentPlanMode := chat.PlanMode isPlanModeTurn := currentPlanMode.Valid && currentPlanMode.ChatPlanMode == database.ChatPlanModePlan @@ -197,8 +189,7 @@ func (server *Server) prepareGeneration( ctx, chat, advisorCfg, - model, - callConfig, + resolved, modelOpts, logger, ) @@ -580,12 +571,6 @@ func (server *Server) prepareGeneration( } } - var requestedEffort *string - if chat.LastReasoningEffort.Valid { - requestedEffort = new(string(chat.LastReasoningEffort.ChatReasoningEffort)) - } - providerOptions := chatprovider.ProviderOptionsForCall(model, callConfig, requestedEffort) - activeToolNames := activeToolNamesForTurn(tools, currentPlanMode, chat.ParentChatID, approvedPlanMCPConfigIDs) if isExploreSubagent { activeToolNames = allowedExploreToolNames(tools) @@ -601,7 +586,7 @@ func (server *Server) prepareGeneration( triggerMessageID, historyTipMessageID, triggerLabel := deriveChatDebugSeed(promptRows) debugSvc := server.existingDebugService() var debug *generationDebug - if debugEnabled { + if resolved.debugEnabled { if debugSvc == nil { cleanup() return generationPrepared{}, xerrors.New("chat debug service missing after enablement check") @@ -609,8 +594,8 @@ func (server *Server) prepareGeneration( debug = &generationDebug{ Enabled: true, Service: debugSvc, - Provider: resolvedProvider, - Model: debugModel, + Provider: resolved.resolvedProvider, + Model: resolved.resolvedModel, TriggerMessageID: triggerMessageID, HistoryTipMessageID: historyTipMessageID, TriggerLabel: triggerLabel, @@ -653,10 +638,11 @@ func (server *Server) prepareGeneration( DebugSvc: debugSvc, ChatID: chat.ID, HistoryTipMessageID: historyTipMessageID, - ResolvedProvider: resolvedProvider, - ResolvedModel: debugModel, + ResolvedProvider: resolved.resolvedProvider, + ResolvedModel: resolved.resolvedModel, ModelConfigID: modelConfig.ID, StepUsage: compactionStepUsage, + SummaryCall: compactionSummaryCall(resolved), } // workspaceCtx.currentChatSnapshot may carry a freshly persisted @@ -678,10 +664,9 @@ func (server *Server) prepareGeneration( ProviderTools: providerTools, ModelRoute: modelRoute, ModelBuildOptions: modelOpts, - ResolvedProvider: resolvedProvider, + ResolvedProvider: resolved.resolvedProvider, ModelConfigID: modelConfig.ID, - ModelConfig: callConfig, - ProviderOptions: providerOptions, + CallTemplate: resolved.newCall(), ContextLimitFallback: modelConfig.ContextLimit, DynamicToolNames: dynamicToolNames, StopAfterTools: stopAfterBehaviorTools(currentPlanMode, chat.Mode, chat.ParentChatID), @@ -801,18 +786,19 @@ func (server *Server) deriveFinalTurnRunResult( return runChatResult{} } - // resolvedProvider/resolvedModel describe the model the fallback handle was - // built from; they only feed the status-label fallback candidate's labels. apiKeyID, err := server.ensureSyntheticAPIKeyID(ctx, chat.OwnerID) if err != nil { logger.Warn(ctx, "derive final turn status label: ensure synthetic API key", slog.Error(err)) return runChatResult{FinalAssistantText: finalAssistantText, TriggerMessageID: triggerMessageID, HistoryTipMessageID: historyTipMessageID} } modelOpts := modelBuildOptions{ActiveAPIKeyID: apiKeyID} - model, dbConfig, modelRoute, _, resolvedProvider, resolvedModel, err := server.resolveChatModel(ctx, chat, modelOpts) + resolved, err := server.resolveModelCall(ctx, modelCallSpec{ + purpose: "turn_status_label", + chat: chat, + buildOptions: modelOpts, + }) if err != nil { - // Return what we have; generateFinalTurnStatusLabel falls back to a - // generic label when StatusLabelModel is nil. + // Preserve the text and IDs for the generic-label fallback. logger.Warn(ctx, "derive final turn status label: resolve model", slog.Error(err)) return runChatResult{ FinalAssistantText: finalAssistantText, @@ -823,12 +809,7 @@ func (server *Server) deriveFinalTurnRunResult( return runChatResult{ FinalAssistantText: finalAssistantText, - StatusLabelModel: model, - FallbackProvider: resolvedProvider, - FallbackRoute: modelRoute, - FallbackModel: resolvedModel, - ModelBuildOptions: modelOpts, - StatusLabelOptions: dbConfig.Options, + StatusLabelCall: &resolved, TriggerMessageID: triggerMessageID, HistoryTipMessageID: historyTipMessageID, } diff --git a/coderd/x/chatd/generation_preparer_internal_test.go b/coderd/x/chatd/generation_preparer_internal_test.go index 83fd496739bc8..651ad17bbf477 100644 --- a/coderd/x/chatd/generation_preparer_internal_test.go +++ b/coderd/x/chatd/generation_preparer_internal_test.go @@ -104,6 +104,11 @@ func TestPrepareGenerationClampsRequestedReasoningEffortToMax(t *testing.T) { Type: database.AIProviderTypeOpenai, }, "test-key") modelConfigRaw, err := json.Marshal(codersdk.ChatModelCallConfig{ + ProviderOptions: &codersdk.ChatModelProviderOptions{ + OpenAI: &codersdk.ChatModelOpenAIProviderOptions{ + User: ptr.Ref("turn-options-sentinel"), + }, + }, ReasoningEffort: &codersdk.ChatModelReasoningEffortConfig{ Default: ptr.Ref(codersdk.ChatModelReasoningEffortLow), Max: ptr.Ref(codersdk.ChatModelReasoningEffortMedium), @@ -155,10 +160,24 @@ func TestPrepareGenerationClampsRequestedReasoningEffortToMax(t *testing.T) { 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]) + providerOptions, ok := prepared.CallTemplate.ProviderOptions[fantasyopenai.Name].(*fantasyopenai.ResponsesProviderOptions) + require.True(t, ok, "%T", prepared.CallTemplate.ProviderOptions[fantasyopenai.Name]) require.NotNil(t, providerOptions.ReasoningEffort) require.Equal(t, fantasyopenai.ReasoningEffortMedium, *providerOptions.ReasoningEffort) + + require.NotNil(t, providerOptions.User) + require.Equal(t, "turn-options-sentinel", *providerOptions.User) + require.NotNil(t, prepared.CallTemplate.MaxOutputTokens) + require.Equal(t, defaultChatMaxOutputTokens, *prepared.CallTemplate.MaxOutputTokens) + + require.NotNil(t, prepared.Compaction) + summaryCall := prepared.Compaction.Options.SummaryCall + require.Equal(t, prepared.CallTemplate.ProviderOptions, summaryCall.ProviderOptions) + require.NotNil(t, summaryCall.ToolChoice) + require.Equal(t, fantasy.ToolChoiceNone, *summaryCall.ToolChoice) + // Non-streaming summaries must not inherit the default output cap the + // Anthropic SDK rejects. + require.Nil(t, summaryCall.MaxOutputTokens) } func TestPrepareGenerationComputerUseIgnoresChatTransportOverride(t *testing.T) { @@ -248,8 +267,8 @@ func TestPrepareGenerationComputerUseIgnoresChatTransportOverride(t *testing.T) // The computer-use model is Responses-selected by the SDK and its client // ignores the config's forced Chat Completions, so the options must be the // Responses type or the SDK discards them. - _, ok := prepared.ProviderOptions[fantasyopenai.Name].(*fantasyopenai.ResponsesProviderOptions) - require.True(t, ok, "%T", prepared.ProviderOptions[fantasyopenai.Name]) + _, ok := prepared.CallTemplate.ProviderOptions[fantasyopenai.Name].(*fantasyopenai.ResponsesProviderOptions) + require.True(t, ok, "%T", prepared.CallTemplate.ProviderOptions[fantasyopenai.Name]) // File classification must also key on the substituted model: the // Responses transport drops native text file parts, so the attachment @@ -442,10 +461,11 @@ func TestDeriveFinalTurnRunResult(t *testing.T) { require.Equal(t, "the answer is 42", result.FinalAssistantText) require.Equal(t, lastUserID, result.TriggerMessageID) require.Equal(t, tipID, result.HistoryTipMessageID) - require.True(t, result.StatusLabelModel.Valid()) - require.Equal(t, "openai", result.FallbackProvider) - require.Equal(t, "gpt-4o-mini", result.FallbackModel) - require.JSONEq(t, `{"openai_config":{"use_responses_api":false}}`, string(result.StatusLabelOptions)) + require.NotNil(t, result.StatusLabelCall) + require.True(t, result.StatusLabelCall.model.Valid()) + require.Equal(t, "openai", result.StatusLabelCall.resolvedProvider) + require.Equal(t, "gpt-4o-mini", result.StatusLabelCall.resolvedModel) + require.JSONEq(t, `{"openai_config":{"use_responses_api":false}}`, string(result.StatusLabelCall.dbConfig.Options)) }) t.Run("NonWaitingReturnsEmpty", func(t *testing.T) { @@ -481,8 +501,6 @@ func TestDeriveFinalTurnRunResult(t *testing.T) { UserID: user.ID, OrganizationID: org.ID, }) - // A disabled AI provider makes resolveChatModel fail, exercising the - // degraded path that still returns the re-derived text and IDs. provider := insertInternalAIProvider(t, db, database.AIProviderTypeOpenai, "provider-api-key", false) modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ Model: "gpt-4o-mini", @@ -519,9 +537,7 @@ func TestDeriveFinalTurnRunResult(t *testing.T) { require.Equal(t, "the answer is 42", result.FinalAssistantText) require.NotZero(t, result.TriggerMessageID) require.NotZero(t, result.HistoryTipMessageID) - require.False(t, result.StatusLabelModel.Valid()) - require.Empty(t, result.FallbackProvider) - require.Empty(t, result.FallbackModel) + require.Nil(t, result.StatusLabelCall) }) } diff --git a/coderd/x/chatd/model_routing.go b/coderd/x/chatd/model_routing.go index a099b651572f6..c089570846bad 100644 --- a/coderd/x/chatd/model_routing.go +++ b/coderd/x/chatd/model_routing.go @@ -2,7 +2,6 @@ package chatd import ( "context" - "encoding/json" "net/http" "github.com/google/uuid" @@ -18,9 +17,9 @@ type modelClientRequest struct { ModelName string UserAgent string ExtraHeaders map[string]string - // ConfigOptions holds the model config row's Options JSONB; empty for - // paths without a config row. - ConfigOptions json.RawMessage + // CallConfig is the parsed model config row's options; zero for paths + // without a config row. + CallConfig codersdk.ChatModelCallConfig } type modelBuildOptions struct { diff --git a/coderd/x/chatd/model_routing_aibridge.go b/coderd/x/chatd/model_routing_aibridge.go index 0311c597f2002..924dfb87de2a9 100644 --- a/coderd/x/chatd/model_routing_aibridge.go +++ b/coderd/x/chatd/model_routing_aibridge.go @@ -166,11 +166,7 @@ func (p *Server) newModel( } config := fantasyConfigForAIBridge(route.Provider.Type) - callConfig, err := parseModelConfigOptions(req.ConfigOptions) - if err != nil { - return chatprovider.Model{}, err - } - extraHeaders := mergeConfigBetaHeaders(req.ExtraHeaders, config.ProviderHint, callConfig) + extraHeaders := mergeConfigBetaHeaders(req.ExtraHeaders, config.ProviderHint, req.CallConfig) return newLanguageModel( config.ProviderHint, req.ModelName, @@ -178,7 +174,7 @@ func (p *Server) newModel( req.UserAgent, extraHeaders, &http.Client{Transport: baseRT}, - callConfig.OpenAIConfig, + req.CallConfig.OpenAIConfig, ) } diff --git a/coderd/x/chatd/model_routing_internal_test.go b/coderd/x/chatd/model_routing_internal_test.go index 6a6d1324ee6e1..76c41f14b0857 100644 --- a/coderd/x/chatd/model_routing_internal_test.go +++ b/coderd/x/chatd/model_routing_internal_test.go @@ -397,13 +397,10 @@ func TestAIGatewayModelAppliesResponsesAPIOverride(t *testing.T) { return &Server{aibridgeTransportFactory: aibridgeTestFactoryPointer(factory)} } - configOptions := func(t *testing.T, useResponsesAPI *bool) json.RawMessage { - t.Helper() - raw, err := json.Marshal(codersdk.ChatModelCallConfig{ + callConfig := func(useResponsesAPI *bool) codersdk.ChatModelCallConfig { + return codersdk.ChatModelCallConfig{ OpenAIConfig: &codersdk.ChatModelOpenAIConfig{UseResponsesAPI: useResponsesAPI}, - }) - require.NoError(t, err) - return raw + } } forceResponses := true @@ -428,7 +425,7 @@ func TestAIGatewayModelAppliesResponsesAPIOverride(t *testing.T) { server := newServer(t, paths) provider := aibridgeTestAIProvider(uuid.New(), "primary-openai", database.AIProviderTypeOpenai) req := aibridgeTestRequest(database.Chat{ID: uuid.New(), OwnerID: uuid.New()}, tt.model) - req.ConfigOptions = configOptions(t, tt.override) + req.CallConfig = callConfig(tt.override) model, err := server.newModel( t.Context(), @@ -634,38 +631,42 @@ func TestAIBridgeGatewayProviderTypesPreserveSlashModelID(t *testing.T) { } } +func computerUseTestServer(t *testing.T, factory *aibridgeTestFactory) *Server { + t.Helper() + ctrl := gomock.NewController(t) + db := dbmock.NewMockStore(ctrl) + db.EXPECT().GetAIProviders(gomock.Any(), gomock.Any()).Return([]database.AIProvider{ + aibridgeTestAIProvider(uuid.New(), "primary-openai", database.AIProviderTypeOpenai), + }, nil).AnyTimes() + return &Server{db: db, aibridgeTransportFactory: aibridgeTestFactoryPointer(factory)} +} + func TestAIBridgeComputerUseModelUsesRoute(t *testing.T) { t.Parallel() - providerID := uuid.New() apiKeyID := uuid.NewString() factory := &aibridgeTestFactory{rt: roundTripFunc(func(*http.Request) (*http.Response, error) { t.Fatal("computer use model construction must not send a request") return nil, xerrors.New("unreachable") })} chat := database.Chat{ID: uuid.New(), OwnerID: uuid.New()} - server := &Server{ - aibridgeTransportFactory: aibridgeTestFactoryPointer(factory), - } + server := computerUseTestServer(t, factory) provider := codersdk.ChatComputerUseProviderOpenAI modelProvider, modelName, ok := chattool.DefaultComputerUseModel(provider) require.True(t, ok) ctx := aibridge.WithDelegatedAPIKeyID(t.Context(), "context-key-must-be-ignored") - model, debugEnabled, resolvedProvider, resolvedModel, err := server.resolveComputerUseModel( - ctx, - chat, - aibridgeTestRoute(aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai)), - provider, - modelProvider, - modelName, - modelBuildOptions{ActiveAPIKeyID: apiKeyID}, - ) + resolved, err := server.resolveModelCall(ctx, modelCallSpec{ + purpose: "computer_use", + chat: chat, + fixedModel: &fixedModelCall{providerType: modelProvider, modelName: modelName}, + buildOptions: modelBuildOptions{ActiveAPIKeyID: apiKeyID}, + }) require.NoError(t, err) - require.True(t, model.Valid()) - require.False(t, debugEnabled) - require.EqualValues(t, codersdk.ChatComputerUseProviderOpenAI, resolvedProvider) - require.Equal(t, modelName, resolvedModel) + require.True(t, resolved.model.Valid()) + require.False(t, resolved.debugEnabled) + require.EqualValues(t, codersdk.ChatComputerUseProviderOpenAI, resolved.resolvedProvider) + require.Equal(t, modelName, resolved.resolvedModel) gotProvider, gotSource := factory.recorded() require.Equal(t, "primary-openai", gotProvider) require.Equal(t, aibridge.SourceAgents, gotSource) @@ -674,35 +675,30 @@ func TestAIBridgeComputerUseModelUsesRoute(t *testing.T) { // The computer-use model is a hardcoded default with no config of its own, so // its transport must come from its own client rather than inheriting the chat // model's openai_config. Request preparation reads the same value back. -func TestResolveComputerUseModel_TransportIndependentOfChatConfig(t *testing.T) { +func TestComputerUseModelCall_TransportIndependentOfChatConfig(t *testing.T) { t.Parallel() - providerID := uuid.New() factory := &aibridgeTestFactory{rt: roundTripFunc(func(*http.Request) (*http.Response, error) { t.Fatal("computer use model construction must not send a request") return nil, xerrors.New("unreachable") })} chat := database.Chat{ID: uuid.New(), OwnerID: uuid.New()} - server := &Server{aibridgeTransportFactory: aibridgeTestFactoryPointer(factory)} + server := computerUseTestServer(t, factory) provider := codersdk.ChatComputerUseProviderOpenAI modelProvider, modelName, ok := chattool.DefaultComputerUseModel(provider) require.True(t, ok) - //nolint:dogsled // Only the built model matters for the transport assertion. - model, _, _, _, err := server.resolveComputerUseModel( - t.Context(), - chat, - aibridgeTestRoute(aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai)), - provider, - modelProvider, - modelName, - modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, - ) + resolved, err := server.resolveModelCall(t.Context(), modelCallSpec{ + purpose: "computer_use", + chat: chat, + fixedModel: &fixedModelCall{providerType: modelProvider, modelName: modelName}, + buildOptions: modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, + }) require.NoError(t, err) wantTransport := chatopenai.TransportFor(modelProvider, modelName, nil) - require.Equal(t, wantTransport, model.Transport()) + require.Equal(t, wantTransport, resolved.model.Transport()) // The assertion above only has teeth if an override could have changed the // result for this model. @@ -710,40 +706,28 @@ func TestResolveComputerUseModel_TransportIndependentOfChatConfig(t *testing.T) require.NotEqual(t, wantTransport, chatopenai.TransportFor(modelProvider, modelName, &opposite)) } -func TestResolveComputerUseModel_AIGatewayMissingAPIKeyID(t *testing.T) { +func TestComputerUseModelCall_AIGatewayMissingAPIKeyID(t *testing.T) { t.Parallel() - providerID := uuid.New() factory := &aibridgeTestFactory{rt: roundTripFunc(func(*http.Request) (*http.Response, error) { t.Fatal("transport must not be used without an API key ID") return nil, xerrors.New("unreachable") })} chat := database.Chat{ID: uuid.New(), OwnerID: uuid.New()} - server := &Server{ - aibridgeTransportFactory: aibridgeTestFactoryPointer(factory), - } + server := computerUseTestServer(t, factory) provider := codersdk.ChatComputerUseProviderOpenAI modelProvider, modelName, ok := chattool.DefaultComputerUseModel(provider) require.True(t, ok) - model, debugEnabled, resolvedProvider, resolvedModel, err := server.resolveComputerUseModel( - t.Context(), - chat, - aibridgeTestRoute(aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai)), - provider, - modelProvider, - modelName, - modelBuildOptions{}, // no ActiveAPIKeyID - ) + resolved, err := server.resolveModelCall(t.Context(), modelCallSpec{ + purpose: "computer_use", + chat: chat, + fixedModel: &fixedModelCall{providerType: modelProvider, modelName: modelName}, + }) require.Error(t, err) - require.False(t, model.Valid()) - require.False(t, debugEnabled) - require.Empty(t, resolvedProvider) - require.Empty(t, resolvedModel) - require.Contains(t, err.Error(), fmt.Sprintf( - `resolve computer use model for provider "openai" model %q`, - chattool.ComputerUseOpenAIModelName, - )) + require.False(t, resolved.model.Valid()) + require.False(t, resolved.debugEnabled) + require.Contains(t, err.Error(), "create model") require.Contains(t, err.Error(), "active turn API key ID") } diff --git a/coderd/x/chatd/modelcall.go b/coderd/x/chatd/modelcall.go new file mode 100644 index 0000000000000..ec9949d8db6f8 --- /dev/null +++ b/coderd/x/chatd/modelcall.go @@ -0,0 +1,219 @@ +package chatd + +import ( + "context" + + "charm.land/fantasy" + "golang.org/x/xerrors" + + "cdr.dev/slog/v3" + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/database/dbauthz" + "github.com/coder/coder/v2/coderd/util/ptr" + "github.com/coder/coder/v2/coderd/x/chatd/chatdebug" + "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" + "github.com/coder/coder/v2/codersdk" +) + +const defaultChatMaxOutputTokens = int64(32_000) + +// fixedModelCall selects a provider/model pair that has no config row of its +// own (computer use). +type fixedModelCall struct { + providerType string + modelName string + callConfig codersdk.ChatModelCallConfig +} + +type modelCallSpec struct { + // purpose labels resolver logs only; it does not affect call behavior. + purpose string + chat database.Chat + explicitConfig *database.ChatModelConfig + fixedModel *fixedModelCall + // requestedEffort overrides the config's default reasoning effort. + requestedEffort *string + // chatdScopedRoute resolves the route with chatd scope. Deployment-wide + // override models must route for user-owned chats regardless of the + // caller's actor. + chatdScopedRoute bool + buildOptions modelBuildOptions +} + +func chatRequestedEffort(chat database.Chat) *string { + if !chat.LastReasoningEffort.Valid { + return nil + } + return new(string(chat.LastReasoningEffort.ChatReasoningEffort)) +} + +// modelCallConfigParseError lets the advisor distinguish malformed options +// from route and client failures when deciding whether to fall back. +type modelCallConfigParseError struct{ err error } + +func (e modelCallConfigParseError) Error() string { + return "parse model call config: " + e.err.Error() +} + +func (e modelCallConfigParseError) Unwrap() error { return e.err } + +type resolvedModelCall struct { + model chatprovider.Model + dbConfig database.ChatModelConfig + callConfig codersdk.ChatModelCallConfig + providerOptions fantasy.ProviderOptions + resolvedProvider string + resolvedModel string + route aiGatewayModelRoute + debugEnabled bool +} + +// resolveModelCall is the single pipeline from a spec to a ready model +// client plus the call metadata flows need. +func (p *Server) resolveModelCall(ctx context.Context, spec modelCallSpec) (resolvedModelCall, error) { + out := resolvedModelCall{} + + var modelName string + var configOptions []byte + switch { + case spec.fixedModel != nil: + modelName = spec.fixedModel.modelName + case spec.explicitConfig != nil: + out.dbConfig = *spec.explicitConfig + modelName = out.dbConfig.Model + configOptions = out.dbConfig.Options + default: + dbConfig, err := p.resolveModelConfig(ctx, spec.chat) + if err != nil { + return resolvedModelCall{}, xerrors.Errorf("resolve model config: %w", err) + } + if !dbConfig.Enabled { + return resolvedModelCall{}, xerrors.Errorf("chat model config %s is disabled", dbConfig.ID) + } + out.dbConfig = dbConfig + modelName = dbConfig.Model + configOptions = dbConfig.Options + } + + // clientCallConfig drives client construction; out.callConfig drives + // per-call option derivation and comes from the chat model for computer + // use, whose fixed model has no config of its own. + clientCallConfig, err := parseModelConfigOptions(configOptions) + if err != nil { + return resolvedModelCall{}, modelCallConfigParseError{err: err} + } + if spec.fixedModel != nil { + out.callConfig = spec.fixedModel.callConfig + } else { + out.callConfig = clientCallConfig + } + if out.callConfig.MaxOutputTokens == nil { + out.callConfig.MaxOutputTokens = ptr.Ref(defaultChatMaxOutputTokens) + } + + routeCtx := ctx + if spec.chatdScopedRoute { + //nolint:gocritic // Deployment-wide override models need chatd-scoped provider reads for user-owned chats. + routeCtx = dbauthz.AsChatd(ctx) + } + if spec.fixedModel != nil { + out.route, err = p.resolveModelRouteForProviderType(routeCtx, spec.chat.OwnerID, spec.fixedModel.providerType) + } else { + out.route, err = p.resolveModelRouteForConfig(routeCtx, spec.chat.OwnerID, out.dbConfig) + } + if err != nil { + return resolvedModelCall{}, err + } + + // The resolved identity feeds metadata, logs, and debug labels. The + // client is constructed with the configured model string so gateway + // validation sees the name exactly as configured. + out.resolvedProvider, out.resolvedModel, err = chatprovider.ResolveModelWithProviderHint( + modelName, + out.route.ModelProviderHint, + ) + if err != nil { + return resolvedModelCall{}, xerrors.Errorf("resolve model metadata: %w", err) + } + + debugSvc := p.debugService() + out.debugEnabled = debugSvc != nil && debugSvc.IsEnabled(ctx, spec.chat.ID, spec.chat.OwnerID) + + buildOpts := spec.buildOptions + buildOpts.RecordHTTP = out.debugEnabled + model, err := p.newModel(ctx, modelClientRequest{ + Chat: spec.chat, + ModelName: modelName, + UserAgent: chatprovider.UserAgent(), + ExtraHeaders: chatprovider.CoderHeaders(spec.chat), + CallConfig: clientCallConfig, + }, out.route, buildOpts) + if err != nil { + return resolvedModelCall{}, xerrors.Errorf("create model: %w", err) + } + + if out.debugEnabled { + model = model.WithLanguageModel(chatdebug.WrapModel(model.LanguageModel(), debugSvc, chatdebug.RecorderOptions{ + ChatID: spec.chat.ID, + OwnerID: spec.chat.OwnerID, + Provider: out.resolvedProvider, + Model: out.resolvedModel, + })) + } + out.model = model + + out.providerOptions = out.deriveProviderOptions(out.callConfig, spec.requestedEffort) + + p.logger.Debug(ctx, "resolved model call", + slog.F("purpose", spec.purpose), + slog.F("chat_id", spec.chat.ID), + slog.F("provider", out.resolvedProvider), + slog.F("model", out.resolvedModel), + slog.F("debug_enabled", out.debugEnabled), + ) + return out, nil +} + +func (r resolvedModelCall) newCall() fantasy.Call { + return fantasy.Call{ + ProviderOptions: r.providerOptions, + MaxOutputTokens: r.callConfig.MaxOutputTokens, + Temperature: r.callConfig.Temperature, + TopP: r.callConfig.TopP, + TopK: r.callConfig.TopK, + PresencePenalty: r.callConfig.PresencePenalty, + FrequencyPenalty: r.callConfig.FrequencyPenalty, + } +} + +// compactionSummaryCall follows the resolved call template, except summaries +// must not call tools and must not carry the default output cap: the summary +// request is non-streaming, and the Anthropic SDK rejects non-streaming +// requests whose max_tokens implies a completion longer than ten minutes. +func compactionSummaryCall(resolved resolvedModelCall) fantasy.Call { + call := resolved.newCall() + toolChoiceNone := fantasy.ToolChoiceNone + call.ToolChoice = &toolChoiceNone + call.MaxOutputTokens = nil + return call +} + +// deriveProviderOptions is the only production ProviderOptionsForCall call +// site; callers that mutate the call config after resolution re-derive here. +func (r resolvedModelCall) deriveProviderOptions(callConfig codersdk.ChatModelCallConfig, requestedEffort *string) fantasy.ProviderOptions { + return chatprovider.ProviderOptionsForCall(r.model, callConfig, requestedEffort) +} + +func (r resolvedModelCall) newObjectCall(schemaName, schemaDescription string, maxOutputTokens int64) fantasy.ObjectCall { + return fantasy.ObjectCall{ + SchemaName: schemaName, + SchemaDescription: schemaDescription, + MaxOutputTokens: ptr.Ref(maxOutputTokens), + Temperature: r.callConfig.Temperature, + TopP: r.callConfig.TopP, + TopK: r.callConfig.TopK, + PresencePenalty: r.callConfig.PresencePenalty, + FrequencyPenalty: r.callConfig.FrequencyPenalty, + ProviderOptions: r.providerOptions, + } +} diff --git a/coderd/x/chatd/modelcall_internal_test.go b/coderd/x/chatd/modelcall_internal_test.go new file mode 100644 index 0000000000000..c233b613ddf37 --- /dev/null +++ b/coderd/x/chatd/modelcall_internal_test.go @@ -0,0 +1,80 @@ +package chatd + +import ( + "encoding/json" + "testing" + + "charm.land/fantasy" + fantasyopenai "charm.land/fantasy/providers/openai" + "github.com/google/uuid" + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" + + "cdr.dev/slog/v3/sloggers/slogtest" + "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/codersdk" + "github.com/coder/coder/v2/testutil" +) + +func modelCallSentinelOptions(t *testing.T, user string) json.RawMessage { + t.Helper() + raw, err := json.Marshal(codersdk.ChatModelCallConfig{ + ProviderOptions: &codersdk.ChatModelProviderOptions{ + OpenAI: &codersdk.ChatModelOpenAIProviderOptions{ + User: ptr.Ref(user), + }, + }, + }) + require.NoError(t, err) + return raw +} + +// The transport decides which of the two OpenAI option shapes derivation +// produces, so both are accepted. +func requireOpenAIUserOption(t *testing.T, options fantasy.ProviderOptions, user string) { + t.Helper() + switch opts := options[fantasyopenai.Name].(type) { + case *fantasyopenai.ResponsesProviderOptions: + require.NotNil(t, opts.User) + require.Equal(t, user, *opts.User) + case *fantasyopenai.ProviderOptions: + require.NotNil(t, opts.User) + require.Equal(t, user, *opts.User) + default: + t.Fatalf("unexpected openai provider options type %T", opts) + } +} + +func TestResolveModelCallDerivesProviderOptions(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitShort) + ctrl := gomock.NewController(t) + db := dbmock.NewMockStore(ctrl) + logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) + chat, _ := titleOverrideTestChatAndMessages(t) + providerID := uuid.New() + config := titleOverrideModelConfig("gpt-4o-mini", true) + config.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true} + config.Options = modelCallSentinelOptions(t, "summary-options-sentinel") + chat.LastModelConfigID = config.ID + + db.EXPECT().GetChatModelConfigByID(gomock.Any(), config.ID).Return(config, nil) + db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai), nil).AnyTimes() + db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ + ProviderID: providerID, + APIKey: "test-key", + }}, nil).AnyTimes() + + server := titleOverrideTestServer(db, logger) + resolved, err := server.resolveModelCall(ctx, modelCallSpec{ + purpose: "chat_summary", + chat: chat, + buildOptions: modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, + }) + require.NoError(t, err) + requireOpenAIUserOption(t, resolved.providerOptions, "summary-options-sentinel") + requireOpenAIUserOption(t, summaryObjectCall(resolved).ProviderOptions, "summary-options-sentinel") +} diff --git a/coderd/x/chatd/quickgen.go b/coderd/x/chatd/quickgen.go index ccc73b3ed30d6..9a0ab481e9930 100644 --- a/coderd/x/chatd/quickgen.go +++ b/coderd/x/chatd/quickgen.go @@ -2,7 +2,6 @@ package chatd import ( "context" - "encoding/json" "errors" "fmt" "net/http" @@ -137,13 +136,12 @@ var preferredTitleModels = []struct { {fantasyvercel.Name, "anthropic/claude-haiku-4.5"}, } +// Debug attribution uses configured identities for title calls and the +// resolved identity for status-label calls. type shortTextCandidate struct { - provider string - model string - route aiGatewayModelRoute - lm chatprovider.Model - providerOptions fantasy.ProviderOptions - configOptions json.RawMessage + provider string + model string + resolved resolvedModelCall } func selectPreferredConfiguredShortTextModelConfig( @@ -231,7 +229,11 @@ func (p *Server) GenerateChatTitleAsync(ctx context.Context, chat database.Chat) } modelOpts := modelBuildOptions{ActiveAPIKeyID: apiKeyID} turnCtx := titleCtx - model, modelConfig, route, _, _, _, err := p.resolveChatModel(turnCtx, chat, modelOpts) + fallback, err := p.resolveModelCall(turnCtx, modelCallSpec{ + purpose: "title", + chat: chat, + buildOptions: modelOpts, + }) if err != nil { logger.Debug(titleCtx, "failed to resolve model for automatic title generation", slog.Error(err), @@ -243,10 +245,7 @@ func (p *Server) GenerateChatTitleAsync(ctx context.Context, chat database.Chat) chat, messages, pasteText, - string(route.Provider.Type), - modelConfig, - model, - route, + fallback, modelOpts, &generatedChatTitle{}, logger, @@ -274,10 +273,7 @@ func (p *Server) maybeGenerateChatTitle( chat database.Chat, messages []database.ChatMessage, pasteText map[uuid.UUID]string, - fallbackProvider string, - fallbackConfig database.ChatModelConfig, - fallbackModel chatprovider.Model, - fallbackRoute aiGatewayModelRoute, + fallback resolvedModelCall, modelOpts modelBuildOptions, generatedTitle *generatedChatTitle, logger slog.Logger, @@ -292,7 +288,7 @@ func (p *Server) maybeGenerateChatTitle( titleCtx, cancel := context.WithTimeout(ctx, 30*time.Second) defer cancel() - overrideConfig, overrideModel, overrideRoute, overrideSet, overrideErr := p.resolveTitleGenerationModelOverride( + overrideResolved, overrideSet, overrideErr := p.resolveTitleGenerationModelOverride( titleCtx, chat, modelOpts, @@ -313,25 +309,14 @@ func (p *Server) maybeGenerateChatTitle( ) } - var candidate shortTextCandidate + selected := fallback if overrideSet { - candidate = shortTextCandidate{ - provider: string(overrideRoute.Provider.Type), - model: overrideConfig.Model, - route: overrideRoute, - lm: overrideModel, - providerOptions: p.titleGenerationProviderOptions(ctx, overrideModel, overrideConfig), - configOptions: overrideConfig.Options, - } - } else { - candidate = shortTextCandidate{ - provider: fallbackProvider, - model: fallbackConfig.Model, - route: fallbackRoute, - lm: fallbackModel, - providerOptions: p.titleGenerationProviderOptions(ctx, fallbackModel, fallbackConfig), - configOptions: fallbackConfig.Options, - } + selected = overrideResolved + } + candidate := shortTextCandidate{ + provider: string(selected.route.Provider.Type), + model: selected.dbConfig.Model, + resolved: selected, } var historyTipMessageID int64 @@ -355,15 +340,13 @@ func (p *Server) maybeGenerateChatTitle( ) candidateCtx := titleCtx - candidateModel := candidate.lm finishDebugRun := func(error) {} if debugEnabled { - candidateCtx, candidateModel, finishDebugRun = p.prepareQuickgenDebugCandidate( + candidateCtx, finishDebugRun = prepareQuickgenDebugCandidate( titleCtx, chat, debugSvc, candidate, - modelOpts, chatdebug.KindTitleGeneration, triggerMessageID, historyTipMessageID, @@ -372,7 +355,7 @@ func (p *Server) maybeGenerateChatTitle( ) } - title, err := generateTitle(candidateCtx, candidateModel.LanguageModel(), candidate.providerOptions, input) + title, err := generateTitle(candidateCtx, candidate.resolved.model.LanguageModel(), titleObjectCall(candidate.resolved), input) finishDebugRun(err) if err != nil { if overrideSet { @@ -411,91 +394,24 @@ func (p *Server) maybeGenerateChatTitle( p.publishChatPubsubEvent(chat, codersdk.ChatWatchEventKindTitleChange, nil) } -func (p *Server) titleGenerationProviderOptions( - ctx context.Context, - model chatprovider.Model, - 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), - ) - } - } - return chatprovider.ProviderOptionsForCall(model, callConfig, nil) -} +const titleMaxOutputTokens = int64(256) -func (p *Server) newQuickgenDebugModel( - ctx context.Context, - chat database.Chat, - debugSvc *chatdebug.Service, - provider string, - model string, - route aiGatewayModelRoute, - modelOpts modelBuildOptions, - configOptions json.RawMessage, -) (chatprovider.Model, error) { - debugOpts := modelOpts - debugOpts.RecordHTTP = true - debugModel, err := p.newModel(ctx, modelClientRequest{ - Chat: chat, - ModelName: model, - UserAgent: chatprovider.UserAgent(), - ExtraHeaders: chatprovider.CoderHeaders(chat), - ConfigOptions: configOptions, - }, route, debugOpts) - if err != nil { - return chatprovider.Model{}, err - } - - return debugModel.WithLanguageModel(chatdebug.WrapModel(debugModel.LanguageModel(), debugSvc, chatdebug.RecorderOptions{ - ChatID: chat.ID, - OwnerID: chat.OwnerID, - Provider: provider, - Model: model, - })), nil +func titleObjectCall(resolved resolvedModelCall) fantasy.ObjectCall { + return resolved.newObjectCall("propose_title", "Propose a short chat title.", titleMaxOutputTokens) } -func (p *Server) prepareQuickgenDebugCandidate( +func prepareQuickgenDebugCandidate( ctx context.Context, chat database.Chat, debugSvc *chatdebug.Service, candidate shortTextCandidate, - modelOpts modelBuildOptions, kind chatdebug.RunKind, triggerMessageID int64, historyTipMessageID int64, seedSummary map[string]any, logger slog.Logger, -) (context.Context, chatprovider.Model, func(error)) { +) (context.Context, func(error)) { finishDebugRun := func(error) {} - if debugSvc == nil { - return ctx, candidate.lm, finishDebugRun - } - - debugModel, err := p.newQuickgenDebugModel( - ctx, - chat, - debugSvc, - candidate.provider, - candidate.model, - candidate.route, - modelOpts, - candidate.configOptions, - ) - if err != nil { - logger.Warn(ctx, "failed to build short-text debug model", - slog.F("chat_id", chat.ID), - slog.F("run_kind", kind), - slog.F("provider", candidate.provider), - slog.F("model", candidate.model), - slog.Error(err), - ) - return ctx, candidate.lm, finishDebugRun - } // Debug instrumentation must not eat into the quickgen budget // (30s titleCtx / summaryCtx on the caller). Detach and bound @@ -524,7 +440,7 @@ func (p *Server) prepareQuickgenDebugCandidate( slog.F("model", candidate.model), slog.Error(err), ) - return ctx, candidate.lm, finishDebugRun + return ctx, finishDebugRun } runContext := chatdebugRunContext(run) @@ -545,7 +461,24 @@ func (p *Server) prepareQuickgenDebugCandidate( ) } } - return runCtx, debugModel, finishDebugRun + return runCtx, finishDebugRun +} + +func quickgenPrompt(systemPrompt, userInput string) fantasy.Prompt { + return fantasy.Prompt{ + { + Role: fantasy.MessageRoleSystem, + Content: []fantasy.MessagePart{ + fantasy.TextPart{Text: systemPrompt}, + }, + }, + { + Role: fantasy.MessageRoleUser, + Content: []fantasy.MessagePart{ + fantasy.TextPart{Text: userInput}, + }, + }, + } } // generateTitle calls the model with a title-generation system prompt @@ -554,10 +487,10 @@ func (p *Server) prepareQuickgenDebugCandidate( func generateTitle( ctx context.Context, model fantasy.LanguageModel, - providerOptions fantasy.ProviderOptions, + call fantasy.ObjectCall, input string, ) (string, error) { - title, err := generateStructuredTitle(ctx, model, providerOptions, titleGenerationPrompt, input) + title, err := generateStructuredTitle(ctx, model, call, titleGenerationPrompt, input) if err != nil { return "", err } @@ -567,14 +500,14 @@ func generateTitle( func generateStructuredTitle( ctx context.Context, model fantasy.LanguageModel, - providerOptions fantasy.ProviderOptions, + call fantasy.ObjectCall, systemPrompt string, userInput string, ) (string, error) { title, _, err := generateStructuredTitleWithUsage( ctx, model, - providerOptions, + call, systemPrompt, userInput, ) @@ -587,7 +520,7 @@ func generateStructuredTitle( func generateStructuredTitleWithUsage( ctx context.Context, model fantasy.LanguageModel, - providerOptions fantasy.ProviderOptions, + call fantasy.ObjectCall, systemPrompt string, userInput string, ) (string, fantasy.Usage, error) { @@ -596,29 +529,8 @@ func generateStructuredTitleWithUsage( return "", fantasy.Usage{}, xerrors.New("title input was empty") } - prompt := fantasy.Prompt{ - { - Role: fantasy.MessageRoleSystem, - Content: []fantasy.MessagePart{ - fantasy.TextPart{Text: systemPrompt}, - }, - }, - { - Role: fantasy.MessageRoleUser, - Content: []fantasy.MessagePart{ - fantasy.TextPart{Text: userInput}, - }, - }, - } - - var maxOutputTokens int64 = 256 - result, err := generateQuickgenObject[generatedTitle](ctx, model, fantasy.ObjectCall{ - Prompt: prompt, - SchemaName: "propose_title", - SchemaDescription: "Propose a short chat title.", - MaxOutputTokens: &maxOutputTokens, - ProviderOptions: providerOptions, - }) + call.Prompt = quickgenPrompt(systemPrompt, userInput) + result, err := generateQuickgenObject[generatedTitle](ctx, model, call) if err != nil { var usage fantasy.Usage var noObjErr *fantasy.NoObjectGeneratedError @@ -926,7 +838,7 @@ func generateManualTitle( messages []database.ChatMessage, pasteText map[uuid.UUID]string, fallbackModel fantasy.LanguageModel, - providerOptions fantasy.ProviderOptions, + call fantasy.ObjectCall, ) (string, error) { turns := extractManualTitleTurns(messages, pasteText) selected := selectManualTitleTurnIndexes(turns) @@ -957,7 +869,7 @@ func generateManualTitle( title, _, err := generateStructuredTitleWithUsage( titleCtx, fallbackModel, - providerOptions, + call, systemPrompt, userInput, ) @@ -1088,12 +1000,17 @@ func boundTranscriptHeadTail(lines []string, maxRunes int) string { return out.String() } +func summaryObjectCall(resolved resolvedModelCall) fantasy.ObjectCall { + return resolved.newObjectCall("chat_summary", "Summarize the whole chat in 1-3 sentences.", summaryMaxOutputTokens) +} + // generateChatSummary generates a 1-3 sentence whole-chat summary from a // transcript. A blank or invalid result returns an error so callers preserve // any existing summary rather than clearing it. func generateChatSummary( ctx context.Context, model fantasy.LanguageModel, + call fantasy.ObjectCall, transcript string, ) (string, fantasy.Usage, error) { transcript = strings.TrimSpace(transcript) @@ -1101,31 +1018,11 @@ func generateChatSummary( return "", fantasy.Usage{}, xerrors.New("chat summary transcript was empty") } - prompt := fantasy.Prompt{ - { - Role: fantasy.MessageRoleSystem, - Content: []fantasy.MessagePart{ - fantasy.TextPart{Text: chatSummaryGenerationPrompt}, - }, - }, - { - Role: fantasy.MessageRoleUser, - Content: []fantasy.MessagePart{ - fantasy.TextPart{Text: transcript}, - }, - }, - } - - maxOutputTokens := int64(summaryMaxOutputTokens) + call.Prompt = quickgenPrompt(chatSummaryGenerationPrompt, transcript) var result *fantasy.ObjectResult[generatedChatSummary] err := chatretry.Retry(ctx, func(retryCtx context.Context) error { var genErr error - result, genErr = object.Generate[generatedChatSummary](retryCtx, model, fantasy.ObjectCall{ - Prompt: prompt, - SchemaName: "chat_summary", - SchemaDescription: "Summarize the whole chat in 1-3 sentences.", - MaxOutputTokens: &maxOutputTokens, - }) + result, genErr = object.Generate[generatedChatSummary](retryCtx, model, call) return genErr }, nil) if err != nil { @@ -1321,19 +1218,19 @@ const turnStatusLabelPrompt = "You write compact chat status labels for a sideba "Prefer short action or state phrases such as Finished, Submitted, Fixed, Testing, Still working, or Waiting for. " + "No quotes, emoji, markdown, or trailing punctuation." -// generateTurnStatusLabel produces a short turn status label using the -// caller-supplied fallback model. Returns "" on any failure. -func (p *Server) generateTurnStatusLabel( +const turnStatusLabelMaxOutputTokens = int64(64) + +func turnStatusLabelObjectCall(resolved resolvedModelCall) fantasy.ObjectCall { + return resolved.newObjectCall("propose_turn_status_label", "Propose a compact chat status label.", turnStatusLabelMaxOutputTokens) +} + +// generateTurnStatusLabel returns an empty string if generation fails. +func generateTurnStatusLabel( ctx context.Context, chat database.Chat, status database.ChatStatus, assistantText string, - fallbackProvider string, - fallbackModelName string, - fallbackModel chatprovider.Model, - fallbackRoute aiGatewayModelRoute, - modelOpts modelBuildOptions, - configOptions json.RawMessage, + resolved resolvedModelCall, logger slog.Logger, debugSvc *chatdebug.Service, triggerMessageID int64, @@ -1350,25 +1247,21 @@ func (p *Server) generateTurnStatusLabel( "\n\nAgent's latest message:\n" + assistantText candidate := shortTextCandidate{ - provider: fallbackProvider, - model: fallbackModelName, - route: fallbackRoute, - lm: fallbackModel, - configOptions: configOptions, + provider: resolved.resolvedProvider, + model: resolved.resolvedModel, + resolved: resolved, } statusSeedSummary := chatdebug.SeedSummary("Turn status label") candidateCtx := labelCtx - candidateModel := candidate.lm finishDebugRun := func(error) {} if debugEnabled { - candidateCtx, candidateModel, finishDebugRun = p.prepareQuickgenDebugCandidate( + candidateCtx, finishDebugRun = prepareQuickgenDebugCandidate( labelCtx, chat, debugSvc, candidate, - modelOpts, chatdebug.KindQuickgen, triggerMessageID, historyTipMessageID, @@ -1379,7 +1272,8 @@ func (p *Server) generateTurnStatusLabel( generatedLabel, err := generateStructuredTurnStatusLabel( candidateCtx, - candidateModel.LanguageModel(), + candidate.resolved.model.LanguageModel(), + turnStatusLabelObjectCall(resolved), turnStatusLabelPrompt, input, ) @@ -1396,6 +1290,7 @@ func (p *Server) generateTurnStatusLabel( func generateStructuredTurnStatusLabel( ctx context.Context, model fantasy.LanguageModel, + call fantasy.ObjectCall, systemPrompt string, userInput string, ) (string, error) { @@ -1404,28 +1299,8 @@ func generateStructuredTurnStatusLabel( return "", xerrors.New("turn status label input was empty") } - prompt := fantasy.Prompt{ - { - Role: fantasy.MessageRoleSystem, - Content: []fantasy.MessagePart{ - fantasy.TextPart{Text: systemPrompt}, - }, - }, - { - Role: fantasy.MessageRoleUser, - Content: []fantasy.MessagePart{ - fantasy.TextPart{Text: userInput}, - }, - }, - } - - var maxOutputTokens int64 = 64 - result, err := generateQuickgenObject[generatedTurnStatusLabel](ctx, model, fantasy.ObjectCall{ - Prompt: prompt, - SchemaName: "propose_turn_status_label", - SchemaDescription: "Propose a compact chat status label.", - MaxOutputTokens: &maxOutputTokens, - }) + call.Prompt = quickgenPrompt(systemPrompt, userInput) + result, err := generateQuickgenObject[generatedTurnStatusLabel](ctx, model, call) if err != nil { return "", xerrors.Errorf("generate structured turn status label: %w", err) } diff --git a/coderd/x/chatd/quickgen_internal_test.go b/coderd/x/chatd/quickgen_internal_test.go index a761bebc96de5..af2ee1c32b7e0 100644 --- a/coderd/x/chatd/quickgen_internal_test.go +++ b/coderd/x/chatd/quickgen_internal_test.go @@ -586,10 +586,10 @@ func TestMaybeGenerateChatTitlePreservesUpdatedAt(t *testing.T) { chat, []database.ChatMessage{message}, nil, - "openai", - database.ChatModelConfig{Model: "test-model"}, - chatprovider.NewModel(model, nil), - aiGatewayModelRoute{}, + resolvedModelCall{ + model: chatprovider.NewModel(model, nil), + dbConfig: database.ChatModelConfig{Model: "test-model"}, + }, modelBuildOptions{}, generated, logger, @@ -650,15 +650,20 @@ func TestMaybeGenerateChatTitleAppliesModelConfigReasoningEffort(t *testing.T) { logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) server := titleOverrideTestServer(db, logger) + fallbackModel := chatprovider.NewModel(model, nil) + fallbackConfig := database.ChatModelConfig{Model: "gpt-4o-mini", Options: modelConfigRaw} + callConfig, err := parseModelConfigOptions(fallbackConfig.Options) + require.NoError(t, err) server.maybeGenerateChatTitle( ctx, chat, messages, nil, - fantasyopenai.Name, - database.ChatModelConfig{Model: "gpt-4o-mini", Options: modelConfigRaw}, - chatprovider.NewModel(model, nil), - aiGatewayModelRoute{}, + resolvedModelCall{ + model: fallbackModel, + dbConfig: fallbackConfig, + providerOptions: chatprovider.ProviderOptionsForCall(fallbackModel, callConfig, nil), + }, modelBuildOptions{}, &generatedChatTitle{}, logger, @@ -710,7 +715,7 @@ func Test_generateManualTitle_UsesTimeout(t *testing.T) { messages, nil, model, - nil, + titleObjectCall(resolvedModelCall{}), ) require.NoError(t, err) require.Equal(t, "Refresh title", title) @@ -748,7 +753,7 @@ func Test_generateManualTitle_TruncatesFirstUserInput(t *testing.T) { messages, nil, model, - nil, + titleObjectCall(resolvedModelCall{}), ) require.NoError(t, err) } @@ -783,7 +788,7 @@ func Test_generateManualTitle_ErrorsOnEmptyNormalizedTitle(t *testing.T) { messages, nil, model, - nil, + titleObjectCall(resolvedModelCall{}), ) require.ErrorContains(t, err, "generated title was empty") } @@ -885,7 +890,7 @@ func TestGenerateStructuredTitleWithUsage_OpenAICompatibleRequiredToolChoice(t * title, _, err := generateStructuredTitleWithUsage( t.Context(), model.LanguageModel(), - nil, + titleObjectCall(resolvedModelCall{}), titleGenerationPrompt, "summarize failed workspace build logs", ) @@ -930,7 +935,7 @@ func TestGenerateStructuredTitleWithUsage_DropsRejectedTemperature(t *testing.T) title, _, err := generateStructuredTitleWithUsage( t.Context(), model, - nil, + titleObjectCall(resolvedModelCall{}), titleGenerationPrompt, "summarize failed workspace build logs", ) @@ -1025,13 +1030,17 @@ func TestGenerateStructuredTurnStatusLabel(t *testing.T) { model := &chattest.FakeModel{ GenerateObjectFn: func(_ context.Context, call fantasy.ObjectCall) (*fantasy.ObjectResponse, error) { require.Equal(t, "propose_turn_status_label", call.SchemaName) + require.NotNil(t, call.MaxOutputTokens) + require.Equal(t, turnStatusLabelMaxOutputTokens, *call.MaxOutputTokens) + require.NotNil(t, call.Temperature) + require.Equal(t, quickgenTemperature, *call.Temperature) return &fantasy.ObjectResponse{ Object: map[string]any{"label": "Submitted PR"}, }, nil }, } - label, err := generateStructuredTurnStatusLabel(t.Context(), model, turnStatusLabelPrompt, "done") + label, err := generateStructuredTurnStatusLabel(t.Context(), model, turnStatusLabelObjectCall(resolvedModelCall{}), turnStatusLabelPrompt, "done") require.NoError(t, err) require.Equal(t, "Submitted PR", label) }) @@ -1042,7 +1051,7 @@ func TestGenerateStructuredTurnStatusLabel(t *testing.T) { server, requests := newOpenAICompatStructuredOutputServer(t, "propose_turn_status_label", `{"label":"Submitted PR"}`) model := openAICompatTestModel(t, server.URL) - label, err := generateStructuredTurnStatusLabel(t.Context(), model.LanguageModel(), turnStatusLabelPrompt, "done") + label, err := generateStructuredTurnStatusLabel(t.Context(), model.LanguageModel(), turnStatusLabelObjectCall(resolvedModelCall{}), turnStatusLabelPrompt, "done") require.NoError(t, err) require.Equal(t, "Submitted PR", label) require.Len(t, requests, 1) @@ -1069,7 +1078,7 @@ func TestGenerateStructuredTurnStatusLabel(t *testing.T) { }, } - label, err := generateStructuredTurnStatusLabel(t.Context(), model, turnStatusLabelPrompt, "done") + label, err := generateStructuredTurnStatusLabel(t.Context(), model, turnStatusLabelObjectCall(resolvedModelCall{}), turnStatusLabelPrompt, "done") require.NoError(t, err) require.Equal(t, "Submitted PR", label) require.Equal(t, []bool{true, false}, sawTemperature, @@ -1091,7 +1100,7 @@ func TestGenerateStructuredTurnStatusLabel(t *testing.T) { }, } - _, err := generateStructuredTurnStatusLabel(t.Context(), model, turnStatusLabelPrompt, "done") + _, err := generateStructuredTurnStatusLabel(t.Context(), model, turnStatusLabelObjectCall(resolvedModelCall{}), turnStatusLabelPrompt, "done") require.ErrorContains(t, err, "JSON schema is invalid") require.Equal(t, 1, calls, "bad requests unrelated to temperature should not trigger a second attempt") @@ -1108,7 +1117,7 @@ func TestGenerateStructuredTurnStatusLabel(t *testing.T) { }, } - _, err := generateStructuredTurnStatusLabel(t.Context(), model, turnStatusLabelPrompt, "done") + _, err := generateStructuredTurnStatusLabel(t.Context(), model, turnStatusLabelObjectCall(resolvedModelCall{}), turnStatusLabelPrompt, "done") require.ErrorContains(t, err, "generated turn status label was invalid") }) @@ -1116,7 +1125,7 @@ func TestGenerateStructuredTurnStatusLabel(t *testing.T) { t.Parallel() model := &chattest.FakeModel{} - _, err := generateStructuredTurnStatusLabel(t.Context(), model, turnStatusLabelPrompt, " ") + _, err := generateStructuredTurnStatusLabel(t.Context(), model, turnStatusLabelObjectCall(resolvedModelCall{}), turnStatusLabelPrompt, " ") require.ErrorContains(t, err, "turn status label input was empty") }) } diff --git a/coderd/x/chatd/subagent_internal_test.go b/coderd/x/chatd/subagent_internal_test.go index 9f1e55c54fd89..fca8c14296d12 100644 --- a/coderd/x/chatd/subagent_internal_test.go +++ b/coderd/x/chatd/subagent_internal_test.go @@ -539,13 +539,12 @@ func TestResolveChatModel_AIProviderDisabled(t *testing.T) { LastModelConfigID: modelConfig.ID, }) - model, config, _, debugEnabled, resolvedProvider, resolvedModel, err := server.resolveChatModel(ctx, chat, modelBuildOptions{}) + resolved, err := server.resolveModelCall(ctx, modelCallSpec{ + purpose: "standard_turn", + chat: chat, + }) require.ErrorContains(t, err, "is disabled") - require.False(t, model.Valid()) - require.Equal(t, database.ChatModelConfig{}, config) - require.False(t, debugEnabled) - require.Empty(t, resolvedProvider) - require.Empty(t, resolvedModel) + require.Equal(t, resolvedModelCall{}, resolved) } func TestResolveUserProviderAPIKeys_PreservesAnthropicKeyFromDBProvider(t *testing.T) { diff --git a/coderd/x/chatd/title_override.go b/coderd/x/chatd/title_override.go index 4056fdcfe13a4..e5214c30b3b33 100644 --- a/coderd/x/chatd/title_override.go +++ b/coderd/x/chatd/title_override.go @@ -60,10 +60,10 @@ func (p *Server) resolveTitleGenerationModelOverride( ctx context.Context, chat database.Chat, modelOpts modelBuildOptions, -) (database.ChatModelConfig, chatprovider.Model, aiGatewayModelRoute, bool, error) { +) (resolvedModelCall, bool, error) { raw, err := readTitleGenerationModelOverride(ctx, p.db) if err != nil { - return database.ChatModelConfig{}, chatprovider.Model{}, aiGatewayModelRoute{}, false, xerrors.Errorf( + return resolvedModelCall{}, false, xerrors.Errorf( "read title generation model override: %w", err, ) @@ -81,30 +81,25 @@ func (p *Server) resolveTitleGenerationModelOverride( modelOverrideFailureModeHard, ) if err != nil { - return database.ChatModelConfig{}, chatprovider.Model{}, aiGatewayModelRoute{}, overrideSet, err + return resolvedModelCall{}, overrideSet, err } if !overrideSet { - return database.ChatModelConfig{}, chatprovider.Model{}, aiGatewayModelRoute{}, false, nil + return resolvedModelCall{}, 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) + resolved, err := p.resolveModelCall(ctx, modelCallSpec{ + purpose: "title", + chat: chat, + explicitConfig: &modelConfig, + chatdScopedRoute: true, + buildOptions: modelOpts, + }) if err != nil { - return database.ChatModelConfig{}, chatprovider.Model{}, aiGatewayModelRoute{}, true, err - } - model, err := p.newModel(ctx, modelClientRequest{ - Chat: chat, - ModelName: modelConfig.Model, - UserAgent: chatprovider.UserAgent(), - ExtraHeaders: chatprovider.CoderHeaders(chat), - ConfigOptions: modelConfig.Options, - }, route, modelOpts) - if err != nil { - return database.ChatModelConfig{}, chatprovider.Model{}, aiGatewayModelRoute{}, true, xerrors.Errorf( + return resolvedModelCall{}, true, xerrors.Errorf( "create title generation model override: %w", err, ) } - return modelConfig, model, route, true, nil + return resolved, true, nil } diff --git a/coderd/x/chatd/title_override_internal_test.go b/coderd/x/chatd/title_override_internal_test.go index 36ebef7226036..7499066f65e54 100644 --- a/coderd/x/chatd/title_override_internal_test.go +++ b/coderd/x/chatd/title_override_internal_test.go @@ -69,10 +69,10 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideUnset(t *testing.T) { chat, messages, nil, - "openai", - database.ChatModelConfig{Model: "fallback-chat-model"}, - chatprovider.NewModel(fallbackModel, nil), - aiGatewayModelRoute{}, + resolvedModelCall{ + model: chatprovider.NewModel(fallbackModel, nil), + dbConfig: database.ChatModelConfig{Model: "fallback-chat-model"}, + }, modelBuildOptions{}, generated, logger, @@ -119,10 +119,10 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideReadDBError(t *testing.T) chat, messages, nil, - "openai", - database.ChatModelConfig{Model: "fallback-chat-model"}, - chatprovider.NewModel(fallbackModel, nil), - aiGatewayModelRoute{}, + resolvedModelCall{ + model: chatprovider.NewModel(fallbackModel, nil), + dbConfig: database.ChatModelConfig{Model: "fallback-chat-model"}, + }, modelBuildOptions{}, generated, logger, @@ -168,10 +168,10 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideMalformedFallsThrough(t * chat, messages, nil, - "openai", - database.ChatModelConfig{Model: "fallback-chat-model"}, - chatprovider.NewModel(fallbackModel, nil), - aiGatewayModelRoute{}, + resolvedModelCall{ + model: chatprovider.NewModel(fallbackModel, nil), + dbConfig: database.ChatModelConfig{Model: "fallback-chat-model"}, + }, modelBuildOptions{}, generated, logger, @@ -255,10 +255,10 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideSetUsable(t *testing.T) { chat, messages, nil, - "openai", - database.ChatModelConfig{Model: "fallback-chat-model"}, - chatprovider.NewModel(fallbackModel, nil), - aiGatewayModelRoute{}, + resolvedModelCall{ + model: chatprovider.NewModel(fallbackModel, nil), + dbConfig: database.ChatModelConfig{Model: "fallback-chat-model"}, + }, modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, generated, logger, @@ -297,10 +297,10 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideSetUnusableSkips(t *testi chat, messages, nil, - "openai", - database.ChatModelConfig{Model: "fallback-chat-model"}, - chatprovider.NewModel(fallbackModel, nil), - aiGatewayModelRoute{}, + resolvedModelCall{ + model: chatprovider.NewModel(fallbackModel, nil), + dbConfig: database.ChatModelConfig{Model: "fallback-chat-model"}, + }, modelBuildOptions{}, generated, logger, @@ -351,10 +351,10 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideCallFailureSkipsFallback( chat, messages, nil, - "openai", - database.ChatModelConfig{Model: "fallback-chat-model"}, - chatprovider.NewModel(fallbackModel, nil), - aiGatewayModelRoute{}, + resolvedModelCall{ + model: chatprovider.NewModel(fallbackModel, nil), + dbConfig: database.ChatModelConfig{Model: "fallback-chat-model"}, + }, modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, generated, logger, @@ -390,15 +390,15 @@ func TestResolveManualTitleModel_TitleGenerationOverrideUnset(t *testing.T) { db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai), nil).AnyTimes() server := titleOverrideTestServer(db, logger) - model, gotConfig, err := server.resolveManualTitleModel( + resolved, err := server.resolveManualTitleModel( ctx, db, chat, modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, ) require.NoError(t, err) - require.True(t, model.Valid()) - require.Equal(t, preferredConfig, gotConfig) + require.True(t, resolved.model.Valid()) + require.Equal(t, preferredConfig, resolved.dbConfig) } func TestResolveManualTitleModel_TitleGenerationOverrideUnsetAIProvider(t *testing.T) { @@ -439,15 +439,15 @@ func TestResolveManualTitleModel_TitleGenerationOverrideUnsetAIProvider(t *testi }}, nil).AnyTimes() server := titleOverrideTestServer(db, logger) - model, gotConfig, err := server.resolveManualTitleModel( + resolved, err := server.resolveManualTitleModel( ctx, db, chat, modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, ) require.NoError(t, err) - require.True(t, model.Valid()) - require.Equal(t, preferredConfig, gotConfig) + require.True(t, resolved.model.Valid()) + require.Equal(t, preferredConfig, resolved.dbConfig) } func TestResolveManualTitleModel_TitleGenerationOverrideReadDBError(t *testing.T) { @@ -474,15 +474,15 @@ func TestResolveManualTitleModel_TitleGenerationOverrideReadDBError(t *testing.T db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai), nil).AnyTimes() server := titleOverrideTestServer(db, logger) - model, gotConfig, err := server.resolveManualTitleModel( + resolved, err := server.resolveManualTitleModel( ctx, db, chat, modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, ) require.NoError(t, err) - require.True(t, model.Valid()) - require.Equal(t, preferredConfig, gotConfig) + require.True(t, resolved.model.Valid()) + require.Equal(t, preferredConfig, resolved.dbConfig) } func TestResolveManualTitleModel_TitleGenerationOverrideSetUsable(t *testing.T) { @@ -506,15 +506,15 @@ func TestResolveManualTitleModel_TitleGenerationOverrideSetUsable(t *testing.T) }}, nil).AnyTimes() server := titleOverrideTestServer(db, logger) - model, gotConfig, err := server.resolveManualTitleModel( + resolved, err := server.resolveManualTitleModel( ctx, db, chat, modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, ) require.NoError(t, err) - require.True(t, model.Valid()) - require.Equal(t, overrideConfig, gotConfig) + require.True(t, resolved.model.Valid()) + require.Equal(t, overrideConfig, resolved.dbConfig) } func TestResolveManualTitleModel_TitleGenerationOverrideMissingCredentials(t *testing.T) { @@ -540,7 +540,7 @@ func TestResolveManualTitleModel_TitleGenerationOverrideMissingCredentials(t *te db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return(nil, nil).AnyTimes() server := titleOverrideTestServer(db, logger) - model, gotConfig, err := server.resolveManualTitleModel( + resolved, err := server.resolveManualTitleModel( ctx, db, chat, @@ -549,8 +549,7 @@ func TestResolveManualTitleModel_TitleGenerationOverrideMissingCredentials(t *te require.Error(t, err) require.ErrorContains(t, err, "resolve manual title generation model override") require.ErrorContains(t, err, "credentials are unavailable") - require.False(t, model.Valid()) - require.Equal(t, database.ChatModelConfig{}, gotConfig) + require.Equal(t, resolvedModelCall{}, resolved) } func TestGenerateManualTitleCandidate_UsesSyntheticAPIKey(t *testing.T) { @@ -565,6 +564,7 @@ func TestGenerateManualTitleCandidate_UsesSyntheticAPIKey(t *testing.T) { overrideConfig := titleOverrideModelConfig("gpt-4.1", true) providerID := uuid.New() overrideConfig.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true} + overrideConfig.Options = modelCallSentinelOptions(t, "title-options-sentinel") provider := database.AIProvider{ ID: providerID, Name: "primary-openai", @@ -574,9 +574,13 @@ func TestGenerateManualTitleCandidate_UsesSyntheticAPIKey(t *testing.T) { apiKeyID := uuid.NewString() wantTitle := "Synthetic title" seenAPIKeyID := make(chan string, 1) + seenBody := make(chan []byte, 1) factory := &aibridgeTestFactory{rt: roundTripFunc(func(req *http.Request) (*http.Response, error) { delegatedID, _ := aibridge.DelegatedAPIKeyIDFromContext(req.Context()) seenAPIKeyID <- delegatedID + bodyBytes, err := io.ReadAll(req.Body) + require.NoError(t, err) + seenBody <- bodyBytes 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}}` return &http.Response{ @@ -620,6 +624,10 @@ func TestGenerateManualTitleCandidate_UsesSyntheticAPIKey(t *testing.T) { require.NoError(t, err) require.Equal(t, wantTitle, title) require.Equal(t, apiKeyID, testutil.RequireReceive(ctx, t, seenAPIKeyID)) + + var raw map[string]any + require.NoError(t, json.Unmarshal(testutil.RequireReceive(ctx, t, seenBody), &raw)) + require.Equal(t, "title-options-sentinel", raw["user"]) } func TestResolveManualTitleModel_TitleGenerationOverrideSetUnusable(t *testing.T) { @@ -636,7 +644,7 @@ func TestResolveManualTitleModel_TitleGenerationOverrideSetUnusable(t *testing.T db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil) server := titleOverrideTestServer(db, logger) - model, gotConfig, err := server.resolveManualTitleModel( + resolved, err := server.resolveManualTitleModel( ctx, db, chat, @@ -645,8 +653,7 @@ func TestResolveManualTitleModel_TitleGenerationOverrideSetUnusable(t *testing.T require.Error(t, err) require.ErrorContains(t, err, "resolve manual title generation model override") require.ErrorContains(t, err, "title generation model override is unavailable") - require.False(t, model.Valid()) - require.Equal(t, database.ChatModelConfig{}, gotConfig) + require.Equal(t, resolvedModelCall{}, resolved) } func TestParseModelOverride(t *testing.T) {