diff --git a/coderd/x/chatd/ARCHITECTURE.md b/coderd/x/chatd/ARCHITECTURE.md index a71dbe1e62..bec77910ec 100644 --- a/coderd/x/chatd/ARCHITECTURE.md +++ b/coderd/x/chatd/ARCHITECTURE.md @@ -866,6 +866,8 @@ The transport is decided in more than one place, and those decisions must agree Paths that build their own clients must thread the override too, including the compaction override, quick generation (used by turn status labels and debug models), and the advisor runtime. +Client construction returns a `chatprovider.Model`, which pairs the fantasy client with the transport resolved from that client's own identity as a `chatopenai.Transport`. Its fields are unexported and only the constructor sets the transport, so no caller can pair a client with a transport it does not speak; a nil client yields the invalid zero value, which fails closed. Decorators such as debug recording replace the wrapped client through `Model.WithLanguageModel`, which preserves the resolved transport, because wrapping does not change what the client speaks. Request preparation does not read the carried transport yet; it still recomputes the decision from the override, and the wrapper is the authoritative value those recomputations must agree with. + Azure is deliberately exempt: its provider always enables the Responses API for known models and exposes no equivalent per-model hook, so `UsesResponsesAPI` 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. Both transports read the same `provider_options.openai` config, but not every field applies to both wire formats. The table below records, per field, which transport honors it; `TestProviderOptionsTransportParity` fails when a field is honored on one transport and silently ignored on the other without being recorded there as intentional. diff --git a/coderd/x/chatd/advisor_internal_test.go b/coderd/x/chatd/advisor_internal_test.go index befb74f84a..8c1e979eb8 100644 --- a/coderd/x/chatd/advisor_internal_test.go +++ b/coderd/x/chatd/advisor_internal_test.go @@ -19,6 +19,7 @@ import ( "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/util/ptr" "github.com/coder/coder/v2/coderd/x/chatd/chatadvisor" + "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" "github.com/coder/coder/v2/coderd/x/chatd/chattest" "github.com/coder/coder/v2/codersdk" "github.com/coder/coder/v2/testutil" @@ -107,11 +108,11 @@ func (p *Server) resolveAdvisorModelOverrideOrFallback( ctx context.Context, chat database.Chat, advisorCfg codersdk.AdvisorConfig, - fallbackModel fantasy.LanguageModel, + fallbackModel chatprovider.Model, fallbackCallConfig codersdk.ChatModelCallConfig, modelOpts modelBuildOptions, logger slog.Logger, -) (fantasy.LanguageModel, codersdk.ChatModelCallConfig) { +) (chatprovider.Model, codersdk.ChatModelCallConfig) { model, cfg, err := p.resolveAdvisorModelOverride( ctx, chat, @@ -132,7 +133,7 @@ func (p *Server) newAdvisorRuntimeOrFallback( ctx context.Context, chat database.Chat, advisorCfg codersdk.AdvisorConfig, - fallbackModel fantasy.LanguageModel, + fallbackModel chatprovider.Model, fallbackCallConfig codersdk.ChatModelCallConfig, modelOpts modelBuildOptions, logger slog.Logger, @@ -159,7 +160,7 @@ func (p *Server) newAdvisorRuntimeOrFallback( func TestResolveAdvisorModelOverride(t *testing.T) { t.Parallel() - fallbackModel := &chattest.FakeModel{ProviderName: "stub", ModelName: "stub"} + fallbackModel := chatprovider.NewModel(&chattest.FakeModel{ProviderName: "stub", ModelName: "stub"}, nil) fallbackCallConfig := codersdk.ChatModelCallConfig{} logger := slog.Make() @@ -356,14 +357,14 @@ func TestResolveAdvisorModelOverride(t *testing.T) { modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, logger, ) - require.NotEqual(t, fantasy.LanguageModel(fallbackModel), gotModel, + require.NotEqual(t, fallbackModel.LanguageModel(), gotModel.LanguageModel(), "success path must return the override model, not the fallback") - require.NotNil(t, gotModel) + require.True(t, gotModel.Valid()) require.Equal(t, "openai", gotModel.Provider()) // Guard against ModelFromConfig silently ignoring the model field // and returning a default. The override is only useful if the // model name from the config row actually propagates. - require.Equal(t, "gpt-5.2", gotModel.Model()) + require.Equal(t, "gpt-5.2", gotModel.ModelID()) require.NotNil(t, gotCfg.Temperature) require.InDelta(t, 0.42, *gotCfg.Temperature, 1e-9) require.NotNil(t, gotCfg.ReasoningEffort) @@ -411,10 +412,10 @@ func TestResolveAdvisorModelOverride(t *testing.T) { modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, logger, ) - require.NotEqual(t, fantasy.LanguageModel(fallbackModel), gotModel) - require.NotNil(t, gotModel) + require.NotEqual(t, fallbackModel.LanguageModel(), gotModel.LanguageModel()) + require.True(t, gotModel.Valid()) require.Equal(t, "openai", gotModel.Provider()) - require.Equal(t, "gpt-5.2", gotModel.Model()) + require.Equal(t, "gpt-5.2", gotModel.ModelID()) require.Equal(t, fallbackCallConfig, gotCfg) }) } @@ -449,13 +450,13 @@ func TestResolveAdvisorModelOverridePromotesAIBridgeErrors(t *testing.T) { ctx, database.Chat{ID: uuid.New(), OwnerID: uuid.New()}, codersdk.AdvisorConfig{ModelConfigID: configID}, - &chattest.FakeModel{ProviderName: "stub", ModelName: "stub"}, + chatprovider.NewModel(&chattest.FakeModel{ProviderName: "stub", ModelName: "stub"}, nil), codersdk.ChatModelCallConfig{}, modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, slog.Make(), ) require.ErrorContains(t, err, "AI Gateway transport factory") - require.Nil(t, model) + require.False(t, model.Valid()) } // TestStripAdvisorGuidanceBlock exercises the filter that keeps the advisor @@ -546,7 +547,7 @@ func TestNewAdvisorRuntime(t *testing.T) { t.Parallel() logger := slog.Make() - fallbackModel := &chattest.FakeModel{ProviderName: "openai", ModelName: "gpt-4"} + fallbackModel := chatprovider.NewModel(&chattest.FakeModel{ProviderName: "openai", ModelName: "gpt-4"}, nil) fallbackCallConfig := codersdk.ChatModelCallConfig{} t.Run("ZeroMaxUsesDefaultsToMaxChatSteps", func(t *testing.T) { diff --git a/coderd/x/chatd/chatd.go b/coderd/x/chatd/chatd.go index 6ea922db74..fb66ea4bdd 100644 --- a/coderd/x/chatd/chatd.go +++ b/coderd/x/chatd/chatd.go @@ -278,11 +278,11 @@ func (p *Server) resolveAdvisorModelOverride( ctx context.Context, chat database.Chat, advisorCfg codersdk.AdvisorConfig, - fallbackModel fantasy.LanguageModel, + fallbackModel chatprovider.Model, fallbackCallConfig codersdk.ChatModelCallConfig, modelOpts modelBuildOptions, logger slog.Logger, -) (fantasy.LanguageModel, codersdk.ChatModelCallConfig, error) { +) (chatprovider.Model, codersdk.ChatModelCallConfig, error) { if advisorCfg.ModelConfigID == uuid.Nil { return fallbackModel, fallbackCallConfig, nil } @@ -331,7 +331,7 @@ func (p *Server) resolveAdvisorModelOverride( ) if err != nil { if overrideConfig.AIProviderID.Valid { - return nil, codersdk.ChatModelCallConfig{}, xerrors.Errorf("resolve advisor override route: %w", err) + return chatprovider.Model{}, codersdk.ChatModelCallConfig{}, xerrors.Errorf("resolve advisor override route: %w", err) } logger.Warn( ctx, @@ -350,7 +350,7 @@ func (p *Server) resolveAdvisorModelOverride( }, route, modelOpts) if err != nil { if overrideConfig.AIProviderID.Valid { - return nil, codersdk.ChatModelCallConfig{}, xerrors.Errorf("create advisor override model: %w", err) + return chatprovider.Model{}, codersdk.ChatModelCallConfig{}, xerrors.Errorf("create advisor override model: %w", err) } logger.Warn( ctx, @@ -381,7 +381,7 @@ func (p *Server) newAdvisorRuntime( ctx context.Context, chat database.Chat, advisorCfg codersdk.AdvisorConfig, - fallbackModel fantasy.LanguageModel, + fallbackModel chatprovider.Model, fallbackCallConfig codersdk.ChatModelCallConfig, modelOpts modelBuildOptions, logger slog.Logger, @@ -430,19 +430,19 @@ func (p *Server) newAdvisorRuntime( ) advisorResponsesOverride := chatprovider.OpenAIResponsesAPIOverride(advisorCallConfig.OpenAIConfig) providerOptions := chatprovider.ProviderOptionsFromChatModelConfig( - advisorModel, + advisorModel.LanguageModel(), advisorCallConfig.ProviderOptions, advisorResponsesOverride, ) providerOptions = chatprovider.ApplyReasoningEffort( - advisorModel, + advisorModel.LanguageModel(), providerOptions, advisorReasoningEffort, advisorResponsesOverride, ) rt, err := chatadvisor.NewRuntime(chatadvisor.RuntimeConfig{ - Model: advisorModel, + Model: advisorModel.LanguageModel(), ModelConfig: advisorCallConfig, ProviderOptions: providerOptions, MaxUsesPerRun: maxUsesPerRun, @@ -2602,7 +2602,7 @@ func (p *Server) generateManualTitleCandidate( titleCtx, messages, pasteText, - titleModel, + titleModel.LanguageModel(), p.titleGenerationProviderOptions(ctx, titleModel, modelConfig), ) finishDebugRun(err) @@ -2654,8 +2654,8 @@ func (p *Server) prepareManualTitleDebugRun( modelConfig database.ChatModelConfig, modelOpts modelBuildOptions, messages []database.ChatMessage, - fallbackModel fantasy.LanguageModel, -) (context.Context, fantasy.LanguageModel, func(error)) { + fallbackModel chatprovider.Model, +) (context.Context, chatprovider.Model, func(error)) { titleCtx := ctx titleModel := fallbackModel finishDebugRun := func(error) {} @@ -2674,7 +2674,7 @@ func (p *Server) prepareManualTitleDebugRun( debugOpts := modelOpts debugOpts.RecordHTTP = true var debugModelErr error - var debugModel fantasy.LanguageModel + var debugModel chatprovider.Model if routeErr != nil { debugModelErr = routeErr } else { @@ -2693,18 +2693,18 @@ func (p *Server) prepareManualTitleDebugRun( slog.F("model", modelConfig.Model), slog.Error(debugModelErr), ) - case debugModel == nil: + 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 = chatdebug.WrapModel(debugModel, debugSvc, chatdebug.RecorderOptions{ + titleModel = debugModel.WithLanguageModel(chatdebug.WrapModel(debugModel.LanguageModel(), debugSvc, chatdebug.RecorderOptions{ ChatID: chat.ID, OwnerID: chat.OwnerID, Provider: routeProvider, Model: modelConfig.Model, - }) + })) } var historyTipMessageID int64 @@ -2829,7 +2829,7 @@ func (p *Server) resolveManualTitleModel( store database.Store, chat database.Chat, modelOpts modelBuildOptions, -) (fantasy.LanguageModel, database.ChatModelConfig, error) { +) (chatprovider.Model, database.ChatModelConfig, error) { overrideConfig, overrideModel, _, overrideSet, overrideErr := p.resolveTitleGenerationModelOverride( ctx, chat, @@ -2837,7 +2837,7 @@ func (p *Server) resolveManualTitleModel( ) if overrideErr != nil { if overrideSet { - return nil, database.ChatModelConfig{}, xerrors.Errorf( + return chatprovider.Model{}, database.ChatModelConfig{}, xerrors.Errorf( "resolve manual title generation model override: %w", overrideErr, ) @@ -2896,17 +2896,17 @@ func (p *Server) resolveFallbackManualTitleModel( ctx context.Context, chat database.Chat, modelOpts modelBuildOptions, -) (fantasy.LanguageModel, database.ChatModelConfig, error) { +) (chatprovider.Model, database.ChatModelConfig, error) { config, err := p.resolveModelConfig(ctx, chat) if err != nil { - return nil, database.ChatModelConfig{}, xerrors.Errorf( + return chatprovider.Model{}, database.ChatModelConfig{}, xerrors.Errorf( "resolve fallback manual title model config: %w", err, ) } route, err := p.resolveModelRouteForConfig(ctx, chat.OwnerID, config) if err != nil { - return nil, database.ChatModelConfig{}, err + return chatprovider.Model{}, database.ChatModelConfig{}, err } model, err := p.newModel(ctx, modelClientRequest{ Chat: chat, @@ -2916,7 +2916,7 @@ func (p *Server) resolveFallbackManualTitleModel( ConfigOptions: config.Options, }, route, modelOpts) if err != nil { - return nil, database.ChatModelConfig{}, xerrors.Errorf( + return chatprovider.Model{}, database.ChatModelConfig{}, xerrors.Errorf( "create fallback manual title model: %w", err, ) @@ -3499,7 +3499,7 @@ func (p *Server) trackWorkspaceUsage( type runChatResult struct { FinalAssistantText string - StatusLabelModel fantasy.LanguageModel + StatusLabelModel chatprovider.Model FallbackProvider string FallbackRoute aiGatewayModelRoute FallbackModel string @@ -4085,7 +4085,7 @@ func (p *Server) resolveChatModel( chat database.Chat, modelOpts modelBuildOptions, ) ( - model fantasy.LanguageModel, + model chatprovider.Model, dbConfig database.ChatModelConfig, route aiGatewayModelRoute, debugEnabled bool, @@ -4095,16 +4095,16 @@ func (p *Server) resolveChatModel( ) { dbConfig, err = p.resolveModelConfig(ctx, chat) if err != nil { - return nil, database.ChatModelConfig{}, aiGatewayModelRoute{}, false, "", "", xerrors.Errorf("resolve model config: %w", err) + return chatprovider.Model{}, database.ChatModelConfig{}, aiGatewayModelRoute{}, false, "", "", xerrors.Errorf("resolve model config: %w", err) } if !dbConfig.Enabled { - return nil, database.ChatModelConfig{}, aiGatewayModelRoute{}, false, "", "", xerrors.Errorf("chat model config %s is disabled", dbConfig.ID) + 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 nil, database.ChatModelConfig{}, aiGatewayModelRoute{}, false, "", "", err + return chatprovider.Model{}, database.ChatModelConfig{}, aiGatewayModelRoute{}, false, "", "", err } providerHint := route.ModelProviderHint @@ -4113,7 +4113,7 @@ func (p *Server) resolveChatModel( providerHint, ) if err != nil { - return nil, database.ChatModelConfig{}, aiGatewayModelRoute{}, false, "", "", xerrors.Errorf( + return chatprovider.Model{}, database.ChatModelConfig{}, aiGatewayModelRoute{}, false, "", "", xerrors.Errorf( "resolve model metadata: %w", err, ) } @@ -4126,7 +4126,7 @@ func (p *Server) resolveChatModel( ConfigOptions: dbConfig.Options, }, route, modelOpts) if err != nil { - return nil, database.ChatModelConfig{}, aiGatewayModelRoute{}, false, "", "", xerrors.Errorf( + return chatprovider.Model{}, database.ChatModelConfig{}, aiGatewayModelRoute{}, false, "", "", xerrors.Errorf( "create model: %w", err, ) } @@ -4680,7 +4680,7 @@ func (p *Server) generateFinalTurnStatusLabel( } assistantText := strings.TrimSpace(runResult.FinalAssistantText) - if assistantText == "" || runResult.StatusLabelModel == nil { + if assistantText == "" || !runResult.StatusLabelModel.Valid() { return fallbackTurnStatusLabel(status) } @@ -4946,7 +4946,7 @@ func (p *Server) resolveChatSummaryModel( slog.F("chat_id", chat.ID), slog.Error(err)) return nil, database.ChatModelConfig{}, false } - return model, dbConfig, true + return model.LanguageModel(), dbConfig, 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 9b8e33de3c..8fdf19c6a2 100644 --- a/coderd/x/chatd/chatd_debug.go +++ b/coderd/x/chatd/chatd_debug.go @@ -4,8 +4,6 @@ import ( "context" "time" - "charm.land/fantasy" - "cdr.dev/slog/v3" "github.com/coder/coder/v2/coderd/x/chatd/chatdebug" "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" @@ -118,10 +116,10 @@ func (p *Server) newDebugAwareModel( req modelClientRequest, route aiGatewayModelRoute, opts modelBuildOptions, -) (fantasy.LanguageModel, bool, error) { +) (chatprovider.Model, bool, error) { provider, resolvedModel, err := chatprovider.ResolveModelWithProviderHint(req.ModelName, route.ModelProviderHint) if err != nil { - return nil, false, err + return chatprovider.Model{}, false, err } route.ModelProviderHint = provider req.ModelName = resolvedModel @@ -132,16 +130,16 @@ func (p *Server) newDebugAwareModel( model, err := p.newModel(ctx, req, route, opts) if err != nil { - return nil, debugEnabled, err + return chatprovider.Model{}, debugEnabled, err } if !debugEnabled { return model, false, nil } - return chatdebug.WrapModel(model, debugSvc, chatdebug.RecorderOptions{ + return model.WithLanguageModel(chatdebug.WrapModel(model.LanguageModel(), debugSvc, chatdebug.RecorderOptions{ ChatID: req.Chat.ID, OwnerID: req.Chat.OwnerID, Provider: provider, Model: resolvedModel, - }), true, nil + })), true, nil } diff --git a/coderd/x/chatd/chatd_internal_test.go b/coderd/x/chatd/chatd_internal_test.go index 0ed9206878..4d16d0b5f3 100644 --- a/coderd/x/chatd/chatd_internal_test.go +++ b/coderd/x/chatd/chatd_internal_test.go @@ -3709,7 +3709,7 @@ func TestPrepareManualTitleDebugRun_RouteFailureDerivesProviderFromConfig(t *tes allowBYOK: true, } debugSvc := chatdebug.NewService(db, logger, nil) - fallbackModel := &chattest.FakeModel{ProviderName: "stub", ModelName: "stub"} + fallbackModel := chatprovider.NewModel(&chattest.FakeModel{ProviderName: "stub", ModelName: "stub"}, nil) server.prepareManualTitleDebugRun( ctx, diff --git a/coderd/x/chatd/chatopenai/options.go b/coderd/x/chatd/chatopenai/options.go index d8851ce187..6f4c53c290 100644 --- a/coderd/x/chatd/chatopenai/options.go +++ b/coderd/x/chatd/chatopenai/options.go @@ -5,7 +5,6 @@ import ( "strings" "charm.land/fantasy" - fantasyazure "charm.land/fantasy/providers/azure" fantasyopenai "charm.land/fantasy/providers/openai" "github.com/coder/coder/v2/coderd/x/chatd/chatutil" @@ -120,17 +119,7 @@ func EnsureResponseIncludes( // UsesResponsesAPI reports whether a model uses the OpenAI Responses API. // Callers must pass the same override the client was built with. func UsesResponsesAPI(provider, modelID string, override *bool) bool { - switch provider { - case fantasyopenai.Name: - if override != nil { - return *override - } - return fantasyopenai.IsResponsesModel(modelID) - case fantasyazure.Name: - return fantasyopenai.IsResponsesModel(modelID) - default: - return false - } + return TransportFor(provider, modelID, override).UsesResponses() } // UsesResponsesOptions reports whether the model should use OpenAI Responses diff --git a/coderd/x/chatd/chatopenai/transport.go b/coderd/x/chatd/chatopenai/transport.go new file mode 100644 index 0000000000..57ee4eff6c --- /dev/null +++ b/coderd/x/chatd/chatopenai/transport.go @@ -0,0 +1,64 @@ +package chatopenai + +import ( + "fmt" + + fantasyazure "charm.land/fantasy/providers/azure" + fantasyopenai "charm.land/fantasy/providers/openai" + fantasyopenaicompat "charm.land/fantasy/providers/openaicompat" +) + +// Transport identifies the OpenAI wire format a model's client speaks. It is +// resolved once when the client is built so that request preparation cannot +// disagree with the client it prepares for. +type Transport int + +const ( + // TransportInvalid is the zero value, meaning no transport was ever + // resolved. + TransportInvalid Transport = iota + // TransportNotApplicable marks providers outside the OpenAI wire formats. + TransportNotApplicable + TransportChatCompletions + TransportResponses +) + +// UsesResponses reports whether t selects the Responses API. It panics on +// TransportInvalid instead of defaulting, because an unresolved transport is a +// construction bug rather than user input. +func (t Transport) UsesResponses() bool { + switch t { + case TransportResponses: + return true + case TransportChatCompletions, TransportNotApplicable: + return false + default: + panic(fmt.Sprintf("chatopenai: unresolved transport %d", int(t))) + } +} + +// TransportFor resolves the wire format for a provider and model. override, +// when non-nil, forces the choice instead of consulting the provider SDK's +// known-model list. Azure follows the known-model list because its provider +// exposes no equivalent hook. +func TransportFor(provider, modelID string, override *bool) Transport { + var useResponses bool + switch provider { + case fantasyopenai.Name: + useResponses = fantasyopenai.IsResponsesModel(modelID) + if override != nil { + useResponses = *override + } + case fantasyazure.Name: + useResponses = fantasyopenai.IsResponsesModel(modelID) + case fantasyopenaicompat.Name: + // chatd never builds an openai-compat client with Responses enabled. + return TransportChatCompletions + default: + return TransportNotApplicable + } + if useResponses { + return TransportResponses + } + return TransportChatCompletions +} diff --git a/coderd/x/chatd/chatopenai/transport_test.go b/coderd/x/chatd/chatopenai/transport_test.go new file mode 100644 index 0000000000..28f23675f9 --- /dev/null +++ b/coderd/x/chatd/chatopenai/transport_test.go @@ -0,0 +1,60 @@ +package chatopenai_test + +import ( + "testing" + + fantasyazure "charm.land/fantasy/providers/azure" + fantasyopenai "charm.land/fantasy/providers/openai" + fantasyopenaicompat "charm.land/fantasy/providers/openaicompat" + "github.com/stretchr/testify/require" + + "github.com/coder/coder/v2/coderd/x/chatd/chatopenai" +) + +func TestTransportFor(t *testing.T) { + t.Parallel() + + // Taken from opposite sides of the provider SDK's known-model list. + const responsesModel = "gpt-4o" + const nonResponsesModel = "babbage-002" + + forceResponses := true + forceCompletions := false + + tests := []struct { + name string + provider string + modelID string + override *bool + want chatopenai.Transport + }{ + {"OpenAIKnownModel", fantasyopenai.Name, responsesModel, nil, chatopenai.TransportResponses}, + {"OpenAIUnknownModel", fantasyopenai.Name, nonResponsesModel, nil, chatopenai.TransportChatCompletions}, + {"OpenAIForceResponses", fantasyopenai.Name, nonResponsesModel, &forceResponses, chatopenai.TransportResponses}, + {"OpenAIForceCompletions", fantasyopenai.Name, responsesModel, &forceCompletions, chatopenai.TransportChatCompletions}, + {"AzureIgnoresOverride", fantasyazure.Name, responsesModel, &forceCompletions, chatopenai.TransportResponses}, + {"AzureKnownModelList", fantasyazure.Name, nonResponsesModel, nil, chatopenai.TransportChatCompletions}, + {"OpenAICompat", fantasyopenaicompat.Name, responsesModel, &forceResponses, chatopenai.TransportChatCompletions}, + {"Anthropic", "anthropic", "claude-sonnet-4-5", &forceResponses, chatopenai.TransportNotApplicable}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + require.Equal(t, tt.want, chatopenai.TransportFor(tt.provider, tt.modelID, tt.override)) + }) + } +} + +func TestTransportUsesResponses(t *testing.T) { + t.Parallel() + + require.True(t, chatopenai.TransportResponses.UsesResponses()) + require.False(t, chatopenai.TransportChatCompletions.UsesResponses()) + require.False(t, chatopenai.TransportNotApplicable.UsesResponses()) + + // The zero value must not silently mean Chat Completions. + require.Panics(t, func() { + _ = chatopenai.TransportInvalid.UsesResponses() + }) +} diff --git a/coderd/x/chatd/chatprovider/chatprovider.go b/coderd/x/chatd/chatprovider/chatprovider.go index 4929a1580c..051c0d0f4d 100644 --- a/coderd/x/chatd/chatprovider/chatprovider.go +++ b/coderd/x/chatd/chatprovider/chatprovider.go @@ -920,16 +920,16 @@ func ModelFromConfig( extraHeaders map[string]string, httpClient *http.Client, openAIResponsesOverride *bool, -) (fantasy.LanguageModel, error) { +) (Model, error) { provider, modelID, err := ResolveModelWithProviderHint(modelName, providerHint) if err != nil { - return nil, err + return Model{}, err } apiKey := providerKeys.APIKey(provider) if apiKey == "" && !(ProviderAllowsAmbientCredentials(provider) && providerKeys.HasProvider(provider)) { - return nil, missingProviderAPIKeyError(provider) + return Model{}, missingProviderAPIKeyError(provider) } baseURL := providerKeys.BaseURL(provider) @@ -952,7 +952,7 @@ func ModelFromConfig( providerClient, err = fantasyanthropic.New(options...) case fantasyazure.Name: if baseURL == "" { - return nil, xerrors.New("AZURE_OPENAI_BASE_URL is not set") + return Model{}, xerrors.New("AZURE_OPENAI_BASE_URL is not set") } azureOpts := []fantasyazure.Option{ fantasyazure.WithAPIKey(apiKey), @@ -1068,17 +1068,17 @@ func ModelFromConfig( } providerClient, err = fantasyvercel.New(options...) default: - return nil, xerrors.Errorf("unsupported model provider %q", provider) + return Model{}, xerrors.Errorf("unsupported model provider %q", provider) } if err != nil { - return nil, providerCreationError(provider, err) + return Model{}, providerCreationError(provider, err) } model, err := providerClient.LanguageModel(context.Background(), modelID) if err != nil { - return nil, xerrors.Errorf("load %s model: %w", provider, err) + return Model{}, xerrors.Errorf("load %s model: %w", provider, err) } - return model, nil + return NewModel(model, openAIResponsesOverride), nil } func providerCreationError(provider string, err error) error { diff --git a/coderd/x/chatd/chatprovider/chatprovider_test.go b/coderd/x/chatd/chatprovider/chatprovider_test.go index 08433b2636..6a1188c3da 100644 --- a/coderd/x/chatd/chatprovider/chatprovider_test.go +++ b/coderd/x/chatd/chatprovider/chatprovider_test.go @@ -918,7 +918,7 @@ func TestModelFromConfig_Bedrock(t *testing.T) { nil, ) require.NoError(t, err) - require.NotNil(t, model) + require.True(t, model.Valid()) require.Equal(t, fantasybedrock.Name, model.Provider()) }) @@ -934,7 +934,7 @@ func TestModelFromConfig_Bedrock(t *testing.T) { nil, nil, ) - require.Nil(t, model) + require.False(t, model.Valid()) require.EqualError(t, err, "API key for provider \"bedrock\" is not set") }) @@ -978,9 +978,9 @@ func TestModelFromConfig_Bedrock(t *testing.T) { nil, ) require.NoError(t, err) - require.NotNil(t, model) + require.True(t, model.Valid()) - _, err = model.Generate(ctx, fantasy.Call{ + _, err = model.LanguageModel().Generate(ctx, fantasy.Call{ Prompt: []fantasy.Message{ { Role: fantasy.MessageRoleUser, @@ -1034,7 +1034,7 @@ func TestModelFromConfig_Bedrock(t *testing.T) { nil, nil, ) - require.Nil(t, model) + require.False(t, model.Valid()) require.EqualError(t, err, tt.wantErr) }) } @@ -1103,9 +1103,9 @@ func TestModelFromConfig_BedrockStripsAnthropicHeaders(t *testing.T) { nil, ) require.NoError(t, err) - require.NotNil(t, model) + require.True(t, model.Valid()) - _, err = model.Generate(ctx, fantasy.Call{ + _, err = model.LanguageModel().Generate(ctx, fantasy.Call{ Prompt: []fantasy.Message{ { Role: fantasy.MessageRoleUser, @@ -1188,9 +1188,9 @@ func TestModelFromConfig_BedrockStreamingHeaders(t *testing.T) { nil, ) require.NoError(t, err) - require.NotNil(t, model) + require.True(t, model.Valid()) - stream, err := model.Stream(ctx, fantasy.Call{ + stream, err := model.LanguageModel().Stream(ctx, fantasy.Call{ Prompt: []fantasy.Message{ { Role: fantasy.MessageRoleUser, @@ -1338,7 +1338,7 @@ func TestModelFromConfig_ExtraHeaders(t *testing.T) { model, err := chatprovider.ModelFromConfig("openai", "gpt-4", keys, chatprovider.UserAgent(), headers, nil, nil) require.NoError(t, err) - _, err = model.Generate(ctx, fantasy.Call{ + _, err = model.LanguageModel().Generate(ctx, fantasy.Call{ Prompt: []fantasy.Message{ { Role: fantasy.MessageRoleUser, @@ -1369,7 +1369,7 @@ func TestModelFromConfig_ExtraHeaders(t *testing.T) { model, err := chatprovider.ModelFromConfig("anthropic", "claude-sonnet-4-20250514", keys, chatprovider.UserAgent(), headers, nil, nil) require.NoError(t, err) - _, err = model.Generate(ctx, fantasy.Call{ + _, err = model.LanguageModel().Generate(ctx, fantasy.Call{ Prompt: []fantasy.Message{ { Role: fantasy.MessageRoleUser, @@ -1422,8 +1422,8 @@ func TestBetaHeadersFromCallConfig(t *testing.T) { } } -func generateHello(ctx context.Context, model fantasy.LanguageModel) error { - _, err := model.Generate(ctx, fantasy.Call{ +func generateHello(ctx context.Context, model chatprovider.Model) error { + _, err := model.LanguageModel().Generate(ctx, fantasy.Call{ Prompt: []fantasy.Message{ { Role: fantasy.MessageRoleUser, @@ -1584,7 +1584,7 @@ func TestModelFromConfig_AnthropicPDFFilePartReachesProvider(t *testing.T) { model, err := chatprovider.ModelFromConfig("anthropic", "claude-sonnet-4-20250514", keys, chatprovider.UserAgent(), nil, nil, nil) require.NoError(t, err) - _, err = model.Generate(ctx, fantasy.Call{ + _, err = model.LanguageModel().Generate(ctx, fantasy.Call{ Prompt: []fantasy.Message{ { Role: fantasy.MessageRoleUser, @@ -1625,7 +1625,7 @@ func TestModelFromConfig_NilExtraHeaders(t *testing.T) { model, err := chatprovider.ModelFromConfig("openai", "gpt-4", keys, chatprovider.UserAgent(), nil, nil, nil) require.NoError(t, err) - _, err = model.Generate(ctx, fantasy.Call{ + _, err = model.LanguageModel().Generate(ctx, fantasy.Call{ Prompt: []fantasy.Message{ { Role: fantasy.MessageRoleUser, @@ -1670,7 +1670,7 @@ func TestModelFromConfig_HTTPClient(t *testing.T) { ) require.NoError(t, err) - _, err = model.Generate(ctx, fantasy.Call{ + _, err = model.LanguageModel().Generate(ctx, fantasy.Call{ Prompt: []fantasy.Message{{ Role: fantasy.MessageRoleUser, Content: []fantasy.MessagePart{fantasy.TextPart{Text: "hello"}}, diff --git a/coderd/x/chatd/chatprovider/model.go b/coderd/x/chatd/chatprovider/model.go new file mode 100644 index 0000000000..09414202ac --- /dev/null +++ b/coderd/x/chatd/chatprovider/model.go @@ -0,0 +1,59 @@ +package chatprovider + +import ( + "charm.land/fantasy" + + "github.com/coder/coder/v2/coderd/x/chatd/chatopenai" +) + +// Model pairs a language model client with the facts resolved when it was +// built. Its fields are unexported and nothing sets transport except NewModel, +// which derives it from the client, so no caller can choose a transport that +// disagrees with the client it wraps. +type Model struct { + lm fantasy.LanguageModel + transport chatopenai.Transport +} + +// NewModel pairs an already-built client with the transport resolved from that +// client's own identity. ModelFromConfig is the only production caller; +// callers must pass the same override the client was built with, which is the +// one degree of freedom Model cannot police. +func NewModel(lm fantasy.LanguageModel, openAIResponsesOverride *bool) Model { + if lm == nil { + // The invalid zero value lets callers report a nil client as an + // error instead of dereferencing it here. + return Model{} + } + return Model{ + lm: lm, + transport: chatopenai.TransportFor(lm.Provider(), lm.Model(), openAIResponsesOverride), + } +} + +func (m Model) LanguageModel() fantasy.LanguageModel { return m.lm } + +func (m Model) Provider() string { return m.lm.Provider() } + +func (m Model) ModelID() string { return m.lm.Model() } + +func (m Model) Transport() chatopenai.Transport { return m.transport } + +// Valid reports whether m wraps a client rather than being the zero value. +func (m Model) Valid() bool { return m.lm != nil } + +// WithLanguageModel replaces the wrapped client and keeps the resolved +// transport, because decorators such as debug recording do not change what the +// client speaks. It panics on an invalid receiver or a replacement that +// reports a different identity, because both would pair the kept transport +// with a client it was not resolved from. +func (m Model) WithLanguageModel(lm fantasy.LanguageModel) Model { + if !m.Valid() { + panic("chatprovider: WithLanguageModel on an invalid Model") + } + if lm == nil || lm.Provider() != m.lm.Provider() || lm.Model() != m.lm.Model() { + panic("chatprovider: WithLanguageModel replacement changes the model identity") + } + m.lm = lm + return m +} diff --git a/coderd/x/chatd/chatprovider/model_test.go b/coderd/x/chatd/chatprovider/model_test.go new file mode 100644 index 0000000000..e5015ddf8a --- /dev/null +++ b/coderd/x/chatd/chatprovider/model_test.go @@ -0,0 +1,87 @@ +package chatprovider_test + +import ( + "testing" + + fantasyopenai "charm.land/fantasy/providers/openai" + "github.com/stretchr/testify/require" + + "github.com/coder/coder/v2/coderd/x/chatd/chatopenai" + "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" + "github.com/coder/coder/v2/coderd/x/chatd/chattest" +) + +func TestModelResolvesTransportFromClient(t *testing.T) { + t.Parallel() + + forceResponses := true + + model := chatprovider.NewModel( + &chattest.FakeModel{ProviderName: fantasyopenai.Name, ModelName: "babbage-002"}, + &forceResponses, + ) + require.Equal(t, chatopenai.TransportResponses, model.Transport()) + require.Equal(t, fantasyopenai.Name, model.Provider()) + require.Equal(t, "babbage-002", model.ModelID()) + require.True(t, model.Valid()) +} + +func TestModelZeroValueFailsClosed(t *testing.T) { + t.Parallel() + + var model chatprovider.Model + require.False(t, model.Valid()) + require.Equal(t, chatopenai.TransportInvalid, model.Transport()) + require.Panics(t, func() { + _ = model.Transport().UsesResponses() + }) +} + +// A fantasy provider can return a nil client without an error; the +// constructor must yield the invalid zero value so newLanguageModel reports +// it instead of panicking on the nil dereference. +func TestNewModelNilClientIsInvalid(t *testing.T) { + t.Parallel() + + model := chatprovider.NewModel(nil, nil) + require.False(t, model.Valid()) + require.Equal(t, chatopenai.TransportInvalid, model.Transport()) +} + +func TestModelWithLanguageModelPreservesTransport(t *testing.T) { + t.Parallel() + + forceResponses := true + model := chatprovider.NewModel( + &chattest.FakeModel{ProviderName: fantasyopenai.Name, ModelName: "babbage-002"}, + &forceResponses, + ) + + replacement := &chattest.FakeModel{ProviderName: fantasyopenai.Name, ModelName: "babbage-002"} + wrapped := model.WithLanguageModel(replacement) + + require.Equal(t, model.Transport(), wrapped.Transport()) + require.Same(t, replacement, wrapped.LanguageModel()) +} + +// The kept transport is only correct for the identity it was resolved from, +// so replacing the client may not launder an invalid wrapper into a valid one +// or swap in a client with a different identity. +func TestModelWithLanguageModelRejectsIdentityChanges(t *testing.T) { + t.Parallel() + + client := &chattest.FakeModel{ProviderName: fantasyopenai.Name, ModelName: "babbage-002"} + model := chatprovider.NewModel(client, nil) + + require.Panics(t, func() { + chatprovider.Model{}.WithLanguageModel(client) + }) + require.Panics(t, func() { + model.WithLanguageModel(nil) + }) + require.Panics(t, func() { + model.WithLanguageModel( + &chattest.FakeModel{ProviderName: fantasyopenai.Name, ModelName: "gpt-4.1"}, + ) + }) +} diff --git a/coderd/x/chatd/chatprovider/openai_compat_patches_test.go b/coderd/x/chatd/chatprovider/openai_compat_patches_test.go index 4b864c2924..e8d5194ef8 100644 --- a/coderd/x/chatd/chatprovider/openai_compat_patches_test.go +++ b/coderd/x/chatd/chatprovider/openai_compat_patches_test.go @@ -71,7 +71,7 @@ func generateOpenAICompatRequest(t *testing.T, baseURL string, modelID string) m ) require.NoError(t, err) - _, err = model.Generate(t.Context(), fantasy.Call{ + _, err = model.LanguageModel().Generate(t.Context(), fantasy.Call{ Prompt: geminiOpenAICompatToolPrompt(), }) require.NoError(t, err) diff --git a/coderd/x/chatd/chatprovider/responses_api_test.go b/coderd/x/chatd/chatprovider/responses_api_test.go index 30572568fc..87991c8d36 100644 --- a/coderd/x/chatd/chatprovider/responses_api_test.go +++ b/coderd/x/chatd/chatprovider/responses_api_test.go @@ -65,7 +65,7 @@ func TestModelFromConfig_OpenAIResponsesAPIOverride(t *testing.T) { ) require.NoError(t, err) - _, err = model.Generate(context.Background(), fantasy.Call{ + _, err = model.LanguageModel().Generate(context.Background(), fantasy.Call{ Prompt: []fantasy.Message{{ Role: fantasy.MessageRoleUser, Content: []fantasy.MessagePart{fantasy.TextPart{Text: "Test message"}}, diff --git a/coderd/x/chatd/chatprovider/useragent_test.go b/coderd/x/chatd/chatprovider/useragent_test.go index fbe48f68f7..3954295249 100644 --- a/coderd/x/chatd/chatprovider/useragent_test.go +++ b/coderd/x/chatd/chatprovider/useragent_test.go @@ -53,7 +53,7 @@ func TestModelFromConfig_UserAgent(t *testing.T) { // Make a real call so Fantasy sends an HTTP request to the // fake server, which asserts the User-Agent header. - _, err = model.Generate(ctx, fantasy.Call{ + _, err = model.LanguageModel().Generate(ctx, fantasy.Call{ Prompt: []fantasy.Message{ { Role: fantasy.MessageRoleUser, diff --git a/coderd/x/chatd/compaction_override.go b/coderd/x/chatd/compaction_override.go index f784f074f6..8d224073fb 100644 --- a/coderd/x/chatd/compaction_override.go +++ b/coderd/x/chatd/compaction_override.go @@ -36,7 +36,7 @@ func readCompactionModelOverride( // the identity metadata debug runs and prompt sanitization need. type compactionModelOverride struct { modelConfig database.ChatModelConfig - model fantasy.LanguageModel + model chatprovider.Model resolvedProvider string resolvedModel string // providerOptions include the override's reasoning effort for the @@ -169,7 +169,7 @@ func (p *Server) buildCompactionOverrideModel( // options, including the admin-resolved reasoning effort, into provider // options for the summary call. func compactionOverrideProviderOptions( - model fantasy.LanguageModel, + model chatprovider.Model, modelConfig database.ChatModelConfig, ) (fantasy.ProviderOptions, *bool, error) { callConfig := codersdk.ChatModelCallConfig{} @@ -183,7 +183,7 @@ func compactionOverrideProviderOptions( } responsesOverride := chatprovider.OpenAIResponsesAPIOverride(callConfig.OpenAIConfig) providerOptions := chatprovider.ProviderOptionsFromChatModelConfig( - model, + model.LanguageModel(), callConfig.ProviderOptions, responsesOverride, ) @@ -192,7 +192,7 @@ func compactionOverrideProviderOptions( callConfig.ReasoningEffort, ) return chatprovider.ApplyReasoningEffort( - model, + model.LanguageModel(), providerOptions, reasoningEffort, responsesOverride, diff --git a/coderd/x/chatd/compaction_override_internal_test.go b/coderd/x/chatd/compaction_override_internal_test.go index d325b91a89..390801d7fb 100644 --- a/coderd/x/chatd/compaction_override_internal_test.go +++ b/coderd/x/chatd/compaction_override_internal_test.go @@ -13,6 +13,7 @@ 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" @@ -21,7 +22,7 @@ import ( func TestCompactionOverrideProviderOptions(t *testing.T) { t.Parallel() - model := &chattest.FakeModel{ProviderName: "anthropic", ModelName: "claude-3-5-haiku"} + model := chatprovider.NewModel(&chattest.FakeModel{ProviderName: "anthropic", ModelName: "claude-3-5-haiku"}, nil) t.Run("NoOptions", func(t *testing.T) { t.Parallel() diff --git a/coderd/x/chatd/compaction_sanitize.go b/coderd/x/chatd/compaction_sanitize.go index 1f0b939cab..43b0f0f5bf 100644 --- a/coderd/x/chatd/compaction_sanitize.go +++ b/coderd/x/chatd/compaction_sanitize.go @@ -27,7 +27,7 @@ func sanitizeCompactionPrompt( ctx context.Context, logger slog.Logger, prompt []fantasy.Message, - compactionModel fantasy.LanguageModel, + compactionModel chatprovider.Model, chatConfig database.ChatModelConfig, overrideConfig database.ChatModelConfig, openAIResponsesOverride *bool, @@ -39,7 +39,7 @@ func sanitizeCompactionPrompt( messages = replaceUnsupportedFileParts(ctx, logger, messages, func(mediaType string) bool { return chatprovider.AcceptsFilePartMediaType( compactionModel.Provider(), - compactionModel.Model(), + compactionModel.ModelID(), mediaType, openAIResponsesOverride, ) @@ -53,7 +53,7 @@ func sanitizeCompactionPrompt( logger, "compaction_prompt", compactionModel.Provider(), - compactionModel.Model(), + compactionModel.ModelID(), stats, ) return sanitized diff --git a/coderd/x/chatd/compaction_sanitize_internal_test.go b/coderd/x/chatd/compaction_sanitize_internal_test.go index 7e18498bae..f39e650fef 100644 --- a/coderd/x/chatd/compaction_sanitize_internal_test.go +++ b/coderd/x/chatd/compaction_sanitize_internal_test.go @@ -9,6 +9,7 @@ import ( "cdr.dev/slog/v3/sloggers/slogtest" "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" "github.com/coder/coder/v2/coderd/x/chatd/chattest" "github.com/coder/coder/v2/testutil" ) @@ -70,7 +71,7 @@ func TestSanitizeCompactionPrompt_FlattensForeignProviderExecutedToolParts(t *te }, } - compactionModel := &chattest.FakeModel{ProviderName: "openai", ModelName: "gpt-4.1-mini"} + compactionModel := chatprovider.NewModel(&chattest.FakeModel{ProviderName: "openai", ModelName: "gpt-4.1-mini"}, nil) sanitized := sanitizeCompactionPrompt(ctx, logger, prompt, compactionModel, configWithProvider(uuid.New()), configWithProvider(uuid.New()), nil) require.Len(t, sanitized, 3) @@ -119,7 +120,7 @@ func TestSanitizeCompactionPrompt_DropsNonAssistantProviderExecutedParts(t *test }, } - compactionModel := &chattest.FakeModel{ProviderName: "openai", ModelName: "gpt-4.1-mini"} + compactionModel := chatprovider.NewModel(&chattest.FakeModel{ProviderName: "openai", ModelName: "gpt-4.1-mini"}, nil) sanitized := sanitizeCompactionPrompt(ctx, logger, prompt, compactionModel, configWithProvider(uuid.New()), configWithProvider(uuid.New()), nil) require.Len(t, sanitized, 1) @@ -148,7 +149,7 @@ func TestSanitizeCompactionPrompt_ReplacesUnsupportedFileParts(t *testing.T) { // Mistral accepts images but not PDFs, so the PDF part must become a // placeholder while the prompt stays otherwise intact. - compactionModel := &chattest.FakeModel{ProviderName: "mistral", ModelName: "mistral-large"} + compactionModel := chatprovider.NewModel(&chattest.FakeModel{ProviderName: "mistral", ModelName: "mistral-large"}, nil) sharedProviderID := uuid.New() sanitized := sanitizeCompactionPrompt(ctx, logger, prompt, compactionModel, configWithProvider(sharedProviderID), configWithProvider(sharedProviderID), nil) @@ -188,7 +189,7 @@ func TestSanitizeCompactionPrompt_SameProviderKeepsProviderExecutedParts(t *test }, } - compactionModel := &chattest.FakeModel{ProviderName: "openai", ModelName: "gpt-4.1-mini"} + compactionModel := chatprovider.NewModel(&chattest.FakeModel{ProviderName: "openai", ModelName: "gpt-4.1-mini"}, nil) sharedProviderID := uuid.New() sanitized := sanitizeCompactionPrompt(ctx, logger, prompt, compactionModel, configWithProvider(sharedProviderID), configWithProvider(sharedProviderID), nil) diff --git a/coderd/x/chatd/computer_use.go b/coderd/x/chatd/computer_use.go index 93e6c5479f..249f4b3ae3 100644 --- a/coderd/x/chatd/computer_use.go +++ b/coderd/x/chatd/computer_use.go @@ -4,7 +4,6 @@ import ( "context" "strings" - "charm.land/fantasy" "golang.org/x/xerrors" "cdr.dev/slog/v3" @@ -66,7 +65,7 @@ func (p *Server) resolveComputerUseModel( computerUseModelName string, modelOpts modelBuildOptions, ) ( - model fantasy.LanguageModel, + model chatprovider.Model, debugEnabled bool, resolvedProvider string, resolvedModel string, @@ -77,7 +76,7 @@ func (p *Server) resolveComputerUseModel( computerUseModelProvider, ) if err != nil { - return nil, false, "", "", xerrors.Errorf( + return chatprovider.Model{}, false, "", "", xerrors.Errorf( "resolve computer use model metadata for provider %q model %q: %w", computerUseProvider, computerUseModelName, @@ -92,7 +91,7 @@ func (p *Server) resolveComputerUseModel( ExtraHeaders: chatprovider.CoderHeaders(chat), }, route, modelOpts) if err != nil { - return nil, false, "", "", xerrors.Errorf( + return chatprovider.Model{}, false, "", "", xerrors.Errorf( "resolve computer use model for provider %q model %q: %w", computerUseProvider, computerUseModelName, diff --git a/coderd/x/chatd/generation.go b/coderd/x/chatd/generation.go index 4cdc6e1a0e..1f47b1decf 100644 --- a/coderd/x/chatd/generation.go +++ b/coderd/x/chatd/generation.go @@ -20,6 +20,7 @@ import ( "github.com/coder/coder/v2/coderd/x/chatd/chathooks" "github.com/coder/coder/v2/coderd/x/chatd/chatloop" "github.com/coder/coder/v2/coderd/x/chatd/chatprompt" + "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" "github.com/coder/coder/v2/coderd/x/chatd/chatretry" "github.com/coder/coder/v2/coderd/x/chatd/chatstate" "github.com/coder/coder/v2/coderd/x/chatd/messagepartbuffer" @@ -39,7 +40,7 @@ type generationPrepared struct { Chat database.Chat Messages []database.ChatMessage - Model fantasy.LanguageModel + Model chatprovider.Model Prompt []fantasy.Message Tools []fantasy.AgentTool ActiveTools []string @@ -723,7 +724,7 @@ func (s *taskStarter) generateAssistant( defer attempt.closeEpisode() runCtx := input.DebugTurn.Ensure(ctx, prepared.Chat, prepared.Debug) outcome, err := chatloop.GenerateAssistant(runCtx, chatloop.GenerateAssistantOptions{ - Model: prepared.Model, + Model: prepared.Model.LanguageModel(), ErrorProvider: prepared.ResolvedProvider, Messages: prepared.Prompt, Tools: prepared.Tools, @@ -816,9 +817,9 @@ func (s *taskStarter) executeLocalTools( defer attempt.closeEpisode() provider := "" modelName := "" - if prepared.Model != nil { + if prepared.Model.Valid() { provider = prepared.Model.Provider() - modelName = prepared.Model.Model() + modelName = prepared.Model.ModelID() } var outcome chatloop.ToolExecutionOutcome var spawnDispatchErr error @@ -927,7 +928,7 @@ func (s *taskStarter) generateCompaction( slog.F("chat_id", prepared.Chat.ID), slog.F("owner_id", prepared.Chat.OwnerID), ) - compactionOpts.Model = overrideModel.model + compactionOpts.Model = overrideModel.model.LanguageModel() compactionOpts.ResolvedProvider = overrideModel.resolvedProvider compactionOpts.ResolvedModel = overrideModel.resolvedModel compactionOpts.ModelConfigID = overrideModel.modelConfig.ID diff --git a/coderd/x/chatd/generation_preparer.go b/coderd/x/chatd/generation_preparer.go index d3552e83cc..4112384f31 100644 --- a/coderd/x/chatd/generation_preparer.go +++ b/coderd/x/chatd/generation_preparer.go @@ -36,7 +36,7 @@ func (server *Server) prepareGeneration( ) var ( - model fantasy.LanguageModel + model chatprovider.Model modelConfig database.ChatModelConfig modelRoute aiGatewayModelRoute modelOpts modelBuildOptions @@ -300,7 +300,7 @@ func (server *Server) prepareGeneration( acceptsFilePart := func(mediaType string) bool { return chatprovider.AcceptsFilePartMediaType( model.Provider(), - model.Model(), + model.ModelID(), mediaType, chatprovider.OpenAIResponsesAPIOverride(callConfig.OpenAIConfig), ) @@ -371,7 +371,7 @@ func (server *Server) prepareGeneration( logger, "persisted_history_replay", model.Provider(), - model.Model(), + model.ModelID(), sanitizeStats, ) @@ -559,12 +559,12 @@ func (server *Server) prepareGeneration( ) responsesOverride := chatprovider.OpenAIResponsesAPIOverride(callConfig.OpenAIConfig) providerOptions := chatprovider.ProviderOptionsFromChatModelConfig( - model, + model.LanguageModel(), callConfig.ProviderOptions, responsesOverride, ) providerOptions = chatprovider.ApplyReasoningEffort( - model, + model.LanguageModel(), providerOptions, reasoningEffort, responsesOverride, @@ -627,7 +627,7 @@ func (server *Server) prepareGeneration( // The options carry the chat model; generateCompaction swaps in the // override client when one is configured. compactionOptions := chatloop.GenerateCompactionOptions{ - Model: model, + Model: model.LanguageModel(), Messages: prompt, ThresholdPercent: effectiveThreshold, ContextLimit: compactionContextLimit, diff --git a/coderd/x/chatd/generation_preparer_internal_test.go b/coderd/x/chatd/generation_preparer_internal_test.go index 9f3059aea2..f144da0b38 100644 --- a/coderd/x/chatd/generation_preparer_internal_test.go +++ b/coderd/x/chatd/generation_preparer_internal_test.go @@ -441,7 +441,7 @@ 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.NotNil(t, result.StatusLabelModel) + 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)) @@ -518,7 +518,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.Nil(t, result.StatusLabelModel) + require.False(t, result.StatusLabelModel.Valid()) require.Empty(t, result.FallbackProvider) require.Empty(t, result.FallbackModel) }) diff --git a/coderd/x/chatd/model_routing.go b/coderd/x/chatd/model_routing.go index 5a07ace623..d66f9111d1 100644 --- a/coderd/x/chatd/model_routing.go +++ b/coderd/x/chatd/model_routing.go @@ -5,7 +5,6 @@ import ( "encoding/json" "net/http" - "charm.land/fantasy" "github.com/google/uuid" "golang.org/x/xerrors" @@ -47,7 +46,7 @@ func newLanguageModel( extraHeaders map[string]string, httpClient *http.Client, openAIResponsesOverride *bool, -) (fantasy.LanguageModel, error) { +) (chatprovider.Model, error) { model, err := chatprovider.ModelFromConfig( providerHint, modelName, @@ -58,14 +57,14 @@ func newLanguageModel( openAIResponsesOverride, ) if err != nil { - return nil, err + return chatprovider.Model{}, err } - if model == nil { + if !model.Valid() { provider, resolvedModel, resolveErr := chatprovider.ResolveModelWithProviderHint(modelName, providerHint) if resolveErr != nil { - return nil, resolveErr + return chatprovider.Model{}, resolveErr } - return nil, xerrors.Errorf( + return chatprovider.Model{}, xerrors.Errorf( "create model for %s/%s returned nil", provider, resolvedModel, diff --git a/coderd/x/chatd/model_routing_aibridge.go b/coderd/x/chatd/model_routing_aibridge.go index 85b4476ec2..77af154250 100644 --- a/coderd/x/chatd/model_routing_aibridge.go +++ b/coderd/x/chatd/model_routing_aibridge.go @@ -7,7 +7,6 @@ import ( "net/http" "strings" - "charm.land/fantasy" fantasyanthropic "charm.land/fantasy/providers/anthropic" fantasyopenai "charm.land/fantasy/providers/openai" fantasyopenaicompat "charm.land/fantasy/providers/openaicompat" @@ -116,15 +115,15 @@ func (p *Server) newModel( req modelClientRequest, route aiGatewayModelRoute, opts modelBuildOptions, -) (fantasy.LanguageModel, error) { +) (chatprovider.Model, error) { if route.Provider.ID == uuid.Nil { - return nil, xerrors.New("AI Gateway routing requires a concrete AI provider") + return chatprovider.Model{}, xerrors.New("AI Gateway routing requires a concrete AI provider") } if route.Provider.Name == "" { - return nil, xerrors.New("AI Gateway routing requires an AI provider name") + return chatprovider.Model{}, xerrors.New("AI Gateway routing requires an AI provider name") } if opts.ActiveAPIKeyID == "" { - return nil, chaterror.WithClassification( + return chatprovider.Model{}, chaterror.WithClassification( xerrors.New("AI Gateway routing requires the active turn API key ID"), chaterror.ClassifiedError{ Kind: codersdk.ChatErrorKindMissingKey, @@ -135,7 +134,7 @@ func (p *Server) newModel( } if err := ValidateAIGatewayProviderModel(route.Provider, req.ModelName); err != nil { - return nil, chaterror.WithClassification( + return chatprovider.Model{}, chaterror.WithClassification( err, chaterror.ClassifiedError{ Kind: codersdk.ChatErrorKindConfig, @@ -147,15 +146,15 @@ func (p *Server) newModel( factoryPtr := p.aibridgeTransportFactory if factoryPtr == nil { - return nil, xerrors.New("AI Gateway transport factory is not configured") + return chatprovider.Model{}, xerrors.New("AI Gateway transport factory is not configured") } factory := factoryPtr.Load() if factory == nil || *factory == nil { - return nil, xerrors.New("AI Gateway transport factory is not configured") + return chatprovider.Model{}, xerrors.New("AI Gateway transport factory is not configured") } rt, err := (*factory).TransportFor(route.Provider.Name, aibridge.SourceAgents) if err != nil { - return nil, xerrors.Errorf("create AI Gateway transport: %w", err) + return chatprovider.Model{}, xerrors.Errorf("create AI Gateway transport: %w", err) } baseRT := http.RoundTripper(&aiGatewayRoundTripper{ base: rt, @@ -169,7 +168,7 @@ func (p *Server) newModel( config := fantasyConfigForAIBridge(route.Provider.Type) callConfig, err := parseModelConfigOptions(req.ConfigOptions) if err != nil { - return nil, err + return chatprovider.Model{}, err } extraHeaders := mergeConfigBetaHeaders(req.ExtraHeaders, config.ProviderHint, callConfig) return newLanguageModel( diff --git a/coderd/x/chatd/model_routing_internal_test.go b/coderd/x/chatd/model_routing_internal_test.go index 317160c7c9..c85ebe5aab 100644 --- a/coderd/x/chatd/model_routing_internal_test.go +++ b/coderd/x/chatd/model_routing_internal_test.go @@ -313,7 +313,7 @@ func TestAIGatewayModelForwardsProviderAuth(t *testing.T) { apiKeyID := uuid.NewString() model, err := server.newModel(t.Context(), aibridgeTestRequest(database.Chat{ID: uuid.New(), OwnerID: uuid.New()}, "gpt-4"), route, modelBuildOptions{ActiveAPIKeyID: apiKeyID, RecordHTTP: true}) require.NoError(t, err) - _, err = model.Generate(t.Context(), fantasy.Call{Prompt: []fantasy.Message{{Role: fantasy.MessageRoleUser, Content: []fantasy.MessagePart{fantasy.TextPart{Text: "hello"}}}}}) + _, err = model.LanguageModel().Generate(t.Context(), fantasy.Call{Prompt: []fantasy.Message{{Role: fantasy.MessageRoleUser, Content: []fantasy.MessagePart{fantasy.TextPart{Text: "hello"}}}}}) require.NoError(t, err) got := <-seen @@ -335,7 +335,7 @@ func TestAIGatewayModelForwardsProviderAuth(t *testing.T) { apiKeyID := uuid.NewString() model, err := server.newModel(t.Context(), aibridgeTestRequest(database.Chat{ID: uuid.New(), OwnerID: uuid.New()}, "claude-haiku-4-5"), route, modelBuildOptions{ActiveAPIKeyID: apiKeyID}) require.NoError(t, err) - _, err = model.Generate(t.Context(), fantasy.Call{Prompt: []fantasy.Message{{Role: fantasy.MessageRoleUser, Content: []fantasy.MessagePart{fantasy.TextPart{Text: "hello"}}}}}) + _, err = model.LanguageModel().Generate(t.Context(), fantasy.Call{Prompt: []fantasy.Message{{Role: fantasy.MessageRoleUser, Content: []fantasy.MessagePart{fantasy.TextPart{Text: "hello"}}}}}) require.NoError(t, err) got := <-seen @@ -354,7 +354,7 @@ func TestAIGatewayModelForwardsProviderAuth(t *testing.T) { apiKeyID := uuid.NewString() model, err := server.newModel(t.Context(), aibridgeTestRequest(database.Chat{ID: uuid.New(), OwnerID: uuid.New()}, "gpt-4"), route, modelBuildOptions{ActiveAPIKeyID: apiKeyID}) require.NoError(t, err) - _, err = model.Generate(t.Context(), fantasy.Call{Prompt: []fantasy.Message{{Role: fantasy.MessageRoleUser, Content: []fantasy.MessagePart{fantasy.TextPart{Text: "hello"}}}}}) + _, err = model.LanguageModel().Generate(t.Context(), fantasy.Call{Prompt: []fantasy.Message{{Role: fantasy.MessageRoleUser, Content: []fantasy.MessagePart{fantasy.TextPart{Text: "hello"}}}}}) require.NoError(t, err) got := <-seen @@ -426,7 +426,7 @@ func TestAIGatewayModelAppliesResponsesAPIOverride(t *testing.T) { modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, ) require.NoError(t, err) - _, err = model.Generate(t.Context(), fantasy.Call{Prompt: []fantasy.Message{{ + _, err = model.LanguageModel().Generate(t.Context(), fantasy.Call{Prompt: []fantasy.Message{{ Role: fantasy.MessageRoleUser, Content: []fantasy.MessagePart{fantasy.TextPart{Text: "hello"}}, }}}) @@ -607,7 +607,7 @@ func TestAIBridgeGatewayProviderTypesPreserveSlashModelID(t *testing.T) { modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, ) require.NoError(t, err) - _, err = model.Generate(t.Context(), fantasy.Call{Prompt: []fantasy.Message{{ + _, err = model.LanguageModel().Generate(t.Context(), fantasy.Call{Prompt: []fantasy.Message{{ Role: fantasy.MessageRoleUser, Content: []fantasy.MessagePart{fantasy.TextPart{Text: "hello"}}, }}}) @@ -650,7 +650,7 @@ func TestAIBridgeComputerUseModelUsesRoute(t *testing.T) { modelBuildOptions{ActiveAPIKeyID: apiKeyID}, ) require.NoError(t, err) - require.NotNil(t, model) + require.True(t, model.Valid()) require.False(t, debugEnabled) require.EqualValues(t, codersdk.ChatComputerUseProviderOpenAI, resolvedProvider) require.Equal(t, modelName, resolvedModel) @@ -684,7 +684,7 @@ func TestResolveComputerUseModel_AIGatewayMissingAPIKeyID(t *testing.T) { modelBuildOptions{}, // no ActiveAPIKeyID ) require.Error(t, err) - require.Nil(t, model) + require.False(t, model.Valid()) require.False(t, debugEnabled) require.Empty(t, resolvedProvider) require.Empty(t, resolvedModel) @@ -726,7 +726,7 @@ func TestAIBridgeDelegatedContextPropagation(t *testing.T) { ctx := aibridge.WithDelegatedAPIKeyID(t.Context(), "context-key-must-be-ignored") model, err := server.newModel(ctx, aibridgeTestRequest(chat, "gpt-4"), aibridgeTestRoute(aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai)), modelBuildOptions{ActiveAPIKeyID: apiKeyID, RecordHTTP: true}) require.NoError(t, err) - _, err = model.Generate(t.Context(), fantasy.Call{Prompt: []fantasy.Message{{ + _, err = model.LanguageModel().Generate(t.Context(), fantasy.Call{Prompt: []fantasy.Message{{ Role: fantasy.MessageRoleUser, Content: []fantasy.MessagePart{fantasy.TextPart{Text: "hello"}}, }}}) diff --git a/coderd/x/chatd/quickgen.go b/coderd/x/chatd/quickgen.go index fc18006bfe..43374403c3 100644 --- a/coderd/x/chatd/quickgen.go +++ b/coderd/x/chatd/quickgen.go @@ -141,7 +141,7 @@ type shortTextCandidate struct { provider string model string route aiGatewayModelRoute - lm fantasy.LanguageModel + lm chatprovider.Model providerOptions fantasy.ProviderOptions configOptions json.RawMessage } @@ -276,7 +276,7 @@ func (p *Server) maybeGenerateChatTitle( pasteText map[uuid.UUID]string, fallbackProvider string, fallbackConfig database.ChatModelConfig, - fallbackModel fantasy.LanguageModel, + fallbackModel chatprovider.Model, fallbackRoute aiGatewayModelRoute, modelOpts modelBuildOptions, generatedTitle *generatedChatTitle, @@ -372,7 +372,7 @@ func (p *Server) maybeGenerateChatTitle( ) } - title, err := generateTitle(candidateCtx, candidateModel, candidate.providerOptions, input) + title, err := generateTitle(candidateCtx, candidateModel.LanguageModel(), candidate.providerOptions, input) finishDebugRun(err) if err != nil { if overrideSet { @@ -413,7 +413,7 @@ func (p *Server) maybeGenerateChatTitle( func (p *Server) titleGenerationProviderOptions( ctx context.Context, - model fantasy.LanguageModel, + model chatprovider.Model, config database.ChatModelConfig, ) fantasy.ProviderOptions { callConfig := codersdk.ChatModelCallConfig{} @@ -427,12 +427,12 @@ func (p *Server) titleGenerationProviderOptions( } responsesOverride := chatprovider.OpenAIResponsesAPIOverride(callConfig.OpenAIConfig) providerOptions := chatprovider.ProviderOptionsFromChatModelConfig( - model, + model.LanguageModel(), callConfig.ProviderOptions, responsesOverride, ) return chatprovider.ApplyReasoningEffort( - model, + model.LanguageModel(), providerOptions, chatprovider.ResolveReasoningEffort(nil, callConfig.ReasoningEffort), responsesOverride, @@ -448,7 +448,7 @@ func (p *Server) newQuickgenDebugModel( route aiGatewayModelRoute, modelOpts modelBuildOptions, configOptions json.RawMessage, -) (fantasy.LanguageModel, error) { +) (chatprovider.Model, error) { debugOpts := modelOpts debugOpts.RecordHTTP = true debugModel, err := p.newModel(ctx, modelClientRequest{ @@ -459,15 +459,15 @@ func (p *Server) newQuickgenDebugModel( ConfigOptions: configOptions, }, route, debugOpts) if err != nil { - return nil, err + return chatprovider.Model{}, err } - return chatdebug.WrapModel(debugModel, debugSvc, chatdebug.RecorderOptions{ + return debugModel.WithLanguageModel(chatdebug.WrapModel(debugModel.LanguageModel(), debugSvc, chatdebug.RecorderOptions{ ChatID: chat.ID, OwnerID: chat.OwnerID, Provider: provider, Model: model, - }), nil + })), nil } func (p *Server) prepareQuickgenDebugCandidate( @@ -481,7 +481,7 @@ func (p *Server) prepareQuickgenDebugCandidate( historyTipMessageID int64, seedSummary map[string]any, logger slog.Logger, -) (context.Context, fantasy.LanguageModel, func(error)) { +) (context.Context, chatprovider.Model, func(error)) { finishDebugRun := func(error) {} if debugSvc == nil { return ctx, candidate.lm, finishDebugRun @@ -1341,7 +1341,7 @@ func (p *Server) generateTurnStatusLabel( assistantText string, fallbackProvider string, fallbackModelName string, - fallbackModel fantasy.LanguageModel, + fallbackModel chatprovider.Model, fallbackRoute aiGatewayModelRoute, modelOpts modelBuildOptions, configOptions json.RawMessage, @@ -1390,7 +1390,7 @@ func (p *Server) generateTurnStatusLabel( generatedLabel, err := generateStructuredTurnStatusLabel( candidateCtx, - candidateModel, + candidateModel.LanguageModel(), turnStatusLabelPrompt, input, ) diff --git a/coderd/x/chatd/quickgen_internal_test.go b/coderd/x/chatd/quickgen_internal_test.go index 615b7215a7..a761bebc96 100644 --- a/coderd/x/chatd/quickgen_internal_test.go +++ b/coderd/x/chatd/quickgen_internal_test.go @@ -588,7 +588,7 @@ func TestMaybeGenerateChatTitlePreservesUpdatedAt(t *testing.T) { nil, "openai", database.ChatModelConfig{Model: "test-model"}, - model, + chatprovider.NewModel(model, nil), aiGatewayModelRoute{}, modelBuildOptions{}, generated, @@ -657,7 +657,7 @@ func TestMaybeGenerateChatTitleAppliesModelConfigReasoningEffort(t *testing.T) { nil, fantasyopenai.Name, database.ChatModelConfig{Model: "gpt-4o-mini", Options: modelConfigRaw}, - model, + chatprovider.NewModel(model, nil), aiGatewayModelRoute{}, modelBuildOptions{}, &generatedChatTitle{}, @@ -884,7 +884,7 @@ func TestGenerateStructuredTitleWithUsage_OpenAICompatibleRequiredToolChoice(t * title, _, err := generateStructuredTitleWithUsage( t.Context(), - model, + model.LanguageModel(), nil, titleGenerationPrompt, "summarize failed workspace build logs", @@ -993,7 +993,7 @@ func newOpenAICompatStructuredOutputServer( return server, requests } -func openAICompatTestModel(t *testing.T, baseURL string) fantasy.LanguageModel { +func openAICompatTestModel(t *testing.T, baseURL string) chatprovider.Model { t.Helper() model, err := chatprovider.ModelFromConfig( @@ -1042,7 +1042,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, turnStatusLabelPrompt, "done") + label, err := generateStructuredTurnStatusLabel(t.Context(), model.LanguageModel(), turnStatusLabelPrompt, "done") require.NoError(t, err) require.Equal(t, "Submitted PR", label) require.Len(t, requests, 1) diff --git a/coderd/x/chatd/subagent_internal_test.go b/coderd/x/chatd/subagent_internal_test.go index b89d7218dc..1d2aeb86fc 100644 --- a/coderd/x/chatd/subagent_internal_test.go +++ b/coderd/x/chatd/subagent_internal_test.go @@ -540,7 +540,7 @@ func TestResolveChatModel_AIProviderDisabled(t *testing.T) { model, config, _, debugEnabled, resolvedProvider, resolvedModel, err := server.resolveChatModel(ctx, chat, modelBuildOptions{}) require.ErrorContains(t, err, "is disabled") - require.Nil(t, model) + require.False(t, model.Valid()) require.Equal(t, database.ChatModelConfig{}, config) require.False(t, debugEnabled) require.Empty(t, resolvedProvider) diff --git a/coderd/x/chatd/title_override.go b/coderd/x/chatd/title_override.go index 9abe961023..4056fdcfe1 100644 --- a/coderd/x/chatd/title_override.go +++ b/coderd/x/chatd/title_override.go @@ -4,7 +4,6 @@ import ( "context" "strings" - "charm.land/fantasy" "github.com/google/uuid" "golang.org/x/xerrors" @@ -61,10 +60,10 @@ func (p *Server) resolveTitleGenerationModelOverride( ctx context.Context, chat database.Chat, modelOpts modelBuildOptions, -) (database.ChatModelConfig, fantasy.LanguageModel, aiGatewayModelRoute, bool, error) { +) (database.ChatModelConfig, chatprovider.Model, aiGatewayModelRoute, bool, error) { raw, err := readTitleGenerationModelOverride(ctx, p.db) if err != nil { - return database.ChatModelConfig{}, nil, aiGatewayModelRoute{}, false, xerrors.Errorf( + return database.ChatModelConfig{}, chatprovider.Model{}, aiGatewayModelRoute{}, false, xerrors.Errorf( "read title generation model override: %w", err, ) @@ -82,17 +81,17 @@ func (p *Server) resolveTitleGenerationModelOverride( modelOverrideFailureModeHard, ) if err != nil { - return database.ChatModelConfig{}, nil, aiGatewayModelRoute{}, overrideSet, err + return database.ChatModelConfig{}, chatprovider.Model{}, aiGatewayModelRoute{}, overrideSet, err } if !overrideSet { - return database.ChatModelConfig{}, nil, aiGatewayModelRoute{}, false, nil + return database.ChatModelConfig{}, chatprovider.Model{}, aiGatewayModelRoute{}, false, nil } modelConfig = withResolvedReasoningEffort(modelConfig, overrideEffort) //nolint:gocritic // Title overrides need chatd-scoped provider reads for user-owned chats. route, err := p.resolveModelRouteForConfig(dbauthz.AsChatd(ctx), chat.OwnerID, modelConfig) if err != nil { - return database.ChatModelConfig{}, nil, aiGatewayModelRoute{}, true, err + return database.ChatModelConfig{}, chatprovider.Model{}, aiGatewayModelRoute{}, true, err } model, err := p.newModel(ctx, modelClientRequest{ Chat: chat, @@ -102,7 +101,7 @@ func (p *Server) resolveTitleGenerationModelOverride( ConfigOptions: modelConfig.Options, }, route, modelOpts) if err != nil { - return database.ChatModelConfig{}, nil, aiGatewayModelRoute{}, true, xerrors.Errorf( + return database.ChatModelConfig{}, chatprovider.Model{}, aiGatewayModelRoute{}, true, xerrors.Errorf( "create title generation model override: %w", err, ) diff --git a/coderd/x/chatd/title_override_internal_test.go b/coderd/x/chatd/title_override_internal_test.go index 204639bc2a..36ebef7226 100644 --- a/coderd/x/chatd/title_override_internal_test.go +++ b/coderd/x/chatd/title_override_internal_test.go @@ -26,6 +26,7 @@ import ( "github.com/coder/coder/v2/coderd/database/dbmock" "github.com/coder/coder/v2/coderd/util/ptr" "github.com/coder/coder/v2/coderd/x/chatd/chatprompt" + "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" "github.com/coder/coder/v2/coderd/x/chatd/chattest" "github.com/coder/coder/v2/codersdk" "github.com/coder/coder/v2/testutil" @@ -70,7 +71,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideUnset(t *testing.T) { nil, "openai", database.ChatModelConfig{Model: "fallback-chat-model"}, - fallbackModel, + chatprovider.NewModel(fallbackModel, nil), aiGatewayModelRoute{}, modelBuildOptions{}, generated, @@ -120,7 +121,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideReadDBError(t *testing.T) nil, "openai", database.ChatModelConfig{Model: "fallback-chat-model"}, - fallbackModel, + chatprovider.NewModel(fallbackModel, nil), aiGatewayModelRoute{}, modelBuildOptions{}, generated, @@ -169,7 +170,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideMalformedFallsThrough(t * nil, "openai", database.ChatModelConfig{Model: "fallback-chat-model"}, - fallbackModel, + chatprovider.NewModel(fallbackModel, nil), aiGatewayModelRoute{}, modelBuildOptions{}, generated, @@ -256,7 +257,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideSetUsable(t *testing.T) { nil, "openai", database.ChatModelConfig{Model: "fallback-chat-model"}, - fallbackModel, + chatprovider.NewModel(fallbackModel, nil), aiGatewayModelRoute{}, modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, generated, @@ -298,7 +299,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideSetUnusableSkips(t *testi nil, "openai", database.ChatModelConfig{Model: "fallback-chat-model"}, - fallbackModel, + chatprovider.NewModel(fallbackModel, nil), aiGatewayModelRoute{}, modelBuildOptions{}, generated, @@ -352,7 +353,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideCallFailureSkipsFallback( nil, "openai", database.ChatModelConfig{Model: "fallback-chat-model"}, - fallbackModel, + chatprovider.NewModel(fallbackModel, nil), aiGatewayModelRoute{}, modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, generated, @@ -396,7 +397,7 @@ func TestResolveManualTitleModel_TitleGenerationOverrideUnset(t *testing.T) { modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, ) require.NoError(t, err) - require.NotNil(t, model) + require.True(t, model.Valid()) require.Equal(t, preferredConfig, gotConfig) } @@ -445,7 +446,7 @@ func TestResolveManualTitleModel_TitleGenerationOverrideUnsetAIProvider(t *testi modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, ) require.NoError(t, err) - require.NotNil(t, model) + require.True(t, model.Valid()) require.Equal(t, preferredConfig, gotConfig) } @@ -480,7 +481,7 @@ func TestResolveManualTitleModel_TitleGenerationOverrideReadDBError(t *testing.T modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, ) require.NoError(t, err) - require.NotNil(t, model) + require.True(t, model.Valid()) require.Equal(t, preferredConfig, gotConfig) } @@ -512,7 +513,7 @@ func TestResolveManualTitleModel_TitleGenerationOverrideSetUsable(t *testing.T) modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, ) require.NoError(t, err) - require.NotNil(t, model) + require.True(t, model.Valid()) require.Equal(t, overrideConfig, gotConfig) } @@ -548,7 +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.Nil(t, model) + require.False(t, model.Valid()) require.Equal(t, database.ChatModelConfig{}, gotConfig) } @@ -644,7 +645,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.Nil(t, model) + require.False(t, model.Valid()) require.Equal(t, database.ChatModelConfig{}, gotConfig) }