mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
refactor(coderd/x/chatd): carry the resolved OpenAI transport on a model wrapper (#27703)
Stacked on #27683. The OpenAI wire format is decided when the client is built, then thrown away, so downstream sites recompute it from `(provider, modelID, override)`. Any disagreement fails silently: the SDK type-asserts the concrete provider options struct and discards every OpenAI option, and text attachments are dropped because Responses natively accepts only images and PDFs. This adds `chatprovider.Model`, which pairs a fantasy client with the transport resolved from that client's own identity. Its fields are unexported and only the constructor sets the transport, deriving it from the client, so no caller can pick a transport that disagrees with the client it wraps. `chatopenai.Transport`'s zero value is invalid and panics when read rather than defaulting to a wire format, following the existing precedent for construction invariants. `Model` is threaded through construction, the resolve paths, and the four struct fields that store a model for later request preparation. Terminal call sites keep taking `fantasy.LanguageModel` and receive `LanguageModel()`, which avoids a new package edge from `chatloop` and `chatadvisor` into `chatprovider`. No decisions move yet. The consumers still recompute the transport, and `UsesResponsesAPI` now delegates to `TransportFor` so the two agree by construction. #27704 makes the consumers read it from the model. > Mux prepared this PR on Mike's behalf.
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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) {
|
||||
|
||||
+30
-30
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
})
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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"}},
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"},
|
||||
)
|
||||
})
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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"}},
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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"}},
|
||||
}}})
|
||||
|
||||
+13
-13
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user