refactor: deprecate AIGatewayRoutingEnabled, remove direct chat routing (#26862)

This PR removes the now-dead direct-routing code:

- Deletes the direct routing implementation.
- Collapses the resolvedModelRoute discriminated union into aiGatewayModelRoute.
- Removes the dead providerKeys cascade.
- Deletes the preferredShortTextCandidates quickgen function.
- Simplifies the advisor override error handling.
- Deprecates the AIGatewayRoutingEnabled deployment option. It is now a no-op so as to not break existing deployments on upgrade.

Once direct routing was gone, the AI Gateway became mandatory for chat, which surfaced gaps in how the product behaves with the gateway disabled:

- Exposes ai-gateway-enabled to the frontend via embedded page metadata.
- Disables the chat composer via the existing AgentSetupNotice when the gateway is disabled, for both new and existing chats.
- Fixes nil/typed-nil chatDaemon panics on startup and shutdown when gateway is disabled.
- Fixes chat WebSocket from retrying the still-gated stream endpoint forever when the gateway is disabled.
This commit is contained in:
Cian Johnston
2026-07-01 20:15:03 +01:00
committed by GitHub
parent 3e0875d236
commit 4936ff9808
47 changed files with 1075 additions and 869 deletions
+2 -3
View File
@@ -796,9 +796,8 @@ chat:
# opt-in settings.
# (default: false, type: bool)
debugLoggingEnabled: false
# Route chat model requests through AI Gateway when both chat routing and AI
# Gateway are enabled. Otherwise, chat calls AI providers directly. Pending chats
# without API key metadata may need a retry or temporary direct routing.
# Deprecated: AI Gateway routing is now the only routing path. Setting this value
# has no effect. This option will be removed in a future release.
# (default: true, type: bool)
aiGatewayRoutingEnabled: true
aibridge:
+56 -34
View File
@@ -714,6 +714,7 @@ func New(options *Options) *API {
Telemetry: options.Telemetry,
Logger: options.Logger.Named("site"),
HideAITasks: options.DeploymentValues.HideAITasks.Value(),
AIGatewayEnabled: options.DeploymentValues.AI.BridgeConfig.Enabled.Value(),
})
if err != nil {
options.Logger.Fatal(ctx, "failed to initialize site handler", slog.Error(err))
@@ -825,37 +826,38 @@ func New(options *Options) *API {
providerAPIKeys = *options.ChatProviderAPIKeys
}
chatAIGatewayRoutingEnabled := options.DeploymentValues.AI.BridgeConfig.Enabled.Value() &&
options.DeploymentValues.AI.Chat.AIGatewayRoutingEnabled.Value()
api.chatDaemon = chatd.New(options.Pubsub, chatd.Config{
Logger: options.Logger.Named("chatd"),
Database: options.Database,
ReplicaID: api.ID,
StreamPartsDialer: options.ChatStreamPartsDialer,
MaxChatsPerAcquire: int32(maxChatsPerAcquire), //nolint:gosec // maxChatsPerAcquire is clamped to int32 range above.
ProviderAPIKeys: providerAPIKeys,
AllowBYOK: options.DeploymentValues.AI.BridgeConfig.AllowBYOK.Value(),
AllowBYOKSet: true,
AIBridgeTransportFactory: &api.AIBridgeTransportFactory,
AIGatewayRoutingEnabled: chatAIGatewayRoutingEnabled,
AlwaysEnableDebugLogs: options.DeploymentValues.AI.Chat.DebugLoggingEnabled.Value(),
Experiments: experiments,
AgentConn: api.agentProvider.AgentConn,
AgentInactiveDisconnectTimeout: api.AgentInactiveDisconnectTimeout,
InstructionLookupTimeout: options.ChatdInstructionLookupTimeout,
CreateWorkspace: api.chatCreateWorkspace,
StartWorkspace: api.chatStartWorkspace,
StopWorkspace: api.chatStopWorkspace,
WebpushDispatcher: options.WebPushDispatcher,
UsageTracker: options.WorkspaceUsageTracker,
PrometheusRegistry: options.PrometheusRegistry,
OIDCTokenSource: oidcMCPSrc,
NotificationsEnqueuer: options.NotificationsEnqueuer,
Auditor: &api.Auditor,
})
if !options.ChatWorkerDisabled {
api.chatDaemon.Start()
// AI Gateway is mandatory for chat. When the bridge is disabled
// the chat daemon stays nil and chat HTTP handlers return a
// service-unavailable error with a clear remediation message.
if options.DeploymentValues.AI.BridgeConfig.Enabled.Value() {
api.chatDaemon = chatd.New(options.Pubsub, chatd.Config{
Logger: options.Logger.Named("chatd"),
Database: options.Database,
ReplicaID: api.ID,
StreamPartsDialer: options.ChatStreamPartsDialer,
MaxChatsPerAcquire: int32(maxChatsPerAcquire), //nolint:gosec // maxChatsPerAcquire is clamped to int32 range above.
ProviderAPIKeys: providerAPIKeys,
AllowBYOK: options.DeploymentValues.AI.BridgeConfig.AllowBYOK.Value(),
AllowBYOKSet: true,
AIBridgeTransportFactory: &api.AIBridgeTransportFactory,
AlwaysEnableDebugLogs: options.DeploymentValues.AI.Chat.DebugLoggingEnabled.Value(),
Experiments: experiments,
AgentConn: api.agentProvider.AgentConn,
AgentInactiveDisconnectTimeout: api.AgentInactiveDisconnectTimeout,
InstructionLookupTimeout: options.ChatdInstructionLookupTimeout,
CreateWorkspace: api.chatCreateWorkspace,
StartWorkspace: api.chatStartWorkspace,
StopWorkspace: api.chatStopWorkspace,
WebpushDispatcher: options.WebPushDispatcher,
UsageTracker: options.WorkspaceUsageTracker,
PrometheusRegistry: options.PrometheusRegistry,
OIDCTokenSource: oidcMCPSrc,
NotificationsEnqueuer: options.NotificationsEnqueuer,
Auditor: &api.Auditor,
})
if !options.ChatWorkerDisabled {
api.chatDaemon.Start()
}
}
gitSyncLogger := options.Logger.Named("gitsync")
refresher := gitsync.NewRefresher(
@@ -864,9 +866,10 @@ func New(options *Options) *API {
gitSyncLogger.Named("refresher"),
quartz.NewReal(),
)
publishDiffStatusChange := chatDaemonPublishDiffStatusChangeFunc(api.chatDaemon)
api.gitSyncWorker = gitsync.NewWorker(options.Database,
refresher,
api.chatDaemon.PublishDiffStatusChange,
publishDiffStatusChange,
quartz.NewReal(),
gitSyncLogger,
)
@@ -2312,6 +2315,23 @@ type API struct {
workspaceAgentConnWatcher *workspaceconnwatcher.Watcher
}
// chatDaemonPublishDiffStatusChangeFunc returns chatDaemon's
// PublishDiffStatusChange method bound as a gitsync.PublishDiffStatusChangeFunc,
// or a true nil func value when chatDaemon is nil (AI Gateway disabled).
//
// This must not be inlined as chatDaemon.PublishDiffStatusChange: a method
// value on a nil pointer receiver is itself non-nil (it captures the
// receiver, it doesn't call the method), so gitsync.Worker's own "if
// publishDiffStatusChangeFn != nil" check would not catch a nil chatDaemon,
// and invoking the returned func would panic dereferencing the nil
// receiver.
func chatDaemonPublishDiffStatusChangeFunc(chatDaemon *chatd.Server) gitsync.PublishDiffStatusChangeFunc {
if chatDaemon == nil {
return nil
}
return chatDaemon.PublishDiffStatusChange
}
// Close waits for all WebSocket connections to drain before returning.
func (api *API) Close() error {
select {
@@ -2346,8 +2366,10 @@ func (api *API) Close() error {
api.Logger.Warn(context.Background(),
"chat diff refresh worker did not exit in time")
}
if err := api.chatDaemon.Close(); err != nil {
api.Logger.Warn(api.ctx, "close chat processor", slog.Error(err))
if api.chatDaemon != nil {
if err := api.chatDaemon.Close(); err != nil {
api.Logger.Warn(api.ctx, "close chat processor", slog.Error(err))
}
}
api.metricsCache.Close()
if api.updateChecker != nil {
+13
View File
@@ -8,6 +8,7 @@ import (
"github.com/go-chi/chi/v5"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestStripSlashesMW(t *testing.T) {
@@ -65,3 +66,15 @@ func TestStripSlashesMW(t *testing.T) {
})
}
}
// TestChatDaemonPublishDiffStatusChangeFunc verifies that
// chatDaemonPublishDiffStatusChangeFunc returns a true nil, not a method
// value bound to a nil receiver, when chatDaemon is nil. See that function
// for why the distinction matters. The non-nil path is covered by
// coderd/exp_chats_test.go.
func TestChatDaemonPublishDiffStatusChangeFunc(t *testing.T) {
t.Parallel()
fn := chatDaemonPublishDiffStatusChangeFunc(nil)
require.Nil(t, fn, "func value must be a true nil, not a bound method on a nil receiver")
}
+23
View File
@@ -263,6 +263,29 @@ func TestHealthz(t *testing.T) {
assert.Equal(t, "OK", string(body))
}
// TestAIGatewayDisabledStartupAndShutdown verifies the server starts and
// shuts down cleanly when the AI Gateway is disabled, leaving api.chatDaemon
// nil. It exercises the startup path (git sync worker callback binding, see
// chatDaemonPublishDiffStatusChangeFunc) and the shutdown path (Close on a
// nil daemon), both of which panicked before the nil guards were added.
func TestAIGatewayDisabledStartupAndShutdown(t *testing.T) {
t.Parallel()
dv := coderdtest.DeploymentValues(t)
require.NoError(t, dv.AI.BridgeConfig.Enabled.Set("false"))
// Constructing the client starts the server; t.Cleanup (registered by
// coderdtest.New) closes it at the end of the test, exercising the
// shutdown path that used to panic.
client := coderdtest.New(t, &coderdtest.Options{
DeploymentValues: dv,
})
res, err := client.Request(context.Background(), http.MethodGet, "/healthz", nil)
require.NoError(t, err)
defer res.Body.Close()
require.Equal(t, http.StatusOK, res.StatusCode)
}
func TestSwagger(t *testing.T) {
t.Parallel()
+71
View File
@@ -115,6 +115,22 @@ func maybeWriteLimitErr(ctx context.Context, rw http.ResponseWriter, err error)
return false
}
// requireChatDaemon reports whether the chat daemon exists, writing a 503
// Service Unavailable with a remediation message when it does not. The
// daemon is nil when the in-memory AI Gateway is disabled by deployment
// config. Operations that depend on it (creating, mutating, or streaming a
// chat) must call this; pure reads (e.g. getChat) do not.
func (api *API) requireChatDaemon(ctx context.Context, rw http.ResponseWriter) bool {
if api.chatDaemon != nil {
return true
}
httpapi.Write(ctx, rw, http.StatusServiceUnavailable, codersdk.Response{
Message: "AI Gateway must be enabled for Coder Agents functionality. Please contact your deployment administrator.",
Detail: "Set CODER_AI_GATEWAY_ENABLED=true (or ai-gateway-enabled in deployment YAML) to enable.",
})
return false
}
func publishChatConfigEvent(logger slog.Logger, ps dbpubsub.Pubsub, kind pubsub.ChatConfigEventKind, entityID uuid.UUID) {
payload, err := json.Marshal(pubsub.ChatConfigEvent{
Kind: kind,
@@ -1028,6 +1044,10 @@ func (api *API) postChats(rw http.ResponseWriter, r *http.Request) {
ctx := r.Context()
apiKey := httpmw.APIKey(r)
if !api.requireChatDaemon(ctx, rw) {
return
}
// Cap the raw request body to prevent excessive memory use
// from large dynamic tool schemas.
r.Body = http.MaxBytesReader(rw, r.Body, int64(2*maxSystemPromptLenBytes))
@@ -2648,6 +2668,10 @@ func (api *API) refreshChatContext(rw http.ResponseWriter, r *http.Request) {
return
}
if !api.requireChatDaemon(ctx, rw) {
return
}
updated, err := api.chatDaemon.RefreshChatContext(ctx, chat)
if err != nil {
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
@@ -2701,6 +2725,10 @@ func (api *API) patchChat(rw http.ResponseWriter, r *http.Request) {
return
}
if !api.requireChatDaemon(ctx, rw) {
return
}
aReq, commitAudit := audit.InitRequest[database.Chat](rw, &audit.RequestParams{
Audit: *api.Auditor.Load(),
Log: api.Logger,
@@ -3023,6 +3051,10 @@ func (api *API) postChatMessages(rw http.ResponseWriter, r *http.Request) {
chat := httpmw.ChatParam(r)
chatID := chat.ID
if !api.requireChatDaemon(ctx, rw) {
return
}
// Sending a message triggers LLM inference, requiring update
// permission on the org-scoped chat resource.
if !api.Authorize(r, policy.ActionUpdate, chat.RBACObject()) {
@@ -3232,6 +3264,10 @@ func (api *API) patchChatMessage(rw http.ResponseWriter, r *http.Request) {
apiKey := httpmw.APIKey(r)
chat := httpmw.ChatParam(r)
if !api.requireChatDaemon(ctx, rw) {
return
}
if !api.Authorize(r, policy.ActionUpdate, chat.RBACObject()) {
httpapi.ResourceNotFound(rw)
return
@@ -3353,6 +3389,10 @@ func (api *API) deleteChatQueuedMessage(rw http.ResponseWriter, r *http.Request)
chat := httpmw.ChatParam(r)
chatID := chat.ID
if !api.requireChatDaemon(ctx, rw) {
return
}
if !api.Authorize(r, policy.ActionUpdate, chat.RBACObject()) {
httpapi.ResourceNotFound(rw)
return
@@ -3403,6 +3443,10 @@ func (api *API) promoteChatQueuedMessage(rw http.ResponseWriter, r *http.Request
chat := httpmw.ChatParam(r)
chatID := chat.ID
if !api.requireChatDaemon(ctx, rw) {
return
}
// Promoting a queued message triggers LLM inference,
// requiring update permission on the org-scoped chat resource.
if !api.Authorize(r, policy.ActionUpdate, chat.RBACObject()) {
@@ -3528,6 +3572,10 @@ func (api *API) streamChat(rw http.ResponseWriter, r *http.Request) {
chatID := chat.ID
logger := api.Logger.Named("chat_streamer").With(slog.F("chat_id", chatID))
if !api.requireChatDaemon(ctx, rw) {
return
}
var afterMessageID int64
if v := r.URL.Query().Get("after_id"); v != "" {
var err error
@@ -3665,6 +3713,10 @@ func (api *API) interruptChat(rw http.ResponseWriter, r *http.Request) {
chatID := chat.ID
logger := api.Logger.Named("chat_interrupt").With(slog.F("chat_id", chatID))
if !api.requireChatDaemon(ctx, rw) {
return
}
if !api.Authorize(r, policy.ActionUpdate, chat.RBACObject()) {
httpapi.ResourceNotFound(rw)
return
@@ -3714,6 +3766,10 @@ func (api *API) reconcileInvalidChatState(rw http.ResponseWriter, r *http.Reques
chatID := chat.ID
logger := api.Logger.Named("chat_reconcile_invalid").With(slog.F("chat_id", chatID))
if !api.requireChatDaemon(ctx, rw) {
return
}
if !api.Authorize(r, policy.ActionUpdate, chat.RBACObject()) {
httpapi.ResourceNotFound(rw)
return
@@ -3761,6 +3817,10 @@ func (api *API) regenerateChatTitle(rw http.ResponseWriter, r *http.Request) {
apiKey := httpmw.APIKey(r)
chat := httpmw.ChatParam(r)
if !api.requireChatDaemon(ctx, rw) {
return
}
if !api.Authorize(r, policy.ActionUpdate, chat.RBACObject()) {
httpapi.ResourceNotFound(rw)
return
@@ -3807,6 +3867,10 @@ func (api *API) proposeChatTitle(rw http.ResponseWriter, r *http.Request) {
apiKey := httpmw.APIKey(r)
chat := httpmw.ChatParam(r)
if !api.requireChatDaemon(ctx, rw) {
return
}
if !api.Authorize(r, policy.ActionUpdate, chat.RBACObject()) {
httpapi.ResourceNotFound(rw)
return
@@ -7606,6 +7670,10 @@ func (api *API) postChatToolResults(rw http.ResponseWriter, r *http.Request) {
chat := httpmw.ChatParam(r)
apiKey := httpmw.APIKey(r)
if !api.requireChatDaemon(ctx, rw) {
return
}
// Submitting tool results resumes LLM inference,
// requiring update permission on the org-scoped chat resource.
if !api.Authorize(r, policy.ActionUpdate, chat.RBACObject()) {
@@ -7812,6 +7880,9 @@ func (api *API) getChatDebugRun(rw http.ResponseWriter, r *http.Request) {
func (api *API) streamChatParts(rw http.ResponseWriter, r *http.Request) {
ctx := r.Context()
chat := httpmw.ChatParam(r)
if !api.requireChatDaemon(ctx, rw) {
return
}
if err := api.chatDaemon.ServeStreamPartsAuthorized(rw, r, chat); err != nil {
api.Logger.Named("chat_stream_parts").Debug(ctx, "chat stream parts closed", slog.Error(err))
}
+32
View File
@@ -4562,6 +4562,38 @@ func TestGetChat(t *testing.T) {
requireSDKError(t, err, http.StatusNotFound)
})
// AIGatewayDisabled regression-tests that getChat is a pure DB read that
// still works when the AI Gateway is disabled and api.chatDaemon is nil.
// It builds the server without starting the test AI bridge daemon
// (unlike every other subtest here) and seeds the chat directly into the
// database, since the create-chat route itself requires the daemon.
t.Run("AIGatewayDisabled", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
values := coderdtest.DeploymentValues(t)
require.NoError(t, values.AI.BridgeConfig.Enabled.Set("false"))
opts := newChatTestOptions(t, values)
rawClient, _, api := coderdtest.NewWithAPI(t, opts)
client := codersdk.NewExperimentalClient(rawClient)
firstUser := coderdtest.CreateFirstUser(t, client.Client)
modelConfig := dbgen.ChatModelConfig(t, api.Database, database.ChatModelConfig{})
seededChat := dbgen.Chat(t, api.Database, database.Chat{
OrganizationID: firstUser.OrganizationID,
OwnerID: firstUser.UserID,
LastModelConfigID: modelConfig.ID,
Title: "ai gateway disabled chat",
})
chatResult, err := client.GetChat(ctx, seededChat.ID)
require.NoError(t, err)
require.Equal(t, seededChat.ID, chatResult.ID)
require.Equal(t, firstUser.UserID, chatResult.OwnerID)
require.Equal(t, modelConfig.ID, chatResult.LastModelConfigID)
require.Equal(t, "ai gateway disabled chat", chatResult.Title)
})
t.Run("FilesHydrated", func(t *testing.T) {
t.Parallel()
+12 -3
View File
@@ -143,6 +143,17 @@ func (api *API) workspaceAgentRPC(rw http.ResponseWriter, r *http.Request) {
slog.F("role", role))
}
// api.chatDaemon is a *chatd.Server that stays nil when AI Gateway is
// disabled. Assigning a nil *chatd.Server directly to the
// interface-typed ContextDirtyMarker field below would produce a
// non-nil interface value (a typed nil), defeating agentapi's own
// "if DirtyMarker != nil" check and panicking on first use. Only set
// the field when there's a real chat daemon to call into.
var contextDirtyMarker agentapi.ContextDirtyMarker
if api.chatDaemon != nil {
contextDirtyMarker = api.chatDaemon
}
agentAPI := agentapi.New(agentapi.Options{
AgentID: workspaceAgent.ID,
OwnerID: workspace.OwnerID,
@@ -180,9 +191,7 @@ func (api *API) workspaceAgentRPC(rw http.ResponseWriter, r *http.Request) {
// Optional:
UpdateAgentMetricsFn: api.UpdateAgentMetrics,
// chatDaemon is always constructed (only its worker is gated), so
// this is non-nil; agentapi treats a nil marker as "chatd absent".
ContextDirtyMarker: api.chatDaemon,
ContextDirtyMarker: contextDirtyMarker,
}, workspace, workspaceAgent)
streamID := tailnet.StreamID{
+11 -36
View File
@@ -4,6 +4,7 @@ import (
"context"
"database/sql"
"encoding/json"
"net/http"
"testing"
"time"
@@ -16,7 +17,6 @@ import (
"github.com/coder/coder/v2/coderd/aibridge"
"github.com/coder/coder/v2/coderd/database"
"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,7 +107,6 @@ func (p *Server) resolveAdvisorModelOverrideOrFallback(
advisorCfg codersdk.AdvisorConfig,
fallbackModel fantasy.LanguageModel,
fallbackCallConfig codersdk.ChatModelCallConfig,
providerKeys chatprovider.ProviderAPIKeys,
modelOpts modelBuildOptions,
logger slog.Logger,
) (fantasy.LanguageModel, codersdk.ChatModelCallConfig) {
@@ -117,7 +116,6 @@ func (p *Server) resolveAdvisorModelOverrideOrFallback(
advisorCfg,
fallbackModel,
fallbackCallConfig,
providerKeys,
modelOpts,
logger,
)
@@ -134,7 +132,6 @@ func (p *Server) newAdvisorRuntimeOrFallback(
advisorCfg codersdk.AdvisorConfig,
fallbackModel fantasy.LanguageModel,
fallbackCallConfig codersdk.ChatModelCallConfig,
providerKeys chatprovider.ProviderAPIKeys,
modelOpts modelBuildOptions,
logger slog.Logger,
) *chatadvisor.Runtime {
@@ -144,7 +141,6 @@ func (p *Server) newAdvisorRuntimeOrFallback(
advisorCfg,
fallbackModel,
fallbackCallConfig,
providerKeys,
modelOpts,
logger,
)
@@ -178,7 +174,6 @@ func TestResolveAdvisorModelOverride(t *testing.T) {
codersdk.AdvisorConfig{},
fallbackModel,
fallbackCallConfig,
chatprovider.ProviderAPIKeys{},
modelBuildOptions{},
logger,
)
@@ -202,7 +197,6 @@ func TestResolveAdvisorModelOverride(t *testing.T) {
codersdk.AdvisorConfig{ModelConfigID: uuid.New()},
fallbackModel,
fallbackCallConfig,
chatprovider.ProviderAPIKeys{OpenAI: "sk-test"},
modelBuildOptions{},
logger,
)
@@ -232,7 +226,6 @@ func TestResolveAdvisorModelOverride(t *testing.T) {
codersdk.AdvisorConfig{ModelConfigID: uuid.New()},
fallbackModel,
fallbackCallConfig,
chatprovider.ProviderAPIKeys{OpenAI: "sk-test"},
modelBuildOptions{},
logger,
)
@@ -265,7 +258,6 @@ func TestResolveAdvisorModelOverride(t *testing.T) {
codersdk.AdvisorConfig{ModelConfigID: configID},
fallbackModel,
fallbackCallConfig,
chatprovider.ProviderAPIKeys{OpenAI: "sk-test"},
modelBuildOptions{},
logger,
)
@@ -308,7 +300,6 @@ func TestResolveAdvisorModelOverride(t *testing.T) {
codersdk.AdvisorConfig{ModelConfigID: configID},
fallbackModel,
fallbackCallConfig,
chatprovider.ProviderAPIKeys{},
modelBuildOptions{},
logger,
)
@@ -339,20 +330,13 @@ func TestResolveAdvisorModelOverride(t *testing.T) {
}, nil
},
getAIProviderByID: func(context.Context, uuid.UUID) (database.AIProvider, error) {
return database.AIProvider{
ID: providerID,
Type: database.AIProviderTypeOpenai,
Enabled: true,
}, nil
},
getAIProviderKeysByProviderID: func(context.Context, uuid.UUID) ([]database.AIProviderKey, error) {
return []database.AIProviderKey{{
ProviderID: providerID,
APIKey: "sk-test",
}}, nil
return aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai), nil
},
}
p := newAdvisorTestServer(ctx, t, store)
p.aibridgeTransportFactory = aibridgeTestFactoryPointer(&aibridgeTestFactory{rt: roundTripFunc(func(*http.Request) (*http.Response, error) {
return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody}, nil
})})
gotModel, gotCfg := p.resolveAdvisorModelOverrideOrFallback(
ctx,
@@ -360,8 +344,7 @@ func TestResolveAdvisorModelOverride(t *testing.T) {
codersdk.AdvisorConfig{ModelConfigID: configID},
fallbackModel,
fallbackCallConfig,
chatprovider.ProviderAPIKeys{OpenAI: "sk-test"},
modelBuildOptions{},
modelBuildOptions{ActiveAPIKeyID: uuid.NewString()},
logger,
)
require.NotEqual(t, fantasy.LanguageModel(fallbackModel), gotModel,
@@ -375,7 +358,6 @@ func TestResolveAdvisorModelOverride(t *testing.T) {
require.NotNil(t, gotCfg.Temperature)
require.InDelta(t, 0.42, *gotCfg.Temperature, 1e-9)
})
t.Run("AIProviderIDResolvesOverrideProviderKeys", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitShort)
@@ -394,11 +376,7 @@ func TestResolveAdvisorModelOverride(t *testing.T) {
}, nil
},
getAIProviderByID: func(context.Context, uuid.UUID) (database.AIProvider, error) {
return database.AIProvider{
ID: providerID,
Type: database.AIProviderTypeOpenai,
Enabled: true,
}, nil
return aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai), nil
},
getAIProviderKeysByProviderID: func(context.Context, uuid.UUID) ([]database.AIProviderKey, error) {
return []database.AIProviderKey{{
@@ -408,6 +386,9 @@ func TestResolveAdvisorModelOverride(t *testing.T) {
},
}
p := newAdvisorTestServer(ctx, t, store)
p.aibridgeTransportFactory = aibridgeTestFactoryPointer(&aibridgeTestFactory{rt: roundTripFunc(func(*http.Request) (*http.Response, error) {
return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody}, nil
})})
gotModel, gotCfg := p.resolveAdvisorModelOverrideOrFallback(
ctx,
@@ -415,8 +396,7 @@ func TestResolveAdvisorModelOverride(t *testing.T) {
codersdk.AdvisorConfig{ModelConfigID: configID},
fallbackModel,
fallbackCallConfig,
chatprovider.ProviderAPIKeys{},
modelBuildOptions{},
modelBuildOptions{ActiveAPIKeyID: uuid.NewString()},
logger,
)
require.NotEqual(t, fantasy.LanguageModel(fallbackModel), gotModel)
@@ -451,7 +431,6 @@ func TestResolveAdvisorModelOverridePromotesAIBridgeErrors(t *testing.T) {
},
}
p := newAdvisorTestServer(ctx, t, store)
p.aiGatewayRoutingEnabled = true
ctx = aibridge.WithDelegatedAPIKeyID(ctx, uuid.NewString())
model, _, err := p.resolveAdvisorModelOverride(
@@ -460,7 +439,6 @@ func TestResolveAdvisorModelOverridePromotesAIBridgeErrors(t *testing.T) {
codersdk.AdvisorConfig{ModelConfigID: configID},
&chattest.FakeModel{ProviderName: "stub", ModelName: "stub"},
codersdk.ChatModelCallConfig{},
chatprovider.ProviderAPIKeys{},
modelBuildOptions{ActiveAPIKeyID: uuid.NewString()},
slog.Make(),
)
@@ -575,7 +553,6 @@ func TestNewAdvisorRuntime(t *testing.T) {
},
fallbackModel,
fallbackCallConfig,
chatprovider.ProviderAPIKeys{},
modelBuildOptions{},
logger,
)
@@ -600,7 +577,6 @@ func TestNewAdvisorRuntime(t *testing.T) {
},
fallbackModel,
fallbackCallConfig,
chatprovider.ProviderAPIKeys{},
modelBuildOptions{},
logger,
)
@@ -623,7 +599,6 @@ func TestNewAdvisorRuntime(t *testing.T) {
},
fallbackModel,
fallbackCallConfig,
chatprovider.ProviderAPIKeys{},
modelBuildOptions{},
logger,
)
+34 -65
View File
@@ -189,7 +189,6 @@ type Server struct {
recordingSem chan struct{}
aibridgeTransportFactory *atomic.Pointer[aibridge.TransportFactory]
aiGatewayRoutingEnabled bool
experiments codersdk.Experiments
// Configuration
@@ -270,7 +269,6 @@ func (p *Server) resolveAdvisorModelOverride(
advisorCfg codersdk.AdvisorConfig,
fallbackModel fantasy.LanguageModel,
fallbackCallConfig codersdk.ChatModelCallConfig,
providerKeys chatprovider.ProviderAPIKeys,
modelOpts modelBuildOptions,
logger slog.Logger,
) (fantasy.LanguageModel, codersdk.ChatModelCallConfig, error) {
@@ -319,10 +317,9 @@ func (p *Server) resolveAdvisorModelOverride(
ctx,
chat.OwnerID,
overrideConfig,
providerKeys,
)
if err != nil {
if p.shouldUseAIGatewayRouting() && overrideConfig.AIProviderID.Valid {
if overrideConfig.AIProviderID.Valid {
return nil, codersdk.ChatModelCallConfig{}, xerrors.Errorf("resolve advisor override route: %w", err)
}
logger.Warn(
@@ -340,7 +337,7 @@ func (p *Server) resolveAdvisorModelOverride(
ExtraHeaders: chatprovider.CoderHeaders(chat),
}, route, modelOpts)
if err != nil {
if p.shouldUseAIGatewayRouting() && overrideConfig.AIProviderID.Valid {
if overrideConfig.AIProviderID.Valid {
return nil, codersdk.ChatModelCallConfig{}, xerrors.Errorf("create advisor override model: %w", err)
}
logger.Warn(
@@ -361,7 +358,6 @@ func (p *Server) newAdvisorRuntime(
advisorCfg codersdk.AdvisorConfig,
fallbackModel fantasy.LanguageModel,
fallbackCallConfig codersdk.ChatModelCallConfig,
providerKeys chatprovider.ProviderAPIKeys,
modelOpts modelBuildOptions,
logger slog.Logger,
) (*chatadvisor.Runtime, error) {
@@ -371,7 +367,6 @@ func (p *Server) newAdvisorRuntime(
advisorCfg,
fallbackModel,
fallbackCallConfig,
providerKeys,
modelOpts,
logger,
)
@@ -2291,10 +2286,6 @@ func (p *Server) RegenerateChatTitle(
// keeping chat ownership authorization at the HTTP layer.
//nolint:gocritic // Non-admin users need chatd-scoped config reads here.
chatdCtx := dbauthz.AsChatd(ctx)
keys, err := p.resolveUserProviderAPIKeys(chatdCtx, chat.OwnerID, uuid.Nil)
if err != nil {
keys = chatprovider.ProviderAPIKeys{}
}
if err := p.acquireManualTitleLock(ctx, chat.ID); err != nil {
return database.Chat{}, err
}
@@ -2304,7 +2295,6 @@ func (p *Server) RegenerateChatTitle(
chatdCtx,
p.db,
chat,
keys,
)
if err != nil {
return database.Chat{}, p.recordManualTitleGenerationFailure(ctx, chat, err)
@@ -2355,16 +2345,12 @@ func (p *Server) ProposeChatTitle(
) (string, error) {
//nolint:gocritic // Non-admin users need chatd-scoped config reads here.
chatdCtx := dbauthz.AsChatd(ctx)
keys, err := p.resolveUserProviderAPIKeys(chatdCtx, chat.OwnerID, uuid.Nil)
if err != nil {
keys = chatprovider.ProviderAPIKeys{}
}
if err := p.acquireManualTitleLock(ctx, chat.ID); err != nil {
return "", err
}
defer p.releaseManualTitleLock(chatdCtx, chat.ID)
title, err := p.proposeChatTitleWithStore(chatdCtx, p.db, chat, keys)
title, err := p.proposeChatTitleWithStore(chatdCtx, p.db, chat)
if err != nil {
return "", p.recordManualTitleGenerationFailure(ctx, chat, err)
}
@@ -2412,7 +2398,6 @@ func (p *Server) generateManualTitleCandidate(
ctx context.Context,
store database.Store,
chat database.Chat,
keys chatprovider.ProviderAPIKeys,
) (manualTitleCandidateResult, error) {
if limitErr := p.checkUsageLimit(ctx, store, chat.OwnerID, uuid.NullUUID{UUID: chat.OrganizationID, Valid: true}); limitErr != nil {
return manualTitleCandidateResult{}, limitErr
@@ -2453,7 +2438,7 @@ func (p *Server) generateManualTitleCandidate(
}
}
model, modelConfig, modelKeys, err := p.resolveManualTitleModel(ctx, store, chat, keys, modelOpts)
model, modelConfig, err := p.resolveManualTitleModel(ctx, store, chat, modelOpts)
result := manualTitleCandidateResult{
modelConfig: modelConfig,
activeAPIKeyID: modelOpts.ActiveAPIKeyID,
@@ -2472,7 +2457,6 @@ func (p *Server) generateManualTitleCandidate(
debugSvc,
chat,
modelConfig,
modelKeys,
modelOpts,
messages,
model,
@@ -2503,9 +2487,8 @@ func (p *Server) proposeChatTitleWithStore(
ctx context.Context,
store database.Store,
chat database.Chat,
keys chatprovider.ProviderAPIKeys,
) (string, error) {
result, err := p.generateManualTitleCandidate(ctx, store, chat, keys)
result, err := p.generateManualTitleCandidate(ctx, store, chat)
if err != nil {
return "", err
}
@@ -2533,9 +2516,8 @@ func (p *Server) regenerateChatTitleWithStore(
ctx context.Context,
store database.Store,
chat database.Chat,
keys chatprovider.ProviderAPIKeys,
) (database.Chat, error) {
result, err := p.generateManualTitleCandidate(ctx, store, chat, keys)
result, err := p.generateManualTitleCandidate(ctx, store, chat)
if err != nil {
return database.Chat{}, err
}
@@ -2574,7 +2556,6 @@ func (p *Server) prepareManualTitleDebugRun(
debugSvc *chatdebug.Service,
chat database.Chat,
modelConfig database.ChatModelConfig,
keys chatprovider.ProviderAPIKeys,
modelOpts modelBuildOptions,
messages []database.ChatMessage,
fallbackModel fantasy.LanguageModel,
@@ -2583,10 +2564,10 @@ func (p *Server) prepareManualTitleDebugRun(
titleModel := fallbackModel
finishDebugRun := func(error) {}
route, routeErr := p.resolveModelRouteForConfig(ctx, chat.OwnerID, modelConfig, keys)
route, routeErr := p.resolveModelRouteForConfig(ctx, chat.OwnerID, modelConfig)
var routeProvider string
if routeErr == nil {
routeProvider, _ = route.providerHint()
routeProvider = string(route.Provider.Type)
} else if modelConfig.AIProviderID.Valid {
// Route resolution failed, but the linked provider still identifies the
// type for the debug run record. Best-effort: leave empty if disabled.
@@ -2750,18 +2731,16 @@ func (p *Server) resolveManualTitleModel(
ctx context.Context,
store database.Store,
chat database.Chat,
keys chatprovider.ProviderAPIKeys,
modelOpts modelBuildOptions,
) (fantasy.LanguageModel, database.ChatModelConfig, chatprovider.ProviderAPIKeys, error) {
overrideConfig, overrideModel, overrideKeys, _, overrideSet, overrideErr := p.resolveTitleGenerationModelOverride(
) (fantasy.LanguageModel, database.ChatModelConfig, error) {
overrideConfig, overrideModel, _, overrideSet, overrideErr := p.resolveTitleGenerationModelOverride(
ctx,
chat,
keys,
modelOpts,
)
if overrideErr != nil {
if overrideSet {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, xerrors.Errorf(
return nil, database.ChatModelConfig{}, xerrors.Errorf(
"resolve manual title generation model override: %w",
overrideErr,
)
@@ -2771,7 +2750,7 @@ func (p *Server) resolveManualTitleModel(
slog.Error(overrideErr),
)
} else if overrideSet {
return overrideModel, overrideConfig, overrideKeys, nil
return overrideModel, overrideConfig, nil
}
configs, err := store.GetEnabledChatModelConfigs(ctx)
@@ -2780,22 +2759,22 @@ func (p *Server) resolveManualTitleModel(
slog.F("chat_id", chat.ID),
slog.Error(err),
)
return p.resolveFallbackManualTitleModel(ctx, chat, keys, modelOpts)
return p.resolveFallbackManualTitleModel(ctx, chat, modelOpts)
}
config, ok := selectPreferredConfiguredShortTextModelConfig(configs)
if !ok {
return p.resolveFallbackManualTitleModel(ctx, chat, keys, modelOpts)
return p.resolveFallbackManualTitleModel(ctx, chat, modelOpts)
}
route, err := p.resolveModelRouteForConfig(ctx, chat.OwnerID, config, keys)
route, err := p.resolveModelRouteForConfig(ctx, chat.OwnerID, config)
if err != nil {
p.logger.Debug(ctx, "manual title preferred model unavailable",
slog.F("chat_id", chat.ID),
slog.F("model", config.Model),
slog.Error(err),
)
return p.resolveFallbackManualTitleModel(ctx, chat, keys, modelOpts)
return p.resolveFallbackManualTitleModel(ctx, chat, modelOpts)
}
model, err := p.newModel(ctx, modelClientRequest{
Chat: chat,
@@ -2809,28 +2788,27 @@ func (p *Server) resolveManualTitleModel(
slog.F("model", config.Model),
slog.Error(err),
)
return p.resolveFallbackManualTitleModel(ctx, chat, keys, modelOpts)
return p.resolveFallbackManualTitleModel(ctx, chat, modelOpts)
}
return model, config, route.directProviderKeys(), nil
return model, config, nil
}
func (p *Server) resolveFallbackManualTitleModel(
ctx context.Context,
chat database.Chat,
keys chatprovider.ProviderAPIKeys,
modelOpts modelBuildOptions,
) (fantasy.LanguageModel, database.ChatModelConfig, chatprovider.ProviderAPIKeys, error) {
) (fantasy.LanguageModel, database.ChatModelConfig, error) {
config, err := p.resolveModelConfig(ctx, chat)
if err != nil {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, xerrors.Errorf(
return nil, database.ChatModelConfig{}, xerrors.Errorf(
"resolve fallback manual title model config: %w",
err,
)
}
route, err := p.resolveModelRouteForConfig(ctx, chat.OwnerID, config, keys)
route, err := p.resolveModelRouteForConfig(ctx, chat.OwnerID, config)
if err != nil {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, err
return nil, database.ChatModelConfig{}, err
}
model, err := p.newModel(ctx, modelClientRequest{
Chat: chat,
@@ -2839,12 +2817,12 @@ func (p *Server) resolveFallbackManualTitleModel(
ExtraHeaders: chatprovider.CoderHeaders(chat),
}, route, modelOpts)
if err != nil {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, xerrors.Errorf(
return nil, database.ChatModelConfig{}, xerrors.Errorf(
"create fallback manual title model: %w",
err,
)
}
return model, config, route.directProviderKeys(), nil
return model, config, nil
}
func mergeManualTitleMessages(
@@ -3167,7 +3145,6 @@ type Config struct {
UsageTracker *workspacestats.UsageTracker
Clock quartz.Clock
AIBridgeTransportFactory *atomic.Pointer[aibridge.TransportFactory]
AIGatewayRoutingEnabled bool
Experiments codersdk.Experiments
PrometheusRegistry prometheus.Registerer
@@ -3264,7 +3241,6 @@ func New(ps pubsub.Pubsub, cfg Config) *Server {
return debugSvc
},
aibridgeTransportFactory: cfg.AIBridgeTransportFactory,
aiGatewayRoutingEnabled: cfg.AIGatewayRoutingEnabled,
experiments: cfg.Experiments,
pendingChatAcquireInterval: pendingChatAcquireInterval,
maxChatsPerAcquire: maxChatsPerAcquire,
@@ -3549,9 +3525,8 @@ func (p *Server) trackWorkspaceUsage(
type runChatResult struct {
FinalAssistantText string
StatusLabelModel fantasy.LanguageModel
ProviderKeys chatprovider.ProviderAPIKeys
FallbackProvider string
FallbackRoute resolvedModelRoute
FallbackRoute aiGatewayModelRoute
FallbackModel string
ModelBuildOptions modelBuildOptions
TriggerMessageID int64
@@ -4157,8 +4132,7 @@ func (p *Server) resolveChatModel(
) (
model fantasy.LanguageModel,
dbConfig database.ChatModelConfig,
keys chatprovider.ProviderAPIKeys,
route resolvedModelRoute,
route aiGatewayModelRoute,
debugEnabled bool,
resolvedProvider string,
resolvedModel string,
@@ -4166,29 +4140,25 @@ func (p *Server) resolveChatModel(
) {
dbConfig, err = p.resolveModelConfig(ctx, chat)
if err != nil {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, resolvedModelRoute{}, false, "", "", xerrors.Errorf("resolve model config: %w", err)
return nil, database.ChatModelConfig{}, aiGatewayModelRoute{}, false, "", "", xerrors.Errorf("resolve model config: %w", err)
}
if !dbConfig.Enabled {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, resolvedModelRoute{}, false, "", "", xerrors.Errorf("chat model config %s is disabled", dbConfig.ID)
return nil, database.ChatModelConfig{}, aiGatewayModelRoute{}, false, "", "", xerrors.Errorf("chat model config %s is disabled", dbConfig.ID)
}
route, err = p.resolveModelRouteForConfig(ctx, chat.OwnerID, dbConfig, chatprovider.ProviderAPIKeys{})
route, err = p.resolveModelRouteForConfig(ctx, chat.OwnerID, dbConfig)
if err != nil {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, resolvedModelRoute{}, false, "", "", err
return nil, database.ChatModelConfig{}, aiGatewayModelRoute{}, false, "", "", err
}
keys = route.directProviderKeys()
providerHint, err := route.providerHint()
if err != nil {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, resolvedModelRoute{}, false, "", "", err
}
providerHint := route.ModelProviderHint
resolvedProvider, resolvedModel, err = chatprovider.ResolveModelWithProviderHint(
dbConfig.Model,
providerHint,
)
if err != nil {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, resolvedModelRoute{}, false, "", "", xerrors.Errorf(
return nil, database.ChatModelConfig{}, aiGatewayModelRoute{}, false, "", "", xerrors.Errorf(
"resolve model metadata: %w", err,
)
}
@@ -4200,11 +4170,11 @@ func (p *Server) resolveChatModel(
ExtraHeaders: chatprovider.CoderHeaders(chat),
}, route, modelOpts)
if err != nil {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, resolvedModelRoute{}, false, "", "", xerrors.Errorf(
return nil, database.ChatModelConfig{}, aiGatewayModelRoute{}, false, "", "", xerrors.Errorf(
"create model: %w", err,
)
}
return model, dbConfig, keys, route, debugEnabled, resolvedProvider, resolvedModel, nil
return model, dbConfig, route, debugEnabled, resolvedProvider, resolvedModel, nil
}
func (p *Server) aiProviderConfig(ctx context.Context, provider database.AIProvider) (chatprovider.ConfiguredProvider, error) {
@@ -4765,7 +4735,6 @@ func (p *Server) generateFinalTurnStatusLabel(
runResult.FallbackModel,
runResult.StatusLabelModel,
runResult.FallbackRoute,
runResult.ProviderKeys,
runResult.ModelBuildOptions,
logger,
p.existingDebugService(),
+24 -6
View File
@@ -44,7 +44,10 @@ func TestActiveServer_ChainBrokenRecovery(t *testing.T) {
user, org, model := seedChatDependenciesWithProvider(t, db, "openai", openAIURL)
model = updateModelForChainMode(t, db, model)
server := newActiveTestServer(t, db, ps)
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "first user")
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
insertProviderResponseID(ctx, t, db, chat.ID, "first assistant", model.ID, previousResponseID)
@@ -93,7 +96,10 @@ func TestActiveServer_ChainBrokenRecoveryAppliesProviderPromptPrep(t *testing.T)
user, org, model := seedAnthropicChatDependencies(t, db, anthropicURL)
model = updateModelForChainMode(t, db, model)
server := newActiveTestServer(t, db, ps)
factory := chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath())
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "hello")
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
insertSystemTextMessage(ctx, t, db, chat.ID, "sys-1", model.ID)
@@ -139,7 +145,10 @@ func TestActiveServer_NonChainBrokenRetryPreservesChainMode(t *testing.T) {
user, org, model := seedChatDependenciesWithProvider(t, db, "openai", openAIURL)
model = updateModelForChainMode(t, db, model)
server := newActiveTestServer(t, db, ps)
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "first user")
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
insertProviderResponseID(ctx, t, db, chat.ID, "first assistant", model.ID, previousResponseID)
@@ -190,7 +199,10 @@ func TestActiveServer_ChainBrokenRecoveryPersistsAcrossGenerationActions(t *test
user, org, model := seedChatDependenciesWithProvider(t, db, "openai", openAIURL)
model = updateModelForChainMode(t, db, model)
server := newActiveTestServer(t, db, ps)
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "first user")
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
insertProviderResponseID(ctx, t, db, chat.ID, "first assistant", model.ID, previousResponseID)
@@ -230,7 +242,10 @@ func TestActiveServer_ChainBrokenWithoutChainModeIsSafe(t *testing.T) {
user, org, model := seedChatDependenciesWithProvider(t, db, "openai", openAIURL)
model = updateModelForChainMode(t, db, model)
server := newActiveTestServer(t, db, ps)
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "only user")
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
@@ -258,7 +273,10 @@ func TestActiveServer_ChainBrokenRecoveryDropsOrphanProviderToolCall(t *testing.
user, org, model := seedAnthropicChatDependencies(t, db, anthropicURL)
model = updateModelForChainMode(t, db, model)
server := newActiveTestServer(t, db, ps)
factory := chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath())
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "first user")
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
insertProviderResponseID(ctx, t, db, chat.ID, "first assistant", model.ID, previousResponseID)
+3 -7
View File
@@ -116,18 +116,14 @@ func (p *Server) scheduleDebugCleanup(
func (p *Server) newDebugAwareModel(
ctx context.Context,
req modelClientRequest,
route resolvedModelRoute,
route aiGatewayModelRoute,
opts modelBuildOptions,
) (fantasy.LanguageModel, bool, error) {
providerHint, err := route.providerHint()
provider, resolvedModel, err := chatprovider.ResolveModelWithProviderHint(req.ModelName, route.ModelProviderHint)
if err != nil {
return nil, false, err
}
provider, resolvedModel, err := chatprovider.ResolveModelWithProviderHint(req.ModelName, providerHint)
if err != nil {
return nil, false, err
}
route = route.withProviderHint(provider)
route.ModelProviderHint = provider
req.ModelName = resolvedModel
debugSvc := p.debugService()
+70 -29
View File
@@ -3,6 +3,10 @@ package chatd
import (
"context"
"database/sql"
"encoding/json"
"io"
"net/http"
"strconv"
"strings"
"sync"
"testing"
@@ -182,7 +186,7 @@ func TestResolveModelRouteForProviderTypeAIGatewayRequiresProvider(t *testing.T)
db.EXPECT().GetAIProviders(gomock.Any(), database.GetAIProvidersParams{}).Return(nil, nil)
server := &Server{db: db, aiGatewayRoutingEnabled: true}
server := &Server{db: db}
_, err := server.resolveModelRouteForProviderType(
ctx,
uuid.New(),
@@ -766,6 +770,23 @@ func withChatMessageAPIKeyID(message database.ChatMessage, apiKeyID string) data
return message
}
// requireOutgoingRequestModel asserts that the outgoing request body
// requests wantModel. This is so that mock transports can still
// verify the outgoing request asked for the expected model.
func requireOutgoingRequestModel(t testing.TB, req *http.Request, wantModel string) {
t.Helper()
body, err := io.ReadAll(req.Body)
require.NoError(t, err)
req.Body = io.NopCloser(strings.NewReader(string(body)))
var decoded struct {
Model string `json:"model"`
}
require.NoError(t, json.Unmarshal(body, &decoded))
require.Equal(t, wantModel, decoded.Model)
}
func TestRegenerateChatTitle_PersistsAndBroadcasts(t *testing.T) {
t.Parallel()
@@ -821,30 +842,41 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts(t *testing.T) {
require.NoError(t, err)
defer cancelSub()
serverURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
require.Equal(t, "gpt-4o-mini", req.Model)
return chattest.OpenAINonStreamingResponse("{\"title\":\"" + wantTitle + "\"}")
})
// Title generation routes through the transport factory, so the model
// response is synthesized by the RoundTripper (see aibridgeTestFactory).
factory := &aibridgeTestFactory{rt: roundTripFunc(func(req *http.Request) (*http.Response, error) {
requireOutgoingRequestModel(t, req, modelConfig.Model)
text := strconv.Quote(`{"title":"` + wantTitle + `"}`)
body := `{"id":"resp_test","object":"response","created_at":0,"status":"completed","model":"gpt-4o-mini","output":[{"id":"msg_test","type":"message","role":"assistant","content":[{"type":"output_text","text":` + text + `}]}],"usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}`
return &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(body)),
Request: req,
}, nil
})}
server := &Server{
db: db,
logger: logger,
pubsub: pubsub,
configCache: newChatConfigCache(context.Background(), db, clock),
db: db,
logger: logger,
pubsub: pubsub,
configCache: newChatConfigCache(context.Background(), db, clock),
aibridgeTransportFactory: aibridgeTestFactoryPointer(factory),
}
db.EXPECT().GetChatModelConfigByID(gomock.Any(), modelConfigID).Return(modelConfig, nil)
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(database.AIProvider{
ID: providerID,
Name: "primary-openai",
Type: database.AIProviderTypeOpenai,
Enabled: true,
BaseUrl: serverURL,
}, nil).AnyTimes()
db.EXPECT().GetAIProviders(gomock.Any(), gomock.Any()).Return([]database.AIProvider{{
ID: providerID,
Name: "primary-openai",
Type: database.AIProviderTypeOpenai,
Enabled: true,
BaseUrl: serverURL,
}}, nil).AnyTimes()
db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes()
db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), gomock.Any()).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes()
@@ -952,7 +984,9 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts_IdleChatReleasesManualLock(t
ownerID := uuid.New()
chatID := uuid.New()
modelConfigID := uuid.New()
providerID := uuid.New()
userPrompt := "review pull request 23633 and fix review threads"
activeAPIKeyID := "key-" + uuid.NewString()
wantTitle := "Review PR 23633"
chat := database.Chat{
@@ -965,7 +999,6 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts_IdleChatReleasesManualLock(t
lockedChat := chat
lockedChat.WorkerID = uuid.NullUUID{UUID: manualTitleLockWorkerID, Valid: true}
lockedChat.StartedAt = sql.NullTime{Time: time.Now(), Valid: true}
providerID := uuid.New()
modelConfig := database.ChatModelConfig{
ID: modelConfigID,
Model: "gpt-4o-mini",
@@ -994,30 +1027,40 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts_IdleChatReleasesManualLock(t
require.NoError(t, err)
defer cancelSub()
serverURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
require.Equal(t, "gpt-4o-mini", req.Model)
return chattest.OpenAINonStreamingResponse("{\"title\":\"" + wantTitle + "\"}")
})
// Title generation routes through the transport factory, so the model
// response is synthesized by the RoundTripper (see aibridgeTestFactory).
factory := &aibridgeTestFactory{rt: roundTripFunc(func(req *http.Request) (*http.Response, error) {
requireOutgoingRequestModel(t, req, modelConfig.Model)
text := strconv.Quote(`{"title":"` + wantTitle + `"}`)
body := `{"id":"resp_test","object":"response","created_at":0,"status":"completed","model":"gpt-4o-mini","output":[{"id":"msg_test","type":"message","role":"assistant","content":[{"type":"output_text","text":` + text + `}]}],"usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}`
return &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(body)),
Request: req,
}, nil
})}
server := &Server{
db: db,
logger: logger,
pubsub: pubsub,
configCache: newChatConfigCache(context.Background(), db, clock),
db: db,
logger: logger,
pubsub: pubsub,
configCache: newChatConfigCache(context.Background(), db, clock),
aibridgeTransportFactory: aibridgeTestFactoryPointer(factory),
}
db.EXPECT().GetChatModelConfigByID(gomock.Any(), modelConfigID).Return(modelConfig, nil)
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(database.AIProvider{
ID: providerID,
Name: "primary-openai",
Type: database.AIProviderTypeOpenai,
Enabled: true,
BaseUrl: serverURL,
}, nil).AnyTimes()
db.EXPECT().GetAIProviders(gomock.Any(), gomock.Any()).Return([]database.AIProvider{{
ID: providerID,
Name: "primary-openai",
Type: database.AIProviderTypeOpenai,
Enabled: true,
BaseUrl: serverURL,
}}, nil).AnyTimes()
db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes()
db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), gomock.Any()).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes()
@@ -1030,12 +1073,12 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts_IdleChatReleasesManualLock(t
LimitVal: manualTitleMessageWindowLimit,
},
).Return([]database.ChatMessage{
mustChatMessage(
withChatMessageAPIKeyID(mustChatMessage(
t,
database.ChatMessageRoleUser,
database.ChatMessageVisibilityBoth,
codersdk.ChatMessageText(userPrompt),
),
), activeAPIKeyID),
mustChatMessage(
t,
database.ChatMessageRoleAssistant,
@@ -3558,10 +3601,9 @@ func TestPrepareManualTitleDebugRun_RouteFailureDerivesProviderFromConfig(t *tes
)
server := &Server{
db: db,
logger: logger,
aiGatewayRoutingEnabled: true,
allowBYOK: true,
db: db,
logger: logger,
allowBYOK: true,
}
debugSvc := chatdebug.NewService(db, logger, nil)
fallbackModel := &chattest.FakeModel{ProviderName: "stub", ModelName: "stub"}
@@ -3571,7 +3613,6 @@ func TestPrepareManualTitleDebugRun_RouteFailureDerivesProviderFromConfig(t *tes
debugSvc,
chat,
modelConfig,
chatprovider.ProviderAPIKeys{},
modelBuildOptions{},
nil,
fallbackModel,
+6
View File
@@ -41,9 +41,11 @@ func TestActiveServer_RetryStatePersistedDuringBackoff(t *testing.T) {
return chattest.OpenAIStreamingResponse(openAITextChunksWithStop("recovered")...)
})
user, org, model := seedChatDependenciesWithProvider(t, db, "openai", openAIURL)
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.Clock = clock
cfg.Logger = sink.Logger()
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "hello")
@@ -108,8 +110,10 @@ func TestActiveServer_RetryStreamSilenceTimeoutAndClassification(t *testing.T) {
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.PrometheusRegistry = reg
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "hello")
@@ -150,12 +154,14 @@ func TestActiveServer_RetryStreamSilenceTimeoutAndClassification(t *testing.T) {
return chattest.OpenAIStreamingResponse(openAITextChunksWithStop("recovered")...)
})
user, org, model := seedChatDependenciesWithProvider(t, db, "openai", openAIURL)
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.Clock = clock
cfg.Logger = sink.Logger()
cfg.PrometheusRegistry = reg
cfg.PendingChatAcquireInterval = 30 * time.Minute
cfg.ChatHeartbeatInterval = 30 * time.Minute
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "hello")
+262 -55
View File
@@ -165,6 +165,7 @@ func newWorkspaceToolTestServer(
ps dbpubsub.Pubsub,
agentID uuid.UUID,
planContent string,
overrides ...func(cfg *chatd.Config),
) *chatd.Server {
t.Helper()
@@ -183,12 +184,15 @@ func newWorkspaceToolTestServer(
return io.NopCloser(strings.NewReader("")), "", nil
}).AnyTimes()
return newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AgentConn = func(_ context.Context, gotAgentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, agentID, gotAgentID)
return mockConn, func() {}, nil
}
})
configOverrides := append([]func(cfg *chatd.Config){
func(cfg *chatd.Config) {
cfg.AgentConn = func(_ context.Context, gotAgentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, agentID, gotAgentID)
return mockConn, func() {}, nil
}
},
}, overrides...)
return newActiveTestServer(t, db, ps, configOverrides...)
}
func TestSubagentChatExcludesWorkspaceProvisioningTools(t *testing.T) {
@@ -813,7 +817,9 @@ func TestExploreChatUsesPersistedMCPSnapshot(t *testing.T) {
mockConn.EXPECT().ReadFile(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).
Return(io.NopCloser(strings.NewReader("")), "", nil).AnyTimes()
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, dbAgent.ID, agentID)
return mockConn, func() {}, nil
@@ -894,7 +900,10 @@ func TestRootExploreChatStaysBuiltinOnlyAtRuntime(t *testing.T) {
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
server := newActiveTestServer(t, db, ps)
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
exploreChat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
@@ -980,7 +989,10 @@ func TestRootExploreChatExcludesWebSearchProviderToolAtRuntime(t *testing.T) {
},
)
server := newActiveTestServer(t, db, ps)
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
exploreChat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
@@ -1104,7 +1116,10 @@ func TestExploreChatSendMessageCannotMutateMCPSnapshot(t *testing.T) {
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
server := newActiveTestServer(t, db, ps)
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
rootChat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
@@ -1303,7 +1318,9 @@ func TestPlanModeRootChatAllowsApprovedExternalMCPTools(t *testing.T) {
mockConn.EXPECT().ReadFile(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).
Return(io.NopCloser(strings.NewReader("")), "", nil).AnyTimes()
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, dbAgent.ID, agentID)
return mockConn, func() {}, nil
@@ -1748,7 +1765,10 @@ func TestPlanTurnPromptContract(t *testing.T) {
err := db.UpsertChatPlanModeInstructions(dbauthz.AsSystemRestricted(ctx), planModeInstructions)
require.NoError(t, err)
ws, dbAgent := seedWorkspaceWithAgent(t, db, user.ID)
server := newWorkspaceToolTestServer(t, db, ps, dbAgent.ID, "# Plan\n")
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newWorkspaceToolTestServer(t, db, ps, dbAgent.ID, "# Plan\n", func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
OwnerID: user.ID,
@@ -2042,7 +2062,9 @@ func TestAutoPromoteQueuedMessagesPreservesPerTurnModelOrder(t *testing.T) {
}
})
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
// Disable periodic polling so chained promotions must be driven by
// signalWake.
cfg.PendingChatAcquireInterval = time.Hour
@@ -2220,7 +2242,9 @@ func TestInterruptAutoPromotionIgnoresLaterUsageLimitIncrease(t *testing.T) {
)
})
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
cfg.Clock = clock
// Keep periodic polling frozen so request handoff is synchronized
// through explicit mock channels.
@@ -2664,7 +2688,8 @@ func TestRecoverStaleRequiresActionChat(t *testing.T) {
db, ps, rawDB := dbtestutil.NewDBWithSQLDB(t)
ctx := testutil.Context(t, testutil.WaitLong)
user, org, model := seedChatDependencies(t, db)
openAIURL := chattest.OpenAI(t)
user, org, model := seedChatDependenciesWithProvider(t, db, "openai", openAIURL)
toolName := "my_dynamic_tool"
dynamicToolsJSON, err := json.Marshal([]mcpgo.Tool{{
@@ -2737,7 +2762,10 @@ func TestRecoverStaleRequiresActionChat(t *testing.T) {
time.Now().Add(-time.Hour), created.Chat.ID)
require.NoError(t, err)
server := newTestServer(t, db, ps, uuid.New())
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newTestServer(t, db, ps, uuid.New(), func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
server.Start()
chatResult := waitForTerminalChat(ctx, t, db, created.Chat.ID)
@@ -2765,7 +2793,8 @@ func TestNewReplicaRecoversStaleChatFromDeadReplica(t *testing.T) {
db, ps, rawDB := dbtestutil.NewDBWithSQLDB(t)
ctx := testutil.Context(t, testutil.WaitLong)
user, org, model := seedChatDependencies(t, db)
openAIURL := chattest.OpenAI(t)
user, org, model := seedChatDependenciesWithProvider(t, db, "openai", openAIURL)
content, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{
codersdk.ChatMessageText("hello"),
@@ -2807,7 +2836,10 @@ func TestNewReplicaRecoversStaleChatFromDeadReplica(t *testing.T) {
require.NoError(t, err)
newWorkerID := uuid.New()
server := newTestServer(t, db, ps, newWorkerID)
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newTestServer(t, db, ps, newWorkerID, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
// Start a new replica. It should recover the stale chat on
// startup.
server.Start()
@@ -3047,7 +3079,9 @@ func TestPersistToolResultWithBinaryData(t *testing.T) {
}, nil).
AnyTimes()
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, dbAgent.ID, agentID)
return mockConn, func() {}, nil
@@ -3158,6 +3192,7 @@ func TestRequiresActionChatPersistsWaitingStatusLabel(t *testing.T) {
mockPush := &mockWebpushDispatcher{}
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := chatd.New(ps, chatd.Config{
Logger: logger,
Database: db,
@@ -3165,6 +3200,7 @@ func TestRequiresActionChatPersistsWaitingStatusLabel(t *testing.T) {
PendingChatAcquireInterval: 10 * time.Millisecond,
InFlightChatStaleAfter: testutil.WaitSuperLong,
WebpushDispatcher: mockPush,
AIBridgeTransportFactory: chatAIGatewayTransportFactoryPointer(factory),
})
t.Cleanup(func() {
require.NoError(t, server.Close())
@@ -3301,6 +3337,7 @@ func TestActiveServer_InterruptionBehavior(t *testing.T) {
mockConn.EXPECT().ReadFileLines(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Times(0)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, dbAgent.ID, agentID)
return mockConn, func() {}, nil
@@ -3407,6 +3444,7 @@ func TestActiveServer_InterruptionBehavior(t *testing.T) {
}).Times(1)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, dbAgent.ID, agentID)
return mockConn, func() {}, nil
@@ -3523,7 +3561,9 @@ func TestActiveServer_InterruptionBehavior(t *testing.T) {
},
})
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "search for coder")
testutil.TryReceive(ctx, t, providerToolStarted)
queued, err := server.SendMessage(ctx, chatd.SendMessageOptions{
@@ -3618,6 +3658,7 @@ func TestActiveServer_InterruptionBehavior(t *testing.T) {
mockConn.EXPECT().ReadFileLines(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Times(0)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, dbAgent.ID, agentID)
return mockConn, func() {}, nil
@@ -3719,7 +3760,9 @@ func TestActiveServer_InterruptionBehavior(t *testing.T) {
},
})
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "think")
testutil.TryReceive(ctx, t, reasoningStarted)
queued, err := server.SendMessage(ctx, chatd.SendMessageOptions{
@@ -3773,7 +3816,10 @@ func TestActiveServer_DynamicToolsAndStopAfterToolBehavior(t *testing.T) {
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
dynamicToolsJSON := dynamicToolJSON(t, "my_dynamic_tool")
server := newActiveTestServer(t, db, ps)
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
OwnerID: user.ID,
@@ -3830,7 +3876,10 @@ func TestActiveServer_DynamicToolsAndStopAfterToolBehavior(t *testing.T) {
})
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
ws, dbAgent := seedWorkspaceWithAgent(t, db, user.ID)
server := newWorkspaceToolTestServer(t, db, ps, dbAgent.ID, "# Plan\n")
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newWorkspaceToolTestServer(t, db, ps, dbAgent.ID, "# Plan\n", func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
@@ -3877,7 +3926,10 @@ func TestActiveServer_DynamicToolsAndStopAfterToolBehavior(t *testing.T) {
})
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
ws, dbAgent := seedWorkspaceWithAgent(t, db, user.ID)
server := newWorkspaceToolTestServer(t, db, ps, dbAgent.ID, "# Plan\n")
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newWorkspaceToolTestServer(t, db, ps, dbAgent.ID, "# Plan\n", func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
@@ -3953,7 +4005,10 @@ func TestDynamicToolCallPausesAndResumes(t *testing.T) {
// server without an agent connection, so the built-in tools
// are never invoked because the only tool call targets our
// dynamic tool.
server := newActiveTestServer(t, db, ps)
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
// Create a chat with a dynamic tool.
dynamicToolsJSON, err := json.Marshal([]mcpgo.Tool{{
@@ -4122,7 +4177,10 @@ func TestDynamicToolNamedProposePlanRemainsAvailableOutsidePlanMode(t *testing.T
})
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
server := newActiveTestServer(t, db, ps)
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
dynamicToolsJSON, err := json.Marshal([]mcpgo.Tool{{
Name: "propose_plan",
@@ -4236,7 +4294,10 @@ func TestDynamicToolCallMixedWithBuiltIn(t *testing.T) {
})
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
server := newActiveTestServer(t, db, ps)
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
// Create a chat with a dynamic tool.
dynamicToolsJSON, err := json.Marshal([]mcpgo.Tool{{
@@ -4376,7 +4437,10 @@ func TestSubmitToolResultsConcurrency(t *testing.T) {
})
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
server := newActiveTestServer(t, db, ps)
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
// Create a chat with a dynamic tool.
dynamicToolsJSON, err := json.Marshal([]mcpgo.Tool{{
@@ -5023,7 +5087,9 @@ func TestStoppedWorkspaceWithPersistedAgentBindingDoesNotBlockChat(t *testing.T)
}).Do()
var dialCalls atomic.Int32
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
_ = newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
cfg.AgentConn = func(ctx context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
dialCalls.Add(1)
require.Equal(t, dbAgent.ID, agentID)
@@ -5231,17 +5297,22 @@ func newTestServer(
db database.Store,
ps dbpubsub.Pubsub,
replicaID uuid.UUID,
overrides ...func(*chatd.Config),
) *chatd.Server {
t.Helper()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
server := chatd.New(ps, chatd.Config{
cfg := chatd.Config{
Logger: logger,
Database: db,
ReplicaID: replicaID,
PendingChatAcquireInterval: testutil.WaitLong,
Experiments: codersdk.ExperimentsKnown,
})
}
for _, o := range overrides {
o(&cfg)
}
server := chatd.New(ps, cfg)
t.Cleanup(func() {
require.NoError(t, server.Close())
})
@@ -5372,7 +5443,6 @@ func TestActiveServer_RoutingPreservesAPIKeyAfterCompaction(t *testing.T) {
_ = newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
cfg.AIGatewayRoutingEnabled = true
cfg.AllowBYOK = true
cfg.AllowBYOKSet = true
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
@@ -5451,6 +5521,7 @@ func TestActiveServer_CompactionRecordsMetric(t *testing.T) {
Times(1)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
cfg.PrometheusRegistry = reg
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, dbAgent.ID, agentID)
@@ -5544,6 +5615,7 @@ func TestActiveServer_Compaction(t *testing.T) {
Times(1)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, dbAgent.ID, agentID)
return mockConn, func() {}, nil
@@ -5614,7 +5686,9 @@ func TestActiveServer_Compaction(t *testing.T) {
user, org, model := seedAnthropicChatDependencies(t, db, anthropicURL)
model = updateChatModelCompressionThreshold(t, db, model, contextLimit, thresholdPercent)
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "finish with high usage")
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
@@ -5667,6 +5741,7 @@ func TestActiveServer_Compaction(t *testing.T) {
reg := prometheus.NewRegistry()
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
cfg.Logger = logSink.Logger()
cfg.PrometheusRegistry = reg
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
@@ -5807,7 +5882,9 @@ func TestActiveServer_BasicAssistantGenerationAndPromptPreparation(t *testing.T)
model.ContextLimit = 4096
model = updateChatModelContextLimit(t, db, model)
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "hello")
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
insertSystemTextMessage(ctx, t, db, chat.ID, "sys-2", model.ID)
@@ -5847,7 +5924,9 @@ func TestActiveServer_BasicAssistantGenerationAndPromptPreparation(t *testing.T)
requireTextPart(t, last, "done")
requests = newAnthropicRequestRecorder()
server = newActiveTestServer(t, db, ps)
server = newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
})
planChat := createPlanSubagentChatWithHistory(ctx, t, db, org.ID, user.ID, model.ID)
_, err = server.SendMessage(ctx, chatd.SendMessageOptions{
ChatID: planChat.ID,
@@ -5896,6 +5975,7 @@ func TestActiveServer_ToolExecutionAndPolicy(t *testing.T) {
mockConn.EXPECT().WriteFile(gomock.Any(), gomock.Any(), gomock.Any()).Times(0)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, dbAgent.ID, agentID)
return mockConn, func() {}, nil
@@ -5948,7 +6028,11 @@ func TestActiveServer_ToolExecutionAndPolicy(t *testing.T) {
model.Model = "gpt-5.5"
model = updateChatModelContextLimit(t, db, model)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) { cfg.AllowBYOKSet = true; cfg.AllowBYOK = false })
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
cfg.AllowBYOKSet = true
cfg.AllowBYOK = false
})
result := codersdk.ChatMessageToolResult(
"computer-call",
"computer",
@@ -6035,6 +6119,7 @@ func TestActiveServer_ToolExecutionAndPolicy(t *testing.T) {
Times(1)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, dbAgent.ID, agentID)
return mockConn, func() {}, nil
@@ -6116,6 +6201,7 @@ func TestActiveServer_ToolExecutionAndPolicy(t *testing.T) {
Times(1)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, dbAgent.ID, agentID)
return mockConn, func() {}, nil
@@ -6168,6 +6254,7 @@ func TestActiveServer_RecordsGenerationMetrics(t *testing.T) {
})
user, org, model := seedChatDependenciesWithProvider(t, db, "openai", openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
cfg.PrometheusRegistry = reg
})
@@ -6279,6 +6366,7 @@ func TestActiveServer_ToolErrorRecordsMetric(t *testing.T) {
tt.setupAgent(mockConn)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
cfg.PrometheusRegistry = reg
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, dbAgent.ID, agentID)
@@ -6568,7 +6656,9 @@ func TestActiveServer_AnthropicUsageMatchesFinalDelta(t *testing.T) {
})
user, org, model := seedAnthropicChatDependencies(t, db, anthropicURL)
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "hello")
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
@@ -6601,6 +6691,7 @@ func TestActiveServer_ChatTurnDebugRunRecordsStreamStep(t *testing.T) {
user, org, model := seedAnthropicChatDependencies(t, db, anthropicURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
cfg.AlwaysEnableDebugLogs = true
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "hello debug")
@@ -6711,6 +6802,7 @@ func TestActiveServer_ChatTurnDebugRunRecordsMultipleStreamSteps(t *testing.T) {
Times(1)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
cfg.AlwaysEnableDebugLogs = true
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, dbAgent.ID, agentID)
@@ -6799,7 +6891,9 @@ func TestActiveServer_AnthropicSanitizesProviderToolBeforeRequest(t *testing.T)
})
user, org, model := seedAnthropicChatDependencies(t, db, anthropicURL)
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "search for coder")
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
insertOrphanProviderToolCall(ctx, t, db, chat.ID, model.ID)
@@ -6849,7 +6943,9 @@ func TestActiveServer_AnthropicProviderToolPreRequestGuard(t *testing.T) {
user, org, model := seedAnthropicChatDependencies(t, db, anthropicURL)
model = updateChatModelCallConfig(t, db, model, callConfig)
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "search")
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
insertProviderToolPairMessageWithLocalTool(ctx, t, db, chat.ID, model.ID, "ws-allowed")
@@ -6884,7 +6980,9 @@ func TestActiveServer_AnthropicProviderToolPreRequestGuard(t *testing.T) {
})
user, org, model := seedAnthropicChatDependencies(t, db, anthropicURL)
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "search and read")
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
insertProviderToolPairMessageWithLocalTool(ctx, t, db, chat.ID, model.ID, "ws-disabled")
@@ -6953,7 +7051,9 @@ func TestActiveServer_AnthropicDropsUnpairedProviderToolBeforePersist(t *testing
user, org, model := seedAnthropicChatDependencies(t, db, anthropicURL)
model = enableAnthropicWebSearchForTest(t, db, model)
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "run provider tool")
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
@@ -6984,7 +7084,9 @@ func TestActiveServer_AnthropicKeepsPairedWebSearchBeforePersist(t *testing.T) {
user, org, model := seedAnthropicChatDependencies(t, db, anthropicURL)
model = enableAnthropicWebSearchForTest(t, db, model)
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "search for coder")
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
@@ -7037,7 +7139,9 @@ func TestActiveServer_AnthropicWebSearchFollowUpHasNoSyntheticCancellation(t *te
user, org, model := seedAnthropicChatDependencies(t, db, anthropicURL)
model = enableAnthropicWebSearchForTest(t, db, model)
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "search for coder")
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
@@ -7114,6 +7218,7 @@ func TestActiveServer_AnthropicSanitizesWebSearchBeforeContinuation(t *testing.T
Times(1)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, dbAgent.ID, agentID)
return mockConn, func() {}, nil
@@ -7186,6 +7291,7 @@ func TestActiveServer_ExclusiveToolPolicy(t *testing.T) {
mockConn.EXPECT().ReadFileLines(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).Times(0)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, dbAgent.ID, agentID)
return mockConn, func() {}, nil
@@ -7245,7 +7351,9 @@ func TestActiveServer_ExclusiveToolPolicy(t *testing.T) {
}})
require.NoError(t, err)
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
})
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
OwnerID: user.ID,
@@ -7294,7 +7402,9 @@ func TestActiveServer_ExclusiveToolPolicy(t *testing.T) {
})
user, org, model := seedChatDependenciesWithProvider(t, db, "openai-compat", openAIURL)
seedAdvisorConfig(ctx, t, db, codersdk.AdvisorConfig{Enabled: true, MaxUsesPerRun: 3, MaxOutputTokens: 1024})
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "advise only")
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
@@ -7337,7 +7447,9 @@ func TestActiveServer_ExclusiveToolPolicy(t *testing.T) {
},
})
seedAdvisorConfig(ctx, t, db, codersdk.AdvisorConfig{Enabled: true, MaxUsesPerRun: 3, MaxOutputTokens: 1024})
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "search then advise")
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
@@ -7378,7 +7490,9 @@ func TestActiveServer_ReasoningTimestamps(t *testing.T) {
},
})
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath()))
})
chat := createChatThroughServer(ctx, t, db, server, org.ID, user.ID, model.ID, "think")
waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting)
@@ -8179,6 +8293,7 @@ func TestActiveServer_GenerationErrorLogged(t *testing.T) {
})
user, org, model := seedChatDependenciesWithProvider(t, db, "openai", openAIURL)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
cfg.Logger = sink.Logger()
})
@@ -8289,6 +8404,7 @@ func TestProposeChatTitle_DebugRun(t *testing.T) {
ReplicaID: uuid.New(),
PendingChatAcquireInterval: testutil.WaitLong,
AlwaysEnableDebugLogs: tt.alwaysEnableDebugLogs,
AIBridgeTransportFactory: chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL)),
})
t.Cleanup(func() {
require.NoError(t, server.Close())
@@ -8660,6 +8776,7 @@ func TestInterruptChatDoesNotSendWebPushNotification(t *testing.T) {
PendingChatAcquireInterval: 10 * time.Millisecond,
InFlightChatStaleAfter: testutil.WaitSuperLong,
WebpushDispatcher: mockPush,
AIBridgeTransportFactory: chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL)),
})
t.Cleanup(func() {
require.NoError(t, server.Close())
@@ -8781,6 +8898,7 @@ func TestSuccessfulChatSendsWebPushWithNavigationData(t *testing.T) {
PendingChatAcquireInterval: 10 * time.Millisecond,
InFlightChatStaleAfter: testutil.WaitSuperLong,
WebpushDispatcher: mockPush,
AIBridgeTransportFactory: chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL)),
})
server.Start()
t.Cleanup(func() {
@@ -8861,12 +8979,14 @@ func TestCloseDuringShutdownContextCanceledShouldRetryOnNewReplica(t *testing.T)
})
loggerA := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
serverA := chatd.New(ps, chatd.Config{
Logger: loggerA,
Database: db,
ReplicaID: uuid.New(),
PendingChatAcquireInterval: 10 * time.Millisecond,
InFlightChatStaleAfter: testutil.WaitLong,
AIBridgeTransportFactory: chatAIGatewayTransportFactoryPointer(factory),
})
serverA.Start()
t.Cleanup(func() {
@@ -8912,6 +9032,7 @@ func TestCloseDuringShutdownContextCanceledShouldRetryOnNewReplica(t *testing.T)
ReplicaID: uuid.New(),
PendingChatAcquireInterval: 10 * time.Millisecond,
InFlightChatStaleAfter: testutil.WaitLong,
AIBridgeTransportFactory: chatAIGatewayTransportFactoryPointer(factory),
})
serverB.Start()
t.Cleanup(func() {
@@ -8966,6 +9087,7 @@ func TestSuccessfulChatSendsWebPushWithSummary(t *testing.T) {
PendingChatAcquireInterval: 10 * time.Millisecond,
InFlightChatStaleAfter: testutil.WaitSuperLong,
WebpushDispatcher: mockPush,
AIBridgeTransportFactory: chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL)),
})
server.Start()
t.Cleanup(func() {
@@ -9027,7 +9149,9 @@ func TestSuccessfulChatPersistsTurnSummaryWithoutWebPush(t *testing.T) {
)
})
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
})
user, org, model := seedChatDependencies(t, db)
setOpenAIProviderBaseURL(ctx, t, db, openAIURL)
@@ -9085,6 +9209,7 @@ func TestSuccessfulChatSendsWebPushFallbackWithoutSummaryForEmptyAssistantText(t
PendingChatAcquireInterval: 10 * time.Millisecond,
InFlightChatStaleAfter: testutil.WaitSuperLong,
WebpushDispatcher: mockPush,
AIBridgeTransportFactory: chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL)),
})
t.Cleanup(func() {
require.NoError(t, server.Close())
@@ -9145,6 +9270,7 @@ func TestErroredChatClearsLastTurnSummaryAndSendsWebPush(t *testing.T) {
PendingChatAcquireInterval: 10 * time.Millisecond,
InFlightChatStaleAfter: testutil.WaitSuperLong,
WebpushDispatcher: mockPush,
AIBridgeTransportFactory: chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL)),
})
t.Cleanup(func() {
require.NoError(t, server.Close())
@@ -9368,6 +9494,18 @@ func TestComputerUseSubagentToolsAndModel(t *testing.T) {
AnyTimes()
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
providers, providersErr := db.GetAIProviders(ctx, database.GetAIProvidersParams{})
require.NoError(t, providersErr)
routes := make(map[string]aibridge.TransportFactory, len(providers))
for _, provider := range providers {
switch provider.Type {
case database.AIProviderTypeOpenaiCompat:
routes[provider.Name] = chattest.NewMockAIBridgeTransport(t, openAIURL)
case database.AIProviderTypeAnthropic:
routes[provider.Name] = chattest.NewMockAIBridgeTransport(t, anthropicSrv.URL)
}
}
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(providerRoutedTransportFactory{routes: routes})
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, dbAgent.ID, agentID)
return mockConn, func() {}, nil
@@ -9532,6 +9670,7 @@ func TestInterruptChatPersistsPartialResponse(t *testing.T) {
PendingChatAcquireInterval: 10 * time.Millisecond,
InFlightChatStaleAfter: testutil.WaitSuperLong,
Experiments: codersdk.ExperimentsKnown,
AIBridgeTransportFactory: chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL)),
})
server.Start()
t.Cleanup(func() {
@@ -9658,7 +9797,9 @@ func TestProcessChat_UserProviderKey_Success(t *testing.T) {
})
require.NoError(t, err)
_ = newActiveTestServer(t, db, ps)
_ = newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
})
chatResult := waitForTerminalChat(ctx, t, db, chat.ID)
require.Equal(t, database.ChatStatusWaiting, chatResult.Status)
@@ -9742,7 +9883,6 @@ func TestProcessChat_RoutingUsesDelegatedAPIKey(t *testing.T) {
_ = newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
cfg.AIGatewayRoutingEnabled = true
cfg.AllowBYOK = true
cfg.AllowBYOKSet = true
})
@@ -9809,7 +9949,6 @@ func TestProcessChat_RoutingPreservesAPIKeyAfterWorkspaceContext(t *testing.T) {
"/home/coder/project/AGENTS.md", contextText)
_ = newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
cfg.AIGatewayRoutingEnabled = true
cfg.AllowBYOK = true
cfg.AllowBYOKSet = true
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
@@ -9932,7 +10071,9 @@ func TestProcessChatPanicRecovery(t *testing.T) {
// Pass the panic wrapper to the server, but use the real
// database for seeding so those operations don't panic.
server := newActiveTestServer(t, panicWrapper, ps)
server := newActiveTestServer(t, panicWrapper, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
})
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
@@ -10111,6 +10252,7 @@ func TestMCPServerToolInvocation(t *testing.T) {
Return(io.NopCloser(strings.NewReader("")), "", nil).AnyTimes()
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, dbAgent.ID, agentID)
return mockConn, func() {}, nil
@@ -10272,7 +10414,9 @@ func TestPlanModeRootChatApprovedExternalMCPToolInvocation(t *testing.T) {
UpdatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
})
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
})
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
@@ -10398,6 +10542,7 @@ func TestPlanModeRootChatApprovedExternalMCPWorkflowCanReachProposePlan(t *testi
}).AnyTimes()
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, dbAgent.ID, agentID)
return mockConn, func() {}, nil
@@ -10617,6 +10762,7 @@ func TestMCPServerOAuth2TokenRefresh(t *testing.T) {
mockConn.EXPECT().ReadFile(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).
Return(io.NopCloser(strings.NewReader("")), "", nil).AnyTimes()
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) {
require.Equal(t, dbAgent.ID, agentID)
return mockConn, func() {}, nil
@@ -10731,7 +10877,9 @@ func TestMCPServerOAuth2TokenRefreshFailureGraceful(t *testing.T) {
require.NoError(t, err)
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
})
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
@@ -10844,6 +10992,7 @@ func TestChatTemplateAllowlistEnforcement(t *testing.T) {
require.NoError(t, err)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
// Provide a CreateWorkspace function so the tool reaches
// the allowlist check instead of bailing with "not
// configured". If the allowlist is enforced correctly
@@ -11001,6 +11150,7 @@ func TestChatAsksUserWhenListTemplatesRequiresSelection(t *testing.T) {
})
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
cfg.CreateWorkspace = func(
context.Context,
uuid.UUID,
@@ -11131,6 +11281,7 @@ func TestCreateChatImmediatelyProcessesNewChat(t *testing.T) {
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.PendingChatAcquireInterval = time.Hour
cfg.InFlightChatStaleAfter = testutil.WaitSuperLong
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
})
user, org, model := seedChatDependencies(t, db)
@@ -11197,6 +11348,7 @@ func TestSendMessageImmediatelyProcessesWaitingChat(t *testing.T) {
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.PendingChatAcquireInterval = time.Hour
cfg.InFlightChatStaleAfter = testutil.WaitSuperLong
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, openAIURL))
})
user, org, model := seedChatDependencies(t, db)
@@ -11664,6 +11816,9 @@ func TestAdvisorGating_ExperimentDisabled(t *testing.T) {
)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.Experiments = experiments
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(
chattest.NewMockAIBridgeTransport(t, openAIURL),
)
})
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
@@ -11760,7 +11915,11 @@ func TestAdvisorGating_RootChat(t *testing.T) {
MaxUsesPerRun: 3,
MaxOutputTokens: 16384,
})
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(
chattest.NewMockAIBridgeTransport(t, openAIURL),
)
})
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
@@ -11935,7 +12094,11 @@ func TestAdvisorHappyPath_RootChat(t *testing.T) {
MaxUsesPerRun: 3,
MaxOutputTokens: 16384,
})
server := newTestServer(t, db, ps, uuid.New())
server := newTestServer(t, db, ps, uuid.New(), func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(
chattest.NewMockAIBridgeTransport(t, openAIURL),
)
})
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
@@ -12125,7 +12288,11 @@ func TestAdvisorGating_ChildChat(t *testing.T) {
Title: "advisor-root-parent",
})
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(
chattest.NewMockAIBridgeTransport(t, openAIURL),
)
})
childChat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
@@ -12204,7 +12371,11 @@ func TestAdvisorGating_PlanMode(t *testing.T) {
MaxUsesPerRun: 3,
MaxOutputTokens: 16384,
})
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(
chattest.NewMockAIBridgeTransport(t, openAIURL),
)
})
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
@@ -12284,7 +12455,11 @@ func TestAdvisorGating_ExploreSubagent(t *testing.T) {
MaxUsesPerRun: 3,
MaxOutputTokens: 16384,
})
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(
chattest.NewMockAIBridgeTransport(t, openAIURL),
)
})
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
@@ -12433,7 +12608,11 @@ func TestAdvisorChainMode_SnapshotKeepsFullHistory(t *testing.T) {
MaxUsesPerRun: 3,
MaxOutputTokens: 16384,
})
server := newOpenAIResponsesTestServer(t, db, ps)
server := newOpenAIResponsesTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(
chattest.NewMockAIBridgeTransport(t, openAIURL),
)
})
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
@@ -12577,7 +12756,18 @@ func TestProviderSwitchSanitizesAndRestoresPEToolHistory(t *testing.T) {
AIProviderID: uuid.NullUUID{UUID: cpB.ID, Valid: true},
})
server := newActiveTestServer(t, db, ps)
server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
provA, err := db.GetAIProviderByID(ctx, cpA.ID)
require.NoError(t, err)
provB, err := db.GetAIProviderByID(ctx, cpB.ID)
require.NoError(t, err)
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(providerRoutedTransportFactory{
routes: map[string]aibridge.TransportFactory{
provA.Name: chattest.NewMockAIBridgeTransport(t, serverAURL),
provB.Name: chattest.NewMockAIBridgeTransport(t, serverBURL),
},
})
})
// Given: an initial conversation turn with model A that produces provider-executed
// tool call results
@@ -12667,6 +12857,23 @@ func TestProviderSwitchSanitizesAndRestoresPEToolHistory(t *testing.T) {
"provider A must receive its own PE tool call when switching back")
}
// providerRoutedTransportFactory implements aibridge.TransportFactory by
// dispatching to a different mock transport per AI provider name. It
// exists for tests that exercise two providers in the same chat (e.g.
// switching models mid-conversation), where a single
// chattest.MockAIBridgeTransport's fixed target can't represent both.
type providerRoutedTransportFactory struct {
routes map[string]aibridge.TransportFactory
}
func (f providerRoutedTransportFactory) TransportFor(providerName string, source aibridge.Source) (http.RoundTripper, error) {
route, ok := f.routes[providerName]
if !ok {
return nil, xerrors.Errorf("no mock transport configured for provider %q", providerName)
}
return route.TransportFor(providerName, source)
}
func seedAdvisorConfig(
ctx context.Context,
t *testing.T,
+1 -1
View File
@@ -60,7 +60,7 @@ func (p *Server) computerUseProviderAndModelFromConfig(
func (p *Server) resolveComputerUseModel(
ctx context.Context,
chat database.Chat,
route resolvedModelRoute,
route aiGatewayModelRoute,
computerUseProvider string,
computerUseModelProvider string,
computerUseModelName string,
+1 -3
View File
@@ -19,7 +19,6 @@ import (
"github.com/coder/coder/v2/coderd/x/chatd/chaterror"
"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"
@@ -44,8 +43,7 @@ type generationPrepared struct {
Tools []fantasy.AgentTool
ActiveTools []string
ProviderTools []chatloop.ProviderTool
ProviderKeys chatprovider.ProviderAPIKeys
ModelRoute resolvedModelRoute
ModelRoute aiGatewayModelRoute
ModelBuildOptions modelBuildOptions
// ResolvedProvider is the configured provider identity used to label
+4 -12
View File
@@ -39,8 +39,7 @@ func (server *Server) prepareGeneration(
var (
model fantasy.LanguageModel
modelConfig database.ChatModelConfig
providerKeys chatprovider.ProviderAPIKeys
modelRoute resolvedModelRoute
modelRoute aiGatewayModelRoute
modelOpts modelBuildOptions
callConfig codersdk.ChatModelCallConfig
promptRows []database.ChatMessage
@@ -86,7 +85,7 @@ func (server *Server) prepareGeneration(
ctx = withActiveTurnAPIKeyID(ctx, modelOpts)
var err error
model, modelConfig, providerKeys, modelRoute, debugEnabled, resolvedProvider, debugModel, err = server.resolveChatModel(ctx, chat, modelOpts)
model, modelConfig, modelRoute, debugEnabled, resolvedProvider, debugModel, err = server.resolveChatModel(ctx, chat, modelOpts)
if err != nil {
return generationPrepared{}, err
}
@@ -130,7 +129,6 @@ func (server *Server) prepareGeneration(
advisorCfg,
model,
callConfig,
providerKeys,
modelOpts,
logger,
)
@@ -262,10 +260,7 @@ func (server *Server) prepareGeneration(
acceptsFilePart := func(mediaType string) bool {
return chatprovider.AcceptsFilePartMediaType(model.Provider(), model.Model(), mediaType)
}
providerType, err := modelRoute.providerHint()
if err != nil {
return xerrors.Errorf("resolve provider type: %w", err)
}
providerType := string(modelRoute.Provider.Type)
prompt, err = chatprompt.ConvertMessagesWithFiles(ctx, promptRows, server.chatFileResolver(providerType), logger, acceptsFilePart)
if err != nil {
return xerrors.Errorf("build chat prompt: %w", err)
@@ -497,7 +492,6 @@ func (server *Server) prepareGeneration(
return generationPrepared{}, xerrors.Errorf("resolve computer use provider route: %w", keyErr)
}
modelRoute = computerUseRoute
providerKeys = computerUseRoute.directProviderKeys()
cuModel, cuDebugEnabled, cuResolvedProvider, cuResolvedModel, cuErr := server.resolveComputerUseModel(
ctx,
chat,
@@ -620,7 +614,6 @@ func (server *Server) prepareGeneration(
Tools: tools,
ActiveTools: activeToolNames,
ProviderTools: providerTools,
ProviderKeys: providerKeys,
ModelRoute: modelRoute,
ModelBuildOptions: modelOpts,
ResolvedProvider: resolvedProvider,
@@ -748,7 +741,7 @@ func (server *Server) deriveFinalTurnRunResult(
// built from; they only feed the status-label fallback candidate's labels.
modelOpts := modelBuildOptionsFromMessages(promptRows)
ctx = withActiveTurnAPIKeyID(ctx, modelOpts)
model, _, providerKeys, modelRoute, _, resolvedProvider, resolvedModel, err := server.resolveChatModel(ctx, chat, modelOpts)
model, _, modelRoute, _, resolvedProvider, resolvedModel, err := server.resolveChatModel(ctx, chat, modelOpts)
if err != nil {
// Return what we have; generateFinalTurnStatusLabel falls back to a
// generic label when StatusLabelModel is nil.
@@ -763,7 +756,6 @@ func (server *Server) deriveFinalTurnRunResult(
return runChatResult{
FinalAssistantText: finalAssistantText,
StatusLabelModel: model,
ProviderKeys: providerKeys,
FallbackProvider: resolvedProvider,
FallbackRoute: modelRoute,
FallbackModel: resolvedModel,
@@ -142,7 +142,10 @@ func TestDeriveFinalTurnRunResult(t *testing.T) {
})
require.NoError(t, err)
server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{})
server := newInternalTestServer(
t, db, ps, chatprovider.ProviderAPIKeys{},
withInternalTestServerTransportFactory(&aibridgeTestFactory{}),
)
return server, created.Chat
}
@@ -192,7 +195,6 @@ func TestDeriveFinalTurnRunResult(t *testing.T) {
require.NotNil(t, result.StatusLabelModel)
require.Equal(t, "openai", result.FallbackProvider)
require.Equal(t, "gpt-4o-mini", result.FallbackModel)
require.False(t, result.ProviderKeys.Empty())
})
t.Run("NonWaitingReturnsEmpty", func(t *testing.T) {
+20 -6
View File
@@ -67,7 +67,10 @@ func TestOpenAIResponsesNoStaleWebSearchReplay(t *testing.T) {
user, org, _ := seedChatDependenciesWithProvider(t, db, "openai", openAIURL)
model := insertOpenAIResponsesModelConfig(t, db, user.ID, false, true)
server := newOpenAIResponsesTestServer(t, db, ps)
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newOpenAIResponsesTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
@@ -152,7 +155,10 @@ func TestOpenAIResponsesFullReplayPairsReasoningAndWebSearch(t *testing.T) {
user, org, _ := seedChatDependenciesWithProvider(t, db, "openai", openAIURL)
firstModel := insertOpenAIResponsesModelConfig(t, db, user.ID, true, true)
secondModel := insertOpenAIResponsesModelConfig(t, db, user.ID, true, true)
server := newOpenAIResponsesTestServer(t, db, ps)
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newOpenAIResponsesTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
OrganizationID: org.ID,
@@ -235,7 +241,10 @@ func TestOpenAIResponsesChainModeSkipsWhenLocalCallPending(t *testing.T) {
},
)
server := newOpenAIResponsesTestServer(t, db, ps)
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newOpenAIResponsesTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
_, err := server.SendMessage(ctx, chatd.SendMessageOptions{
ChatID: chat.ID,
CreatedBy: user.ID,
@@ -318,7 +327,10 @@ func TestOpenAIResponsesChainModeStillFiresForProviderExecutedOnly(t *testing.T)
},
)
server := newOpenAIResponsesTestServer(t, db, ps)
factory := chattest.NewMockAIBridgeTransport(t, openAIURL)
server := newOpenAIResponsesTestServer(t, db, ps, func(cfg *chatd.Config) {
cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(factory)
})
_, err := server.SendMessage(ctx, chatd.SendMessageOptions{
ChatID: chat.ID,
CreatedBy: user.ID,
@@ -394,16 +406,18 @@ func newOpenAIResponsesTestServer(
t *testing.T,
db database.Store,
ps dbpubsub.Pubsub,
overrides ...func(*chatd.Config),
) *chatd.Server {
t.Helper()
return newActiveTestServer(t, db, ps, func(cfg *chatd.Config) {
allOverrides := append([]func(*chatd.Config){func(cfg *chatd.Config) {
// Let CreateChat and SendMessage publish their pending status
// before wake-driven processing starts. The responses tests are
// not exercising periodic polling, and PostgreSQL can otherwise
// deliver that stale pending notification after processChat
// subscribes to control events.
cfg.PendingChatAcquireInterval = testutil.WaitLong
})
}}, overrides...)
return newActiveTestServer(t, db, ps, allOverrides...)
}
func insertOpenAIResponsesModelConfig(
+6 -75
View File
@@ -41,57 +41,6 @@ func withActiveTurnAPIKeyID(ctx context.Context, opts modelBuildOptions) context
return aibridge.WithDelegatedAPIKeyID(ctx, opts.ActiveAPIKeyID)
}
type modelRouteKind int
const (
modelRouteKindDirect modelRouteKind = iota + 1
modelRouteKindAIGateway
)
type resolvedModelRoute struct {
kind modelRouteKind
direct directModelRoute
aiGateway aiGatewayModelRoute
}
func newDirectModelRoute(providerHint string, keys chatprovider.ProviderAPIKeys) resolvedModelRoute {
return resolvedModelRoute{
kind: modelRouteKindDirect,
direct: directModelRoute{
ProviderHint: providerHint,
Keys: keys,
},
}
}
func (r resolvedModelRoute) providerHint() (string, error) {
switch r.kind {
case modelRouteKindDirect:
return r.direct.ProviderHint, nil
case modelRouteKindAIGateway:
return r.aiGateway.ModelProviderHint, nil
default:
return "", xerrors.New("model route is not configured")
}
}
func (r resolvedModelRoute) withProviderHint(providerHint string) resolvedModelRoute {
switch r.kind {
case modelRouteKindDirect:
r.direct.ProviderHint = providerHint
case modelRouteKindAIGateway:
r.aiGateway.ModelProviderHint = providerHint
}
return r
}
func (r resolvedModelRoute) directProviderKeys() chatprovider.ProviderAPIKeys {
if r.kind != modelRouteKindDirect {
return chatprovider.ProviderAPIKeys{}
}
return r.direct.Keys
}
func (p *Server) enabledAIProviderByID(ctx context.Context, providerID uuid.UUID) (database.AIProvider, error) {
provider, err := p.db.GetAIProviderByID(ctx, providerID)
if err != nil {
@@ -103,47 +52,29 @@ func (p *Server) enabledAIProviderByID(ctx context.Context, providerID uuid.UUID
return provider, nil
}
func (p *Server) shouldUseAIGatewayRouting() bool {
return p.aiGatewayRoutingEnabled
}
func (p *Server) resolveModelRouteForConfig(
ctx context.Context,
ownerID uuid.UUID,
modelConfig database.ChatModelConfig,
fallbackKeys chatprovider.ProviderAPIKeys,
) (resolvedModelRoute, error) {
if p.shouldUseAIGatewayRouting() {
return p.resolveAIGatewayModelRouteForConfig(ctx, ownerID, modelConfig)
}
return p.resolveDirectModelRouteForConfig(ctx, ownerID, modelConfig, fallbackKeys)
) (aiGatewayModelRoute, error) {
return p.resolveAIGatewayModelRouteForConfig(ctx, ownerID, modelConfig)
}
func (p *Server) resolveModelRouteForProviderType(
ctx context.Context,
ownerID uuid.UUID,
providerType string,
) (resolvedModelRoute, error) {
if p.shouldUseAIGatewayRouting() {
return p.resolveAIGatewayModelRouteForProviderType(ctx, ownerID, providerType)
}
return p.resolveDirectModelRouteForProviderType(ctx, ownerID, providerType)
) (aiGatewayModelRoute, error) {
return p.resolveAIGatewayModelRouteForProviderType(ctx, ownerID, providerType)
}
func (p *Server) newModel(
ctx context.Context,
req modelClientRequest,
route resolvedModelRoute,
route aiGatewayModelRoute,
opts modelBuildOptions,
) (fantasy.LanguageModel, error) {
switch route.kind {
case modelRouteKindDirect:
return p.newDirectModel(ctx, req, route.direct, opts)
case modelRouteKindAIGateway:
return p.newAIGatewayModel(ctx, req, route.aiGateway, opts)
default:
return nil, xerrors.New("model route is not configured")
}
return p.newAIGatewayModel(ctx, req, route, opts)
}
func newLanguageModel(
+11 -14
View File
@@ -39,14 +39,11 @@ func newAIGatewayModelRoute(
provider database.AIProvider,
modelProviderHint string,
auth aiGatewayProviderAuth,
) resolvedModelRoute {
return resolvedModelRoute{
kind: modelRouteKindAIGateway,
aiGateway: aiGatewayModelRoute{
Provider: provider,
ModelProviderHint: modelProviderHint,
ProviderAuth: auth,
},
) aiGatewayModelRoute {
return aiGatewayModelRoute{
Provider: provider,
ModelProviderHint: modelProviderHint,
ProviderAuth: auth,
}
}
@@ -257,7 +254,7 @@ func (p *Server) resolveAIGatewayRoute(
ownerID uuid.UUID,
provider database.AIProvider,
modelProviderHint string,
) (resolvedModelRoute, error) {
) (aiGatewayModelRoute, error) {
auth, err := p.aiGatewayProviderAuthForUser(
ctx,
ownerID,
@@ -265,7 +262,7 @@ func (p *Server) resolveAIGatewayRoute(
aiGatewayRequestFormatForProviderType(provider.Type),
)
if err != nil {
return resolvedModelRoute{}, xerrors.Errorf("resolve AI Gateway provider auth: %w", err)
return aiGatewayModelRoute{}, xerrors.Errorf("resolve AI Gateway provider auth: %w", err)
}
return newAIGatewayModelRoute(provider, modelProviderHint, auth), nil
}
@@ -274,10 +271,10 @@ func (p *Server) resolveAIGatewayModelRouteForConfig(
ctx context.Context,
ownerID uuid.UUID,
modelConfig database.ChatModelConfig,
) (resolvedModelRoute, error) {
) (aiGatewayModelRoute, error) {
provider, err := p.gatewayProviderForConfig(ctx, modelConfig)
if err != nil {
return resolvedModelRoute{}, err
return aiGatewayModelRoute{}, err
}
return p.resolveAIGatewayRoute(ctx, ownerID, provider, string(provider.Type))
}
@@ -286,10 +283,10 @@ func (p *Server) resolveAIGatewayModelRouteForProviderType(
ctx context.Context,
ownerID uuid.UUID,
providerType string,
) (resolvedModelRoute, error) {
) (aiGatewayModelRoute, error) {
provider, err := p.aiProviderForProviderType(ctx, providerType)
if err != nil {
return resolvedModelRoute{}, err
return aiGatewayModelRoute{}, err
}
return p.resolveAIGatewayRoute(
ctx,
-94
View File
@@ -1,94 +0,0 @@
package chatd
import (
"context"
"database/sql"
"net/http"
"charm.land/fantasy"
"github.com/google/uuid"
"golang.org/x/xerrors"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/x/chatd/chatdebug"
"github.com/coder/coder/v2/coderd/x/chatd/chatprovider"
)
type directModelRoute struct {
ProviderHint string
Keys chatprovider.ProviderAPIKeys
}
func (*Server) newDirectModel(
_ context.Context,
req modelClientRequest,
route directModelRoute,
opts modelBuildOptions,
) (fantasy.LanguageModel, error) {
var httpClient *http.Client
if opts.RecordHTTP {
httpClient = &http.Client{Transport: &chatdebug.RecordingTransport{}}
}
return newLanguageModel(
route.ProviderHint,
req.ModelName,
route.Keys,
req.UserAgent,
req.ExtraHeaders,
httpClient,
)
}
func (p *Server) resolveDirectModelRouteForConfig(
ctx context.Context,
ownerID uuid.UUID,
modelConfig database.ChatModelConfig,
fallbackKeys chatprovider.ProviderAPIKeys,
) (resolvedModelRoute, error) {
providerHint, provider, err := p.directProviderHintAndProviderForConfig(ctx, modelConfig)
if err != nil {
return resolvedModelRoute{}, err
}
if provider == nil {
if !fallbackKeys.Empty() && userCanUseProviderKeys(fallbackKeys, providerHint) {
return newDirectModelRoute(providerHint, fallbackKeys), nil
}
keys, err := p.resolveUserProviderAPIKeys(ctx, ownerID, uuid.Nil)
if err != nil {
return resolvedModelRoute{}, xerrors.Errorf("resolve provider API keys: %w", err)
}
return newDirectModelRoute(providerHint, keys), nil
}
providerKeys, err := p.resolveUserProviderAPIKeysForProvider(ctx, ownerID, *provider)
if err != nil {
return resolvedModelRoute{}, xerrors.Errorf("resolve provider API keys: %w", err)
}
return newDirectModelRoute(providerHint, providerKeys), nil
}
func (p *Server) resolveDirectModelRouteForProviderType(
ctx context.Context,
ownerID uuid.UUID,
providerType string,
) (resolvedModelRoute, error) {
normalizedProviderType := chatprovider.NormalizeProvider(providerType)
keys, _, err := p.resolveUserProviderAPIKeysAndProviderForProviderType(ctx, ownerID, providerType)
if err != nil {
return resolvedModelRoute{}, err
}
return newDirectModelRoute(normalizedProviderType, keys), nil
}
func (p *Server) directProviderHintAndProviderForConfig(
ctx context.Context,
modelConfig database.ChatModelConfig,
) (string, *database.AIProvider, error) {
if !modelConfig.AIProviderID.Valid {
return "", nil, sql.ErrNoRows
}
provider, err := p.enabledAIProviderByID(ctx, modelConfig.AIProviderID.UUID)
if err != nil {
return "", nil, err
}
return string(provider.Type), &provider, nil
}
+15 -30
View File
@@ -64,7 +64,7 @@ func aibridgeTestAIProvider(providerID uuid.UUID, providerName string, providerT
}
}
func aibridgeTestRoute(aiProvider database.AIProvider) resolvedModelRoute {
func aibridgeTestRoute(aiProvider database.AIProvider) aiGatewayModelRoute {
return newAIGatewayModelRoute(aiProvider, string(aiProvider.Type), aiGatewayProviderAuth{})
}
@@ -118,20 +118,15 @@ func TestResolveModelRouteForConfigPreservesBaseURL(t *testing.T) {
Enabled: true,
BaseUrl: baseURL,
}, nil)
db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{
ProviderID: providerID,
APIKey: "provider-key",
}}, nil)
server := &Server{db: db}
route, err := server.resolveModelRouteForConfig(ctx, ownerID, database.ChatModelConfig{
AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true},
}, chatprovider.ProviderAPIKeys{})
})
require.NoError(t, err)
require.Equal(t, modelRouteKindDirect, route.kind)
require.Equal(t, "openai", route.direct.ProviderHint)
require.Equal(t, "provider-key", route.direct.Keys.APIKey("openai"))
require.Equal(t, baseURL, route.direct.Keys.BaseURL("openai"))
require.Equal(t, "openai", route.ModelProviderHint)
require.Equal(t, providerID, route.Provider.ID)
require.Equal(t, baseURL, route.Provider.BaseUrl)
}
func TestAIGatewayProviderAuthForUser(t *testing.T) {
@@ -251,11 +246,10 @@ func TestResolveModelRouteForConfigAIGatewayProviderAuth(t *testing.T) {
AIProviderID: providerID,
}).Return(database.UserAIProviderKey{APIKey: "sk-user"}, nil)
server := &Server{db: db, aiGatewayRoutingEnabled: true, allowBYOK: true}
route, err := server.resolveModelRouteForConfig(ctx, ownerID, modelConfig, chatprovider.ProviderAPIKeys{})
server := &Server{db: db, allowBYOK: true}
route, err := server.resolveModelRouteForConfig(ctx, ownerID, modelConfig)
require.NoError(t, err)
require.Equal(t, modelRouteKindAIGateway, route.kind)
require.Equal(t, "Bearer sk-user", route.aiGateway.ProviderAuth.Headers["Authorization"])
require.Equal(t, "Bearer sk-user", route.ProviderAuth.Headers["Authorization"])
})
t.Run("CentralProviderCredentialsNotForwarded", func(t *testing.T) {
@@ -265,11 +259,10 @@ func TestResolveModelRouteForConfigAIGatewayProviderAuth(t *testing.T) {
db := dbmock.NewMockStore(ctrl)
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil)
server := &Server{db: db, aiGatewayRoutingEnabled: true, allowBYOK: false}
route, err := server.resolveModelRouteForConfig(ctx, ownerID, modelConfig, chatprovider.ProviderAPIKeys{})
server := &Server{db: db, allowBYOK: false}
route, err := server.resolveModelRouteForConfig(ctx, ownerID, modelConfig)
require.NoError(t, err)
require.Equal(t, modelRouteKindAIGateway, route.kind)
require.Empty(t, route.aiGateway.ProviderAuth.Headers)
require.Empty(t, route.ProviderAuth.Headers)
})
}
@@ -283,7 +276,7 @@ func TestAIGatewayModelForwardsProviderAuth(t *testing.T) {
apiKeyID string
path string
}
newServer := func(t *testing.T, provider database.AIProvider, auth aiGatewayProviderAuth, seen chan seenRequest) (*Server, resolvedModelRoute) {
newServer := func(t *testing.T, provider database.AIProvider, auth aiGatewayProviderAuth, seen chan seenRequest) (*Server, aiGatewayModelRoute) {
factory := &aibridgeTestFactory{rt: roundTripFunc(func(req *http.Request) (*http.Response, error) {
apiKeyID, _ := aibridge.DelegatedAPIKeyIDFromContext(req.Context())
seen <- seenRequest{
@@ -305,7 +298,6 @@ func TestAIGatewayModelForwardsProviderAuth(t *testing.T) {
}, nil
})}
server := &Server{
aiGatewayRoutingEnabled: true,
aibridgeTransportFactory: aibridgeTestFactoryPointer(factory),
}
route := newAIGatewayModelRoute(provider, string(provider.Type), auth)
@@ -597,7 +589,7 @@ func TestAIBridgeRoutingFailClosed(t *testing.T) {
t.Run("NilFactory", func(t *testing.T) {
t.Parallel()
server := &Server{aiGatewayRoutingEnabled: true}
server := &Server{}
_, err := server.newModel(t.Context(), aibridgeTestRequest(chat, "gpt-4"), aibridgeTestRoute(aiProvider), modelBuildOptions{ActiveAPIKeyID: uuid.NewString()})
require.ErrorContains(t, err, "transport factory")
})
@@ -606,7 +598,6 @@ func TestAIBridgeRoutingFailClosed(t *testing.T) {
t.Parallel()
factory := &aibridgeTestFactory{err: xerrors.New("boom")}
server := &Server{
aiGatewayRoutingEnabled: true,
aibridgeTransportFactory: aibridgeTestFactoryPointer(factory),
}
_, err := server.newModel(t.Context(), aibridgeTestRequest(chat, "gpt-4"), aibridgeTestRoute(aiProvider), modelBuildOptions{ActiveAPIKeyID: uuid.NewString()})
@@ -615,7 +606,7 @@ func TestAIBridgeRoutingFailClosed(t *testing.T) {
t.Run("MissingProviderName", func(t *testing.T) {
t.Parallel()
server := &Server{aiGatewayRoutingEnabled: true}
server := &Server{}
missingNameProvider := aibridgeTestAIProvider(providerID, "", database.AIProviderTypeOpenai)
_, err := server.newModel(t.Context(), aibridgeTestRequest(chat, "gpt-4"), aibridgeTestRoute(missingNameProvider), modelBuildOptions{ActiveAPIKeyID: uuid.NewString()})
require.ErrorContains(t, err, "AI provider name")
@@ -628,7 +619,6 @@ func TestAIBridgeRoutingFailClosed(t *testing.T) {
return nil, xerrors.New("unreachable")
})}
server := &Server{
aiGatewayRoutingEnabled: true,
aibridgeTransportFactory: aibridgeTestFactoryPointer(factory),
}
_, err := server.newModel(t.Context(), aibridgeTestRequest(chat, "gpt-4"), aibridgeTestRoute(aiProvider), modelBuildOptions{})
@@ -647,7 +637,6 @@ func TestAIBridgeRoutingFailClosed(t *testing.T) {
return nil, xerrors.New("unreachable")
})}
server := &Server{
aiGatewayRoutingEnabled: true,
aibridgeTransportFactory: aibridgeTestFactoryPointer(factory),
}
provider := aibridgeTestAIProvider(providerID, "openrouter", database.AIProviderTypeOpenai)
@@ -665,7 +654,7 @@ func TestAIBridgeRoutingFailClosed(t *testing.T) {
t.Run("StaticModel", func(t *testing.T) {
t.Parallel()
server := &Server{aiGatewayRoutingEnabled: true}
server := &Server{}
_, err := server.newModel(t.Context(), aibridgeTestRequest(chat, "gpt-4"), newAIGatewayModelRoute(database.AIProvider{}, "", aiGatewayProviderAuth{}), modelBuildOptions{ActiveAPIKeyID: uuid.NewString()})
require.ErrorContains(t, err, "concrete AI provider")
})
@@ -751,7 +740,6 @@ func TestAIBridgeGatewayProviderTypesPreserveSlashModelID(t *testing.T) {
})}
chat := database.Chat{ID: uuid.New(), OwnerID: uuid.New()}
server := &Server{
aiGatewayRoutingEnabled: true,
aibridgeTransportFactory: aibridgeTestFactoryPointer(factory),
}
@@ -788,7 +776,6 @@ func TestAIBridgeComputerUseModelUsesRoute(t *testing.T) {
})}
chat := database.Chat{ID: uuid.New(), OwnerID: uuid.New()}
server := &Server{
aiGatewayRoutingEnabled: true,
aibridgeTransportFactory: aibridgeTestFactoryPointer(factory),
}
provider := chattool.ComputerUseProviderOpenAI
@@ -824,7 +811,6 @@ func TestResolveComputerUseModel_AIGatewayMissingAPIKeyID(t *testing.T) {
})}
chat := database.Chat{ID: uuid.New(), OwnerID: uuid.New()}
server := &Server{
aiGatewayRoutingEnabled: true,
aibridgeTransportFactory: aibridgeTestFactoryPointer(factory),
}
provider := chattool.ComputerUseProviderOpenAI
@@ -877,7 +863,6 @@ func TestAIBridgeDelegatedContextPropagation(t *testing.T) {
})}
chat := database.Chat{ID: uuid.New(), OwnerID: uuid.New()}
server := &Server{
aiGatewayRoutingEnabled: true,
aibridgeTransportFactory: aibridgeTestFactoryPointer(factory),
}
+93 -169
View File
@@ -68,39 +68,10 @@ var preferredTitleModels = []struct {
type shortTextCandidate struct {
provider string
model string
route resolvedModelRoute
route aiGatewayModelRoute
lm fantasy.LanguageModel
}
func (p *Server) preferredShortTextCandidates(
chat database.Chat,
keys chatprovider.ProviderAPIKeys,
) []shortTextCandidate {
if p.shouldUseAIGatewayRouting() {
return nil
}
candidates := make([]shortTextCandidate, 0, len(preferredTitleModels)+1)
userAgent := chatprovider.UserAgent()
extraHeaders := chatprovider.CoderHeaders(chat)
for _, candidate := range preferredTitleModels {
model, err := chatprovider.ModelFromConfig(
candidate.provider, candidate.model, keys, userAgent,
extraHeaders,
nil,
)
if err == nil {
candidates = append(candidates, shortTextCandidate{
provider: candidate.provider,
model: candidate.model,
route: newDirectModelRoute(candidate.provider, keys),
lm: model,
})
}
}
return candidates
}
func selectPreferredConfiguredShortTextModelConfig(
configs []database.GetEnabledChatModelConfigsRow,
) (database.ChatModelConfig, bool) {
@@ -174,29 +145,21 @@ func (p *Server) GenerateChatTitleAsync(ctx context.Context, chat database.Chat)
defer stopTitleCtx()
modelOpts := modelBuildOptionsFromMessages(messages)
turnCtx := withActiveTurnAPIKeyID(titleCtx, modelOpts)
model, modelConfig, keys, route, _, _, _, err := p.resolveChatModel(turnCtx, chat, modelOpts)
model, modelConfig, route, _, _, _, err := p.resolveChatModel(turnCtx, chat, modelOpts)
if err != nil {
logger.Debug(turnCtx, "failed to resolve model for automatic title generation",
slog.Error(err),
)
return
}
providerType, err := route.providerHint()
if err != nil {
logger.Debug(titleCtx, "failed to resolve provider type for automatic title generation",
slog.Error(err),
)
return
}
p.maybeGenerateChatTitle(
turnCtx,
chat,
messages,
providerType,
string(route.Provider.Type),
modelConfig.Model,
model,
route,
keys,
modelOpts,
&generatedChatTitle{},
logger,
@@ -226,8 +189,7 @@ func (p *Server) maybeGenerateChatTitle(
fallbackProvider string,
fallbackModelName string,
fallbackModel fantasy.LanguageModel,
fallbackRoute resolvedModelRoute,
keys chatprovider.ProviderAPIKeys,
fallbackRoute aiGatewayModelRoute,
modelOpts modelBuildOptions,
generatedTitle *generatedChatTitle,
logger slog.Logger,
@@ -242,10 +204,9 @@ func (p *Server) maybeGenerateChatTitle(
titleCtx, cancel := context.WithTimeout(ctx, 30*time.Second)
defer cancel()
overrideConfig, overrideModel, _, overrideRoute, overrideSet, overrideErr := p.resolveTitleGenerationModelOverride(
overrideConfig, overrideModel, overrideRoute, overrideSet, overrideErr := p.resolveTitleGenerationModelOverride(
titleCtx,
chat,
keys,
modelOpts,
)
if overrideErr != nil {
@@ -264,30 +225,21 @@ func (p *Server) maybeGenerateChatTitle(
)
}
var candidates []shortTextCandidate
var candidate shortTextCandidate
if overrideSet {
overrideProvider, err := overrideRoute.providerHint()
if err != nil {
logger.Debug(ctx, "failed to resolve provider type for title generation override",
slog.F("chat_id", chat.ID),
slog.Error(err),
)
return
}
candidates = []shortTextCandidate{{
provider: overrideProvider,
candidate = shortTextCandidate{
provider: string(overrideRoute.Provider.Type),
model: overrideConfig.Model,
route: overrideRoute,
lm: overrideModel,
}}
}
} else {
candidates = p.preferredShortTextCandidates(chat, keys)
candidates = append(candidates, shortTextCandidate{
candidate = shortTextCandidate{
provider: fallbackProvider,
model: fallbackModelName,
route: fallbackRoute,
lm: fallbackModel,
})
}
}
var historyTipMessageID int64
@@ -310,81 +262,61 @@ func (p *Server) maybeGenerateChatTitle(
chatdebug.TruncateLabel(input, chatdebug.MaxLabelLength),
)
var lastErr error
for _, candidate := range candidates {
candidateCtx := titleCtx
candidateModel := candidate.lm
finishDebugRun := func(error) {}
if debugEnabled {
candidateCtx, candidateModel, finishDebugRun = p.prepareQuickgenDebugCandidate(
titleCtx,
chat,
debugSvc,
candidate,
modelOpts,
chatdebug.KindTitleGeneration,
triggerMessageID,
historyTipMessageID,
seedSummary,
logger,
candidateCtx := titleCtx
candidateModel := candidate.lm
finishDebugRun := func(error) {}
if debugEnabled {
candidateCtx, candidateModel, finishDebugRun = p.prepareQuickgenDebugCandidate(
titleCtx,
chat,
debugSvc,
candidate,
modelOpts,
chatdebug.KindTitleGeneration,
triggerMessageID,
historyTipMessageID,
seedSummary,
logger,
)
}
title, err := generateTitle(candidateCtx, candidateModel, input)
finishDebugRun(err)
if err != nil {
if overrideSet {
logger.Warn(ctx, "title model candidate failed",
slog.F("chat_id", chat.ID),
slog.F("override_context", titleGenerationOverrideContext),
slog.F("provider", candidate.provider),
slog.F("model", candidate.model),
slog.Error(err),
)
}
title, err := generateTitle(candidateCtx, candidateModel, input)
finishDebugRun(err)
if err != nil {
lastErr = err
if overrideSet {
logger.Warn(ctx, "title model candidate failed",
slog.F("chat_id", chat.ID),
slog.F("override_context", titleGenerationOverrideContext),
slog.F("provider", candidate.provider),
slog.F("model", candidate.model),
slog.Error(err),
)
} else {
logger.Debug(ctx, "title model candidate failed",
slog.F("chat_id", chat.ID),
slog.Error(err),
)
}
continue
}
if title == "" || title == chat.Title {
return
}
_, err = p.db.UpdateChatTitleByID(ctx, database.UpdateChatTitleByIDParams{
ID: chat.ID,
Title: title,
})
if err != nil {
logger.Warn(ctx, "failed to update generated chat title",
} else {
logger.Debug(ctx, "title model candidate failed",
slog.F("chat_id", chat.ID),
slog.Error(err),
)
return
}
chat.Title = title
generatedTitle.Store(title)
p.publishChatPubsubEvent(chat, codersdk.ChatWatchEventKindTitleChange, nil)
return
}
if title == "" || title == chat.Title {
return
}
if lastErr != nil {
if overrideSet {
logger.Warn(ctx, "all title model candidates failed",
slog.F("chat_id", chat.ID),
slog.F("override_context", titleGenerationOverrideContext),
slog.Error(lastErr),
)
} else {
logger.Debug(ctx, "all title model candidates failed",
slog.F("chat_id", chat.ID),
slog.Error(lastErr),
)
}
_, err = p.db.UpdateChatTitleByID(ctx, database.UpdateChatTitleByIDParams{
ID: chat.ID,
Title: title,
})
if err != nil {
logger.Warn(ctx, "failed to update generated chat title",
slog.F("chat_id", chat.ID),
slog.Error(err),
)
return
}
chat.Title = title
generatedTitle.Store(title)
p.publishChatPubsubEvent(chat, codersdk.ChatWatchEventKindTitleChange, nil)
}
func (p *Server) newQuickgenDebugModel(
@@ -393,7 +325,7 @@ func (p *Server) newQuickgenDebugModel(
debugSvc *chatdebug.Service,
provider string,
model string,
route resolvedModelRoute,
route aiGatewayModelRoute,
modelOpts modelBuildOptions,
) (fantasy.LanguageModel, error) {
debugOpts := modelOpts
@@ -899,11 +831,8 @@ const turnStatusLabelPrompt = "You write compact chat status labels for a sideba
"Prefer short action or state phrases such as Finished, Submitted, Fixed, Testing, Still working, or Waiting for. " +
"No quotes, emoji, markdown, or trailing punctuation."
// generateTurnStatusLabel calls a cheap model to produce a short status
// label from the chat title, current state, and last assistant
// message text. It follows the same candidate-selection strategy
// as title generation: try preferred lightweight models first, then
// fall back to the provided model. Returns "" on any failure.
// generateTurnStatusLabel produces a short turn status label using the
// caller-supplied fallback model. Returns "" on any failure.
func (p *Server) generateTurnStatusLabel(
ctx context.Context,
chat database.Chat,
@@ -912,8 +841,7 @@ func (p *Server) generateTurnStatusLabel(
fallbackProvider string,
fallbackModelName string,
fallbackModel fantasy.LanguageModel,
fallbackRoute resolvedModelRoute,
keys chatprovider.ProviderAPIKeys,
fallbackRoute aiGatewayModelRoute,
modelOpts modelBuildOptions,
logger slog.Logger,
debugSvc *chatdebug.Service,
@@ -930,51 +858,47 @@ func (p *Server) generateTurnStatusLabel(
"\nChat title: " + chat.Title +
"\n\nAgent's latest message:\n" + assistantText
candidates := p.preferredShortTextCandidates(chat, keys)
candidates = append(candidates, shortTextCandidate{
candidate := shortTextCandidate{
provider: fallbackProvider,
model: fallbackModelName,
route: fallbackRoute,
lm: fallbackModel,
})
}
statusSeedSummary := chatdebug.SeedSummary("Turn status label")
for _, candidate := range candidates {
candidateCtx := labelCtx
candidateModel := candidate.lm
finishDebugRun := func(error) {}
if debugEnabled {
candidateCtx, candidateModel, finishDebugRun = p.prepareQuickgenDebugCandidate(
labelCtx,
chat,
debugSvc,
candidate,
modelOpts,
chatdebug.KindQuickgen,
triggerMessageID,
historyTipMessageID,
statusSeedSummary,
logger,
)
}
generatedLabel, err := generateStructuredTurnStatusLabel(
candidateCtx,
candidateModel,
turnStatusLabelPrompt,
input,
candidateCtx := labelCtx
candidateModel := candidate.lm
finishDebugRun := func(error) {}
if debugEnabled {
candidateCtx, candidateModel, finishDebugRun = p.prepareQuickgenDebugCandidate(
labelCtx,
chat,
debugSvc,
candidate,
modelOpts,
chatdebug.KindQuickgen,
triggerMessageID,
historyTipMessageID,
statusSeedSummary,
logger,
)
finishDebugRun(err)
if err != nil {
logger.Debug(ctx, "turn status label model candidate failed",
slog.Error(err),
)
continue
}
return generatedLabel
}
return ""
generatedLabel, err := generateStructuredTurnStatusLabel(
candidateCtx,
candidateModel,
turnStatusLabelPrompt,
input,
)
finishDebugRun(err)
if err != nil {
logger.Debug(ctx, "turn status label model candidate failed",
slog.Error(err),
)
return ""
}
return generatedLabel
}
func generateStructuredTurnStatusLabel(
+1 -12
View File
@@ -362,16 +362,6 @@ func Test_renderManualTitlePrompt(t *testing.T) {
}
}
func TestPreferredShortTextCandidatesNilUnderAIGateway(t *testing.T) {
t.Parallel()
server := &Server{aiGatewayRoutingEnabled: true}
candidates := server.preferredShortTextCandidates(database.Chat{}, chatprovider.ProviderAPIKeys{
ByProvider: map[string]string{"openai": "test-key"},
})
require.Nil(t, candidates)
}
func TestMaybeGenerateChatTitlePreservesUpdatedAt(t *testing.T) {
t.Parallel()
@@ -440,8 +430,7 @@ func TestMaybeGenerateChatTitlePreservesUpdatedAt(t *testing.T) {
"openai",
"test-model",
model,
resolvedModelRoute{},
chatprovider.ProviderAPIKeys{},
aiGatewayModelRoute{},
modelBuildOptions{},
generated,
logger,
+59 -18
View File
@@ -6,8 +6,10 @@ import (
"encoding/json"
"net/http"
"net/http/httptest"
"net/url"
"slices"
"sync"
"sync/atomic"
"testing"
"time"
@@ -71,10 +73,11 @@ func TestSubagentFallbackChatTitle(t *testing.T) {
}
type internalTestServerConfig struct {
logger slog.Logger
clock quartz.Clock
startWorker bool
experiments codersdk.Experiments
logger slog.Logger
clock quartz.Clock
startWorker bool
experiments codersdk.Experiments
transportFactory *atomic.Pointer[aibridge.TransportFactory]
}
type internalTestServerOpt func(*internalTestServerConfig)
@@ -103,6 +106,16 @@ func withInternalTestServerExperiments(experiments codersdk.Experiments) interna
}
}
// withInternalTestServerTransportFactory wires an [aibridge.TransportFactory]
// into the server's Config so tests that drive real model generation through
// runSubagentTool or processChat can control the HTTP transport AI Gateway
// routing uses.
func withInternalTestServerTransportFactory(factory aibridge.TransportFactory) internalTestServerOpt {
return func(cfg *internalTestServerConfig) {
cfg.transportFactory = aibridgeTestFactoryPointer(factory)
}
}
func experimentsOrDefault(experiments codersdk.Experiments) codersdk.Experiments {
if experiments == nil {
return codersdk.ExperimentsKnown
@@ -139,6 +152,7 @@ func newInternalTestServer(
PendingChatAcquireInterval: testutil.WaitLong,
ProviderAPIKeys: keys,
Experiments: experimentsOrDefault(cfg.experiments),
AIBridgeTransportFactory: cfg.transportFactory,
})
if cfg.startWorker {
server.Start()
@@ -362,9 +376,8 @@ func TestCreateChildSubagentChatRequiresActiveTurnAPIKeyIDForAIGateway(t *testin
})
server := &Server{
db: db,
logger: slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}),
aiGatewayRoutingEnabled: true,
db: db,
logger: slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}),
}
_, err := server.createChildSubagentChat(ctx, parent, "inspect the workspace", "")
require.ErrorContains(t, err, "active turn API key ID is required for subagent messages")
@@ -375,7 +388,6 @@ func TestSendSubagentMessageRequiresActiveTurnAPIKeyIDForAIGateway(t *testing.T)
db, ps := dbtestutil.NewDB(t)
server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{})
server.aiGatewayRoutingEnabled = true
ctx := chatdTestContext(t)
user, org, model := seedInternalChatDeps(t, db)
@@ -426,7 +438,6 @@ func TestSpawnAgentUsesActiveTurnAPIKeyIDFromContext(t *testing.T) {
db, ps := dbtestutil.NewDB(t)
server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{})
server.aiGatewayRoutingEnabled = true
ctx := chatdTestContext(t)
user, org, model := seedInternalChatDeps(t, db)
@@ -592,11 +603,10 @@ func TestResolveChatModel_AIProviderDisabled(t *testing.T) {
LastModelConfigID: modelConfig.ID,
})
model, config, keys, _, debugEnabled, resolvedProvider, resolvedModel, err := server.resolveChatModel(ctx, chat, modelBuildOptions{})
model, config, _, debugEnabled, resolvedProvider, resolvedModel, err := server.resolveChatModel(ctx, chat, modelBuildOptions{})
require.ErrorContains(t, err, "is disabled")
require.Nil(t, model)
require.Equal(t, database.ChatModelConfig{}, config)
require.Equal(t, chatprovider.ProviderAPIKeys{}, keys)
require.False(t, debugEnabled)
require.Empty(t, resolvedProvider)
require.Empty(t, resolvedModel)
@@ -932,6 +942,15 @@ func createInternalParentChat(
return parentChat
}
// withSubagentDelegatedKey enriches ctx with a delegated API key ID for
// subagent tool callbacks. AI Gateway routing requires this key on the
// context; tests that do not otherwise set it should call this helper
// before invoking runSpawnAgentTool or runSubagentTool with spawn_agent.
func withSubagentDelegatedKey(ctx context.Context, t *testing.T, db database.Store, ownerID uuid.UUID) context.Context {
t.Helper()
return aibridge.WithDelegatedAPIKeyID(ctx, testAPIKeyID(t, db, ownerID))
}
func runSubagentTool(
ctx context.Context,
t *testing.T,
@@ -942,11 +961,6 @@ func runSubagentTool(
args any,
) fantasy.ToolResponse {
t.Helper()
if !server.shouldUseAIGatewayRouting() {
if apiKeyID, ok := aibridge.DelegatedAPIKeyIDFromContext(ctx); !ok || apiKeyID == "" {
ctx = aibridge.WithDelegatedAPIKeyID(ctx, testAPIKeyID(t, server.db, parentChat.OwnerID))
}
}
tools := server.subagentTools(
ctx,
@@ -1084,6 +1098,7 @@ func TestSpawnAgent_GeneralInheritsParentModelWhenOmitted(t *testing.T) {
ctx, t, server, db, org.ID, user.ID, model.ID, "parent-inherited-model",
)
ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID)
resp := runSpawnAgentTool(ctx, t, server, parentChat, spawnAgentArgs{
Type: subagentTypeGeneral,
Prompt: "delegate work",
@@ -1114,6 +1129,7 @@ func TestSpawnAgent_GeneralUsesConfiguredModelOverride(t *testing.T) {
ctx, t, server, db, org.ID, user.ID, model.ID, "parent-general-override",
)
ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID)
resp := runSpawnAgentTool(ctx, t, server, parentChat, spawnAgentArgs{
Type: subagentTypeGeneral,
Prompt: "delegate general work",
@@ -1306,6 +1322,7 @@ func TestSpawnAgent_GeneralHonorsPersonalModelOverrides(t *testing.T) {
"parent-general-personal-override",
)
ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID)
resp := runSpawnAgentTool(ctx, t, server, parentChat, spawnAgentArgs{
Type: subagentTypeGeneral,
Prompt: "delegate general work",
@@ -1367,6 +1384,7 @@ func TestSpawnAgent_GeneralOverrideLogsAndFallsBackWhenCredentialsUnavailable(t
parentChat, err := db.GetChatByID(ctx, parent.ID)
require.NoError(t, err)
ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID)
resp := runSpawnAgentTool(ctx, t, server, parentChat, spawnAgentArgs{
Type: subagentTypeGeneral,
Prompt: "inspect provider credentials",
@@ -1437,6 +1455,7 @@ func TestSpawnAgent_GeneralOverrideLogsAndFallsBackWhenProviderDisabled(t *testi
parentChat, err := db.GetChatByID(ctx, parent.ID)
require.NoError(t, err)
ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID)
resp := runSpawnAgentTool(ctx, t, server, parentChat, spawnAgentArgs{
Type: subagentTypeGeneral,
Prompt: "inspect disabled providers",
@@ -1552,6 +1571,7 @@ func TestSpawnAgent_ExploreUsesConfiguredModelOverride(t *testing.T) {
ctx, t, server, db, org.ID, user.ID, model.ID, "parent-explore-override",
)
ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID)
resp := runSubagentTool(
ctx,
t,
@@ -1589,6 +1609,7 @@ func TestSpawnAgent_ExploreFallsBackToCurrentTurnModel(t *testing.T) {
ctx, t, server, db, org.ID, user.ID, parentModel.ID, "parent-explore-fallback",
)
ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID)
resp := runSubagentTool(
ctx,
t,
@@ -1793,6 +1814,7 @@ func TestSpawnAgent_ExploreHonorsPersonalModelOverrides(t *testing.T) {
"parent-explore-personal-override",
)
ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID)
resp := runSubagentTool(
ctx,
t,
@@ -2069,6 +2091,7 @@ func TestSpawnAgent_ExploreFallsBackOnInvalidUUID(t *testing.T) {
ctx, t, server, db, org.ID, user.ID, parentModel.ID, "parent-explore-invalid-override",
)
ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID)
resp := runSubagentTool(
ctx,
t,
@@ -2104,6 +2127,7 @@ func TestSpawnAgent_ExploreFallsBackWhenOverrideIsUnavailable(t *testing.T) {
ctx, t, server, db, org.ID, user.ID, parentModel.ID, "parent-explore-disabled",
)
ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID)
resp := runSubagentTool(
ctx,
t,
@@ -2151,6 +2175,7 @@ func TestSpawnAgent_ExploreFallsBackWhenOverrideCredentialsAreUnavailable(t *tes
ctx, t, server, db, org.ID, user.ID, parentModel.ID, "parent-explore-missing-user-key",
)
ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID)
resp := runSubagentTool(
ctx,
t,
@@ -2653,6 +2678,7 @@ func TestSubagentLifecycleToolsIncludePersistedSubagentTypeAcrossVariants(t *tes
model.ID,
"parent-lifecycle-"+tt.variant,
)
ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID)
spawnResp := runSpawnAgentTool(ctx, t, server, parentChat, spawnAgentArgs{
Type: tt.variant,
@@ -2800,6 +2826,7 @@ func TestSpawnAgent_ComputerUseUsesComputerUseModelNotParent(t *testing.T) {
parentChat, err := db.GetChatByID(ctx, parent.ID)
require.NoError(t, err)
ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID)
resp := runSubagentTool(
ctx,
t,
@@ -2864,6 +2891,7 @@ func TestSpawnAgent_ComputerUseInheritsMCPServerIDs(t *testing.T) {
parentChat, err := db.GetChatByID(ctx, parent.ID)
require.NoError(t, err)
ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID)
resp := runSubagentTool(
ctx,
t,
@@ -3677,7 +3705,20 @@ func TestAwaitSubagentCompletion(t *testing.T) {
})
db, ps := dbtestutil.NewDB(t)
server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{}, withInternalTestServerWorker())
providerServerURL, err := url.Parse(providerServer.URL)
require.NoError(t, err)
factory := &aibridgeTestFactory{rt: roundTripFunc(func(req *http.Request) (*http.Response, error) {
cloned := req.Clone(req.Context())
cloned.URL.Scheme = providerServerURL.Scheme
cloned.URL.Host = providerServerURL.Host
cloned.Host = providerServerURL.Host
return http.DefaultTransport.RoundTrip(cloned)
})}
server := newInternalTestServer(
t, db, ps, chatprovider.ProviderAPIKeys{},
withInternalTestServerWorker(),
withInternalTestServerTransportFactory(factory),
)
ctx := chatdTestContext(t)
user, org, _ := seedInternalChatDeps(t, db)
provider := dbgen.ChatProvider(t, db, database.ChatProvider{
@@ -3698,7 +3739,7 @@ func TestAwaitSubagentCompletion(t *testing.T) {
shortCtx, cancel := context.WithTimeout(ctx, testutil.IntervalMedium)
defer cancel()
_, _, err := server.awaitSubagentCompletion(
_, _, err = server.awaitSubagentCompletion(
shortCtx, parent.ID, child.ID, 5*time.Second,
)
require.ErrorIs(t, err, context.DeadlineExceeded)
+10 -14
View File
@@ -37,18 +37,16 @@ func readTitleGenerationModelOverride(
func (p *Server) resolveTitleGenerationModelOverride(
ctx context.Context,
chat database.Chat,
keys chatprovider.ProviderAPIKeys,
modelOpts modelBuildOptions,
) (database.ChatModelConfig, fantasy.LanguageModel, chatprovider.ProviderAPIKeys, resolvedModelRoute, bool, error) {
) (database.ChatModelConfig, fantasy.LanguageModel, aiGatewayModelRoute, bool, error) {
raw, err := readTitleGenerationModelOverride(ctx, p.db)
if err != nil {
return database.ChatModelConfig{}, nil, chatprovider.ProviderAPIKeys{}, resolvedModelRoute{}, false, xerrors.Errorf(
return database.ChatModelConfig{}, nil, aiGatewayModelRoute{}, false, xerrors.Errorf(
"read title generation model override: %w",
err,
)
}
overrideProviderKeys := keys
modelConfig, overrideSet, err := p.resolveConfiguredModelOverride(
ctx,
titleGenerationOverrideContext,
@@ -58,32 +56,30 @@ func (p *Server) resolveTitleGenerationModelOverride(
func(ctx context.Context, ownerID uuid.UUID, aiProviderID uuid.UUID) (chatprovider.ProviderAPIKeys, error) {
if aiProviderID == uuid.Nil {
resolvedProviderKeys, err := p.resolveUserProviderAPIKeys(ctx, ownerID, uuid.Nil)
if err != nil || resolvedProviderKeys.Empty() {
resolvedProviderKeys = keys
if err != nil {
return chatprovider.ProviderAPIKeys{}, err
}
overrideProviderKeys = resolvedProviderKeys
return resolvedProviderKeys, nil
}
resolvedProviderKeys, err := p.resolveUserProviderAPIKeys(ctx, ownerID, aiProviderID)
if err != nil {
return chatprovider.ProviderAPIKeys{}, err
}
overrideProviderKeys = resolvedProviderKeys
return resolvedProviderKeys, nil
},
modelOverrideFailureModeHard,
)
if err != nil {
return database.ChatModelConfig{}, nil, chatprovider.ProviderAPIKeys{}, resolvedModelRoute{}, overrideSet, err
return database.ChatModelConfig{}, nil, aiGatewayModelRoute{}, overrideSet, err
}
if !overrideSet {
return database.ChatModelConfig{}, nil, keys, resolvedModelRoute{}, false, nil
return database.ChatModelConfig{}, nil, aiGatewayModelRoute{}, false, nil
}
//nolint:gocritic // Title overrides need chatd-scoped provider reads for user-owned chats.
route, err := p.resolveModelRouteForConfig(dbauthz.AsChatd(ctx), chat.OwnerID, modelConfig, overrideProviderKeys)
route, err := p.resolveModelRouteForConfig(dbauthz.AsChatd(ctx), chat.OwnerID, modelConfig)
if err != nil {
return database.ChatModelConfig{}, nil, chatprovider.ProviderAPIKeys{}, resolvedModelRoute{}, true, err
return database.ChatModelConfig{}, nil, aiGatewayModelRoute{}, true, err
}
model, err := p.newModel(ctx, modelClientRequest{
Chat: chat,
@@ -92,10 +88,10 @@ func (p *Server) resolveTitleGenerationModelOverride(
ExtraHeaders: chatprovider.CoderHeaders(chat),
}, route, modelOpts)
if err != nil {
return database.ChatModelConfig{}, nil, chatprovider.ProviderAPIKeys{}, resolvedModelRoute{}, true, xerrors.Errorf(
return database.ChatModelConfig{}, nil, aiGatewayModelRoute{}, true, xerrors.Errorf(
"create title generation model override: %w",
err,
)
}
return modelConfig, model, route.directProviderKeys(), route, true, nil
return modelConfig, model, route, true, nil
}
+49 -161
View File
@@ -21,7 +21,6 @@ import (
"github.com/coder/coder/v2/coderd/aibridge"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/database/dbmock"
"github.com/coder/coder/v2/coderd/x/chatd/chatprovider"
"github.com/coder/coder/v2/coderd/x/chatd/chattest"
"github.com/coder/coder/v2/codersdk"
"github.com/coder/coder/v2/testutil"
@@ -31,59 +30,6 @@ import (
func TestMaybeGenerateChatTitle_TitleGenerationOverrideUnset(t *testing.T) {
t.Parallel()
t.Run("uses preferred model before fallback", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitShort)
ctrl := gomock.NewController(t)
db := dbmock.NewMockStore(ctrl)
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
chat, messages := titleOverrideTestChatAndMessages(t)
wantTitle := "Preferred title"
var requestCount atomic.Int32
serverURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
requestCount.Add(1)
require.Equal(t, preferredTitleModels[1].model, req.Model)
return chattest.OpenAINonStreamingResponse(`{"title":"` + wantTitle + `"}`)
})
keys := titleOverrideOpenAIKeys(serverURL)
fallbackModel := &chattest.FakeModel{
GenerateObjectFn: func(context.Context, fantasy.ObjectCall) (*fantasy.ObjectResponse, error) {
t.Fatal("fallback model should not be called when preferred model works")
return nil, xerrors.New("unexpected fallback model call")
},
}
db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return("", nil)
db.EXPECT().UpdateChatTitleByID(gomock.Any(), database.UpdateChatTitleByIDParams{
ID: chat.ID,
Title: wantTitle,
}).Return(chatWithTitle(chat, wantTitle), nil)
generated := &generatedChatTitle{}
server := titleOverrideTestServer(db, logger)
server.maybeGenerateChatTitle(
ctx,
chat,
messages,
"openai",
"fallback-chat-model",
fallbackModel,
resolvedModelRoute{},
keys,
modelBuildOptions{},
generated,
logger,
nil,
)
require.Equal(t, int32(1), requestCount.Load())
gotTitle, ok := generated.Load()
require.True(t, ok)
require.Equal(t, wantTitle, gotTitle)
})
t.Run("falls back to chat model when preferred models are unavailable", func(t *testing.T) {
t.Parallel()
@@ -119,8 +65,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideUnset(t *testing.T) {
"openai",
"fallback-chat-model",
fallbackModel,
resolvedModelRoute{},
chatprovider.ProviderAPIKeys{},
aiGatewayModelRoute{},
modelBuildOptions{},
generated,
logger,
@@ -169,8 +114,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideReadDBError(t *testing.T)
"openai",
"fallback-chat-model",
fallbackModel,
resolvedModelRoute{},
chatprovider.ProviderAPIKeys{},
aiGatewayModelRoute{},
modelBuildOptions{},
generated,
logger,
@@ -218,8 +162,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideMalformedFallsThrough(t *
"openai",
"fallback-chat-model",
fallbackModel,
resolvedModelRoute{},
chatprovider.ProviderAPIKeys{},
aiGatewayModelRoute{},
modelBuildOptions{},
generated,
logger,
@@ -246,16 +189,22 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideSetUsable(t *testing.T) {
wantTitle := "Override title"
var requestCount atomic.Int32
serverURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
factory := &aibridgeTestFactory{rt: roundTripFunc(func(req *http.Request) (*http.Response, error) {
requestCount.Add(1)
require.Equal(t, overrideConfig.Model, req.Model)
return chattest.OpenAINonStreamingResponse(`{"title":"` + wantTitle + `"}`)
})
text := strconv.Quote(`{"title":"` + wantTitle + `"}`)
body := `{"id":"resp_test","object":"response","created_at":0,"status":"completed","model":"gpt-4.1","output":[{"id":"msg_test","type":"message","role":"assistant","content":[{"type":"output_text","text":` + text + `}]}],"usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}`
return &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(body)),
Request: req,
}, nil
})}
provider := database.AIProvider{
ID: providerID,
Name: "primary-openai",
Type: database.AIProviderTypeOpenai,
Enabled: true,
BaseUrl: serverURL,
}
fallbackModel := &chattest.FakeModel{
GenerateObjectFn: func(context.Context, fantasy.ObjectCall) (*fantasy.ObjectResponse, error) {
@@ -270,7 +219,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideSetUsable(t *testing.T) {
db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{
ProviderID: providerID,
APIKey: "test-key",
}}, nil).Times(2)
}}, nil).AnyTimes()
db.EXPECT().UpdateChatTitleByID(gomock.Any(), database.UpdateChatTitleByIDParams{
ID: chat.ID,
Title: wantTitle,
@@ -278,6 +227,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideSetUsable(t *testing.T) {
generated := &generatedChatTitle{}
server := titleOverrideTestServer(db, logger)
server.aibridgeTransportFactory = aibridgeTestFactoryPointer(factory)
server.maybeGenerateChatTitle(
ctx,
chat,
@@ -285,9 +235,8 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideSetUsable(t *testing.T) {
"openai",
"fallback-chat-model",
fallbackModel,
resolvedModelRoute{},
chatprovider.ProviderAPIKeys{},
modelBuildOptions{},
aiGatewayModelRoute{},
modelBuildOptions{ActiveAPIKeyID: uuid.NewString()},
generated,
logger,
nil,
@@ -327,8 +276,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideSetUnusableSkips(t *testi
"openai",
"fallback-chat-model",
fallbackModel,
resolvedModelRoute{},
chatprovider.ProviderAPIKeys{},
aiGatewayModelRoute{},
modelBuildOptions{},
generated,
logger,
@@ -352,18 +300,10 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideCallFailureSkipsFallback(
overrideConfig.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true}
var requestCount atomic.Int32
serverURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
factory := &aibridgeTestFactory{rt: roundTripFunc(func(req *http.Request) (*http.Response, error) {
requestCount.Add(1)
require.Equal(t, overrideConfig.Model, req.Model)
return chattest.OpenAINonStreamingResponse(`{"title":""}`)
})
provider := database.AIProvider{
ID: providerID,
Type: database.AIProviderTypeOpenai,
Enabled: true,
BaseUrl: serverURL,
}
keys := titleOverrideOpenAIKeys(serverURL)
return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody}, nil
})}
fallbackModel := &chattest.FakeModel{
GenerateObjectFn: func(context.Context, fantasy.ObjectCall) (*fantasy.ObjectResponse, error) {
t.Fatal("fallback model should not be called after override call failure")
@@ -373,7 +313,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideCallFailureSkipsFallback(
db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil)
db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil)
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes()
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai), nil).AnyTimes()
db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{
ProviderID: providerID,
APIKey: "test-key",
@@ -381,6 +321,7 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideCallFailureSkipsFallback(
generated := &generatedChatTitle{}
server := titleOverrideTestServer(db, logger)
server.aibridgeTransportFactory = aibridgeTestFactoryPointer(factory)
server.maybeGenerateChatTitle(
ctx,
chat,
@@ -388,9 +329,8 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideCallFailureSkipsFallback(
"openai",
"fallback-chat-model",
fallbackModel,
resolvedModelRoute{},
keys,
modelBuildOptions{},
aiGatewayModelRoute{},
modelBuildOptions{ActiveAPIKeyID: uuid.NewString()},
generated,
logger,
nil,
@@ -416,35 +356,20 @@ func TestResolveManualTitleModel_TitleGenerationOverrideUnset(t *testing.T) {
Model: preferredTitleModels[1].model,
Enabled: true,
}
serverURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
t.Fatal("model construction should not call the provider")
return chattest.OpenAIResponse{}
})
provider := database.AIProvider{
ID: providerID,
Type: database.AIProviderTypeOpenai,
Enabled: true,
BaseUrl: serverURL,
}
db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return("", nil)
db.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return([]database.GetEnabledChatModelConfigsRow{
{ChatModelConfig: database.ChatModelConfig{Model: "gpt-4.1", Enabled: true}, Provider: "openai"},
{ChatModelConfig: preferredConfig, Provider: preferredTitleModels[1].provider},
}, nil)
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes()
db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{
ProviderID: providerID,
APIKey: "test-key",
}}, nil).AnyTimes()
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai), nil).AnyTimes()
server := titleOverrideTestServer(db, logger)
model, gotConfig, _, err := server.resolveManualTitleModel(
model, gotConfig, err := server.resolveManualTitleModel(
ctx,
db,
chat,
chatprovider.ProviderAPIKeys{ByProvider: map[string]string{"openai": "test-key"}},
modelBuildOptions{},
modelBuildOptions{ActiveAPIKeyID: uuid.NewString()},
)
require.NoError(t, err)
require.NotNil(t, model)
@@ -472,6 +397,7 @@ func TestResolveManualTitleModel_TitleGenerationOverrideUnsetAIProvider(t *testi
})
provider := database.AIProvider{
ID: providerID,
Name: "primary-openai",
Type: database.AIProviderTypeOpenai,
Enabled: true,
BaseUrl: serverURL,
@@ -485,20 +411,18 @@ func TestResolveManualTitleModel_TitleGenerationOverrideUnsetAIProvider(t *testi
db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{
ProviderID: providerID,
APIKey: "test-key",
}}, nil)
}}, nil).AnyTimes()
server := titleOverrideTestServer(db, logger)
model, gotConfig, gotKeys, err := server.resolveManualTitleModel(
model, gotConfig, err := server.resolveManualTitleModel(
ctx,
db,
chat,
chatprovider.ProviderAPIKeys{},
modelBuildOptions{},
modelBuildOptions{ActiveAPIKeyID: uuid.NewString()},
)
require.NoError(t, err)
require.NotNil(t, model)
require.Equal(t, preferredConfig, gotConfig)
require.Equal(t, "test-key", gotKeys.APIKey("openai"))
}
func TestResolveManualTitleModel_TitleGenerationOverrideReadDBError(t *testing.T) {
@@ -516,35 +440,20 @@ func TestResolveManualTitleModel_TitleGenerationOverrideReadDBError(t *testing.T
Model: preferredTitleModels[1].model,
Enabled: true,
}
serverURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
t.Fatal("model construction should not call the provider")
return chattest.OpenAIResponse{}
})
provider := database.AIProvider{
ID: providerID,
Type: database.AIProviderTypeOpenai,
Enabled: true,
BaseUrl: serverURL,
}
db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return("", sql.ErrConnDone)
db.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return([]database.GetEnabledChatModelConfigsRow{
{ChatModelConfig: database.ChatModelConfig{Model: "gpt-4.1", Enabled: true}, Provider: "openai"},
{ChatModelConfig: preferredConfig, Provider: preferredTitleModels[1].provider},
}, nil)
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes()
db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{
ProviderID: providerID,
APIKey: "test-key",
}}, nil).AnyTimes()
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai), nil).AnyTimes()
server := titleOverrideTestServer(db, logger)
model, gotConfig, _, err := server.resolveManualTitleModel(
model, gotConfig, err := server.resolveManualTitleModel(
ctx,
db,
chat,
chatprovider.ProviderAPIKeys{ByProvider: map[string]string{"openai": "test-key"}},
modelBuildOptions{},
modelBuildOptions{ActiveAPIKeyID: uuid.NewString()},
)
require.NoError(t, err)
require.NotNil(t, model)
@@ -562,32 +471,21 @@ func TestResolveManualTitleModel_TitleGenerationOverrideSetUsable(t *testing.T)
overrideConfig := titleOverrideModelConfig("gpt-4.1", true)
providerID := uuid.New()
overrideConfig.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true}
serverURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
t.Fatal("model construction should not call the provider")
return chattest.OpenAIResponse{}
})
provider := database.AIProvider{
ID: providerID,
Type: database.AIProviderTypeOpenai,
Enabled: true,
BaseUrl: serverURL,
}
db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil)
db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil)
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes()
db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai), nil).AnyTimes()
db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{
ProviderID: providerID,
APIKey: "test-key",
}}, nil).AnyTimes()
server := titleOverrideTestServer(db, logger)
model, gotConfig, _, err := server.resolveManualTitleModel(
model, gotConfig, err := server.resolveManualTitleModel(
ctx,
db,
chat,
chatprovider.ProviderAPIKeys{ByProvider: map[string]string{"openai": "test-key"}},
modelBuildOptions{},
modelBuildOptions{ActiveAPIKeyID: uuid.NewString()},
)
require.NoError(t, err)
require.NotNil(t, model)
@@ -617,11 +515,10 @@ func TestResolveManualTitleModel_TitleGenerationOverrideMissingCredentials(t *te
db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return(nil, nil).AnyTimes()
server := titleOverrideTestServer(db, logger)
model, gotConfig, _, err := server.resolveManualTitleModel(
model, gotConfig, err := server.resolveManualTitleModel(
ctx,
db,
chat,
chatprovider.ProviderAPIKeys{},
modelBuildOptions{},
)
require.Error(t, err)
@@ -721,9 +618,8 @@ func TestGenerateManualTitleCandidate_ActiveAPIKeyIDFallback(t *testing.T) {
}}, nil).AnyTimes()
server := titleOverrideTestServer(db, logger)
server.aiGatewayRoutingEnabled = true
server.aibridgeTransportFactory = aibridgeTestFactoryPointer(factory)
result, err := server.generateManualTitleCandidate(ctx, db, chat, chatprovider.ProviderAPIKeys{})
result, err := server.generateManualTitleCandidate(ctx, db, chat)
if tt.wantErrContains != "" {
require.ErrorContains(t, err, tt.wantErrContains)
return
@@ -751,11 +647,10 @@ func TestResolveManualTitleModel_TitleGenerationOverrideSetUnusable(t *testing.T
db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil)
server := titleOverrideTestServer(db, logger)
model, gotConfig, _, err := server.resolveManualTitleModel(
model, gotConfig, err := server.resolveManualTitleModel(
ctx,
db,
chat,
chatprovider.ProviderAPIKeys{ByProvider: map[string]string{"openai": "test-key"}},
modelBuildOptions{},
)
require.Error(t, err)
@@ -785,10 +680,14 @@ func titleOverrideTestChatAndMessages(t *testing.T) (database.Chat, []database.C
}
func titleOverrideTestServer(db database.Store, logger slog.Logger) *Server {
factory := &aibridgeTestFactory{rt: roundTripFunc(func(*http.Request) (*http.Response, error) {
return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody}, nil
})}
return &Server{
db: db,
logger: logger,
configCache: newChatConfigCache(context.Background(), db, quartz.NewReal()),
db: db,
logger: logger,
configCache: newChatConfigCache(context.Background(), db, quartz.NewReal()),
aibridgeTransportFactory: aibridgeTestFactoryPointer(factory),
}
}
@@ -800,17 +699,6 @@ func titleOverrideModelConfig(model string, enabled bool) database.ChatModelConf
}
}
func titleOverrideOpenAIKeys(serverURL string) chatprovider.ProviderAPIKeys {
return chatprovider.ProviderAPIKeys{
ByProvider: map[string]string{
"openai": "test-key",
},
BaseURLByProvider: map[string]string{
"openai": serverURL,
},
}
}
func chatWithTitle(chat database.Chat, title string) database.Chat {
chat.Title = title
return chat
+6 -4
View File
@@ -4252,7 +4252,7 @@ Write out the current server config as YAML to stdout.`,
},
{
Name: "Chat: AI Gateway Routing Enabled",
Description: "Route chat model requests through AI Gateway when both chat routing and AI Gateway are enabled. Otherwise, chat calls AI providers directly. Pending chats without API key metadata may need a retry or temporary direct routing.",
Description: "Deprecated: AI Gateway routing is now the only routing path. Setting this value has no effect. This option will be removed in a future release.",
Flag: "chat-ai-gateway-routing-enabled",
Env: "CODER_CHAT_AI_GATEWAY_ROUTING_ENABLED",
Value: &c.AI.Chat.AIGatewayRoutingEnabled,
@@ -4930,9 +4930,11 @@ type AIBridgeProxyConfig struct {
}
type ChatConfig struct {
AcquireBatchSize serpent.Int64 `json:"acquire_batch_size" typescript:",notnull"`
DebugLoggingEnabled serpent.Bool `json:"debug_logging_enabled" typescript:",notnull"`
AIGatewayRoutingEnabled serpent.Bool `json:"ai_gateway_routing_enabled" typescript:",notnull" swaggerignore:"true"`
AcquireBatchSize serpent.Int64 `json:"acquire_batch_size" typescript:",notnull"`
DebugLoggingEnabled serpent.Bool `json:"debug_logging_enabled" typescript:",notnull"`
// Deprecated: AI Gateway routing is now the only routing path. Setting this
// value has no effect. This option will be removed in a future release.
AIGatewayRoutingEnabled serpent.Bool `json:"ai_gateway_routing_enabled" typescript:",notnull" swaggerignore:"true"`
}
type AIConfig struct {
+1
View File
@@ -29,6 +29,7 @@
<meta property="docs-url" content="{{ .DocsURL }}" />
<meta property="logo-url" content="{{ .LogoURL }}" />
<meta property="tasks-tab-visible" content="{{ .TasksTabVisible }}" />
<meta property="ai-gateway-enabled" content="{{ .AIGatewayEnabled }}" />
<meta property="permissions" content="{{ .Permissions }}" />
<meta property="organizations" content="{{ .Organizations }}" />
<link
+11 -3
View File
@@ -83,6 +83,7 @@ type Options struct {
Telemetry telemetry.Reporter
Logger slog.Logger
HideAITasks bool
AIGatewayEnabled bool
}
func New(opts *Options) (*Handler, error) {
@@ -266,9 +267,10 @@ type htmlState struct {
Regions string
DocsURL string
TasksTabVisible string
Permissions string
Organizations string
TasksTabVisible string
AIGatewayEnabled string
Permissions string
Organizations string
}
type csrfState struct {
@@ -519,6 +521,12 @@ func (h *Handler) populateHTMLState(
state.TasksTabVisible = html.EscapeString(string(data))
}
})
wg.Go(func() {
data, err := json.Marshal(h.opts.AIGatewayEnabled)
if err == nil {
state.AIGatewayEnabled = html.EscapeString(string(data))
}
})
wg.Go(func() {
sdkOrgs := slice.List(userOrgs, db2sdk.Organization)
data, err := json.Marshal(sdkOrgs)
+4
View File
@@ -1679,6 +1679,10 @@ export interface ChatComputerUseProviderResponse {
export interface ChatConfig {
readonly acquire_batch_size: number;
readonly debug_logging_enabled: boolean;
/**
* @deprecated AI Gateway routing is now the only routing path. Setting this
* value has no effect. This option will be removed in a future release.
*/
readonly ai_gateway_routing_enabled: boolean;
}
@@ -1,6 +1,7 @@
import { act, renderHook } from "@testing-library/react";
import type { Region, User } from "#/api/typesGenerated";
import {
MockAIGatewayEnabled,
MockAppearanceConfig,
MockBuildInfo,
MockEntitlements,
@@ -45,6 +46,7 @@ const mockDataForTags = {
userAppearance: MockUserAppearanceSettings,
regions: MockRegions,
"tasks-tab-visible": MockTasksTabVisible,
"ai-gateway-enabled": MockAIGatewayEnabled,
permissions: MockPermissions,
organizations: [MockOrganization],
} as const satisfies Record<MetadataKey, MetadataValue>;
@@ -82,6 +84,10 @@ const emptyMetadata: RuntimeHtmlMetadata = {
available: false,
value: undefined,
},
"ai-gateway-enabled": {
available: false,
value: undefined,
},
permissions: {
available: false,
value: undefined,
@@ -125,6 +131,10 @@ const populatedMetadata: RuntimeHtmlMetadata = {
available: true,
value: MockTasksTabVisible,
},
"ai-gateway-enabled": {
available: true,
value: MockAIGatewayEnabled,
},
permissions: {
available: true,
value: MockPermissions,
+7
View File
@@ -32,6 +32,7 @@ type AvailableMetadata = Readonly<{
regions: readonly Region[];
"build-info": BuildInfoResponse;
"tasks-tab-visible": boolean;
"ai-gateway-enabled": boolean;
permissions: Permissions;
organizations: Organization[];
}>;
@@ -96,6 +97,7 @@ export class MetadataManager implements MetadataManagerApi {
"build-info": this.registerValue<BuildInfoResponse>("build-info"),
regions: this.registerRegionValue(),
"tasks-tab-visible": this.registerValue<boolean>("tasks-tab-visible"),
"ai-gateway-enabled": this.registerValue<boolean>("ai-gateway-enabled"),
permissions: this.registerValue<Permissions>("permissions"),
organizations: this.registerValue<Organization[]>("organizations"),
};
@@ -249,3 +251,8 @@ export const defaultMetadataManager = new MetadataManager();
export const useEmbeddedMetadata = makeUseEmbeddedMetadata(
defaultMetadataManager,
);
export function useAIGatewayEnabled(): boolean {
const { metadata } = useEmbeddedMetadata();
return metadata["ai-gateway-enabled"].value ?? true;
}
+9 -1
View File
@@ -48,6 +48,7 @@ import type * as TypesGen from "#/api/typesGenerated";
import type { ChatMessagePart } from "#/api/typesGenerated";
import { useProxy } from "#/contexts/ProxyContext";
import { useAuthenticated } from "#/hooks/useAuthenticated";
import { useAIGatewayEnabled } from "#/hooks/useEmbeddedMetadata";
import {
getDefaultOrganizationName,
useDashboard,
@@ -1060,6 +1061,7 @@ const AgentChatPage: FC = () => {
void trackedSync.catch(() => undefined);
};
const aiGatewayDisabled = !useAIGatewayEnabled();
const { store, clearStreamError, upsertCacheMessages } = useChatStore({
chatID: agentId,
chatMessages: chatMessagesList,
@@ -1068,6 +1070,7 @@ const AgentChatPage: FC = () => {
chatQueuedMessages,
setChatErrorReason,
clearChatErrorReason,
aiGatewayDisabled,
});
const liveChatStatus =
useChatSelector(store, selectChatStatus) ?? chatRecord?.status ?? null;
@@ -1159,7 +1162,11 @@ const AgentChatPage: FC = () => {
const isChatSettingsPending =
isUpdateChatPlanModePending || isUpdateChatWorkspacePending;
const isInputDisabled =
!hasModelOptions || isArchived || isChatSettingsPending || isViewerNotOwner;
!hasModelOptions ||
isArchived ||
isChatSettingsPending ||
isViewerNotOwner ||
aiGatewayDisabled;
const canUpdateChatWorkspace = !isArchived && !isViewerNotOwner;
const selectedWorkspaceId = chatQuery.data?.workspace_id ?? null;
@@ -1651,6 +1658,7 @@ const AgentChatPage: FC = () => {
providerCount={providerCount}
modelCount={modelCount}
unsupportedProviderNames={unsupportedProviderNames}
aiGatewayDisabled={aiGatewayDisabled}
hasModelOptions={hasModelOptions}
isModelCatalogLoading={isModelCatalogLoading}
planModeEnabled={planModeEnabled}
@@ -140,6 +140,7 @@ interface AgentChatPageViewProps {
providerCount?: number;
modelCount?: number;
unsupportedProviderNames?: readonly string[];
aiGatewayDisabled?: boolean;
hasModelOptions: boolean;
isModelCatalogLoading?: boolean;
planModeEnabled?: boolean;
@@ -332,6 +333,7 @@ export const AgentChatPageView: FC<AgentChatPageViewProps> = ({
providerCount,
modelCount,
unsupportedProviderNames,
aiGatewayDisabled,
hasModelOptions,
isModelCatalogLoading = false,
planModeEnabled,
@@ -936,6 +938,7 @@ export const AgentChatPageView: FC<AgentChatPageViewProps> = ({
providerCount={providerCount}
modelCount={modelCount}
unsupportedProviderNames={unsupportedProviderNames}
aiGatewayDisabled={aiGatewayDisabled}
selectedModel={effectiveSelectedModel}
onModelChange={setSelectedModel}
modelOptions={modelOptions}
@@ -17,6 +17,7 @@ import { workspaces } from "#/api/queries/workspaces";
import type * as TypesGen from "#/api/typesGenerated";
import { useWebpushNotifications } from "#/contexts/useWebpushNotifications";
import { useAuthenticated } from "#/hooks/useAuthenticated";
import { useAIGatewayEnabled } from "#/hooks/useEmbeddedMetadata";
import {
AgentCreateForm,
type CreateChatOptions,
@@ -40,6 +41,7 @@ const AgentCreatePage: FC = () => {
const location = useLocation();
const navigate = useNavigate();
const { permissions } = useAuthenticated();
const aiGatewayDisabled = !useAIGatewayEnabled();
const chatModelsQuery = useQuery(chatModels());
const chatModelConfigsQuery = useQuery(chatModelConfigs());
@@ -169,6 +171,7 @@ const AgentCreatePage: FC = () => {
providerCount={providerCount}
modelCount={modelCount}
unsupportedProviderNames={unsupportedProviderNames}
aiGatewayDisabled={aiGatewayDisabled}
modelConfigs={chatModelConfigsQuery.data ?? []}
isModelCatalogLoading={isModelCatalogLoading}
isModelConfigsLoading={chatModelConfigsQuery.isLoading}
@@ -335,6 +335,23 @@ export const NoModelOptions: Story = {
},
};
export const AIGatewayDisabledShowsSetupNotice: Story = {
args: {
// canConfigureAgentSetup: false and providerCount/modelCount left
// undefined simulates the model-catalog query still loading, which
// used to make an admin briefly see the wrong copy before this was
// fixed to short-circuit on aiGatewayDisabled directly.
canConfigureAgentSetup: false,
aiGatewayDisabled: true,
},
play: async ({ canvasElement }) => {
const canvas = within(canvasElement);
await expect(
canvas.getByText(/Enable it in your deployment config/),
).toBeInTheDocument();
},
};
export const LoadingSpinner: Story = {
args: {
isDisabled: true,
@@ -192,6 +192,9 @@ interface AgentChatInputProps {
providerCount?: number;
modelCount?: number;
unsupportedProviderNames?: readonly string[];
// AI Gateway is disabled deployment-wide, independent of provider/model
// configuration. Forces the setup notice regardless of the counts above.
aiGatewayDisabled?: boolean;
}
export interface AttachedWorkspaceInfo {
@@ -395,13 +398,16 @@ export const AgentChatInput: FC<AgentChatInputProps> = ({
providerCount,
modelCount,
unsupportedProviderNames = [],
aiGatewayDisabled,
}) => {
const [chatFullWidth] = useChatFullWidth();
const showAgentSetupNotice = canConfigureAgentSetup
? providerCount !== undefined &&
modelCount !== undefined &&
(providerCount === 0 || modelCount === 0)
: modelCount !== undefined && modelCount === 0;
const showAgentSetupNotice =
aiGatewayDisabled ||
(canConfigureAgentSetup
? providerCount !== undefined &&
modelCount !== undefined &&
(providerCount === 0 || modelCount === 0)
: modelCount !== undefined && modelCount === 0);
const internalRef = useRef<ChatMessageInputRef>(null);
const [previewImage, setPreviewImage] = useState<string | null>(null);
const [previewText, setPreviewText] = useState<string | null>(null);
@@ -1064,14 +1070,15 @@ export const AgentChatInput: FC<AgentChatInputProps> = ({
)}
{showAgentSetupNotice && (
<div className="relative z-0 mb-[-2.5rem]">
{canConfigureAgentSetup &&
providerCount !== undefined &&
modelCount !== undefined ? (
{(aiGatewayDisabled ||
(providerCount !== undefined && modelCount !== undefined)) &&
canConfigureAgentSetup ? (
<AgentSetupNotice
isAdmin
providerCount={providerCount}
modelCount={modelCount}
providerCount={providerCount ?? 0}
modelCount={modelCount ?? 0}
unsupportedProviderNames={unsupportedProviderNames}
aiGatewayDisabled={aiGatewayDisabled}
/>
) : (
<AgentSetupNotice
@@ -1079,6 +1086,7 @@ export const AgentChatInput: FC<AgentChatInputProps> = ({
providerCount={0}
modelCount={0}
unsupportedProviderNames={unsupportedProviderNames}
aiGatewayDisabled={aiGatewayDisabled}
/>
)}
</div>
@@ -499,6 +499,20 @@ export const MissingProviderAndModelSetup: Story = {
},
};
export const AIGatewayDisabled: Story = {
args: {
...defaultArgs,
aiGatewayDisabled: true,
},
play: async ({ canvasElement }) => {
const canvas = within(canvasElement);
await expect(canvas.getByRole("textbox")).toHaveAttribute(
"aria-disabled",
"true",
);
},
};
export const PreservesAttachmentsOnFailedSend: Story = {
args: {
...defaultArgs,
@@ -129,6 +129,7 @@ interface AgentCreateFormProps {
providerCount?: number;
modelCount?: number;
unsupportedProviderNames?: readonly string[];
aiGatewayDisabled?: boolean;
isModelCatalogLoading: boolean;
modelConfigs: readonly TypesGen.ChatModelConfig[];
isModelConfigsLoading: boolean;
@@ -154,6 +155,7 @@ export const AgentCreateForm: FC<AgentCreateFormProps> = ({
providerCount,
modelCount,
unsupportedProviderNames,
aiGatewayDisabled,
modelConfigs,
isModelCatalogLoading,
isModelConfigsLoading,
@@ -516,7 +518,8 @@ export const AgentCreateForm: FC<AgentCreateFormProps> = ({
isCreating ||
isForbidden ||
isPersonalModelOverridesLoading ||
!hasModelOptions
!hasModelOptions ||
Boolean(aiGatewayDisabled)
}
isLoading={isCreating}
initialValue={initialInputValue}
@@ -551,6 +554,7 @@ export const AgentCreateForm: FC<AgentCreateFormProps> = ({
providerCount={providerCount}
modelCount={modelCount}
unsupportedProviderNames={unsupportedProviderNames}
aiGatewayDisabled={aiGatewayDisabled}
/>
{modelSelectorHelp ? (
<div className="px-3 pt-1 text-2xs text-content-secondary">
@@ -88,6 +88,26 @@ export const MemberOnlyUnsupportedProvider: Story = {
},
};
// AI Gateway disabled takes precedence even when providers and models are
// configured and the viewer is an admin.
export const AIGatewayDisabled: Story = {
args: {
isAdmin: true,
providerCount: 1,
modelCount: 1,
aiGatewayDisabled: true,
},
play: async ({ canvasElement }) => {
const canvas = within(canvasElement);
await expect(
canvas.getByText(/AI Gateway is disabled/),
).toBeInTheDocument();
await expect(
canvas.getByText(/Enable it in your deployment config/),
).toBeInTheDocument();
},
};
// Both a provider and a model are configured: the notice renders nothing.
export const Configured: Story = {
args: {
@@ -9,6 +9,7 @@ interface AgentSetupNoticeProps {
// Names of configured providers the harness cannot use, populated by
// the page only when no supported provider is configured.
unsupportedProviderNames?: readonly string[];
aiGatewayDisabled?: boolean;
}
const formatProviderList = (names: readonly string[]): string => {
@@ -26,11 +27,26 @@ export const AgentSetupNotice: FC<AgentSetupNoticeProps> = ({
providerCount,
modelCount,
unsupportedProviderNames = [],
aiGatewayDisabled,
}) => {
const hasProvider = providerCount > 0;
const hasModel = modelCount > 0;
const hasUnsupportedProviderNames = unsupportedProviderNames.length > 0;
// AI Gateway can be disabled even when providers/models exist in the DB
// catalog, so check it before the provider/model counts below. Unlike
// the provider/model branches, there is no in-app settings page for
// this deployment-level flag for any audience, so the message doesn't
// vary by isAdmin.
if (aiGatewayDisabled) {
return (
<NoticeContainer>
AI Gateway is disabled. Enable it in your deployment config to chat with
Coder Agents.
</NoticeContainer>
);
}
if (hasProvider && hasModel) {
return null;
}
@@ -343,6 +343,36 @@ describe("useChatStore", () => {
});
});
it("does not open the WebSocket when the AI Gateway is disabled", async () => {
const chatID = "chat-gateway-disabled";
const existingMessage = buildMessage(chatID, 1, "user", "hello");
const queryClient = createTestQueryClient();
const wrapper = createWrapper(queryClient);
const setChatErrorReason = vi.fn();
const clearChatErrorReason = vi.fn();
renderHook(
() =>
useChatStore({
chatID,
chatMessages: [existingMessage],
chatRecord: buildChat(chatID),
chatMessagesData: {
messages: [existingMessage],
queued_messages: [],
has_more: false,
},
chatQueuedMessages: [],
setChatErrorReason,
clearChatErrorReason,
aiGatewayDisabled: true,
}),
{ wrapper },
);
expect(watchChat).not.toHaveBeenCalled();
});
it("keeps create_workspace durable call without result after preview_reset", async () => {
vi.useFakeTimers({ shouldAdvanceTime: true });
@@ -50,6 +50,7 @@ interface UseChatStoreOptions {
chatQueuedMessages: readonly TypesGen.ChatQueuedMessage[] | undefined;
setChatErrorReason: (chatID: string, reason: ChatDetailError) => void;
clearChatErrorReason: (chatID: string) => void;
aiGatewayDisabled?: boolean;
}
export const useChatStore = (
@@ -67,6 +68,7 @@ export const useChatStore = (
chatQueuedMessages,
setChatErrorReason,
clearChatErrorReason,
aiGatewayDisabled = false,
} = options;
const queryClient = useQueryClient();
@@ -358,7 +360,7 @@ export const useChatStore = (
store.resetTransientState();
activeChatIDRef.current = chatID ?? null;
if (!chatID || !initialDataLoaded) {
if (!chatID || !initialDataLoaded || aiGatewayDisabled) {
return;
}
@@ -714,6 +716,7 @@ export const useChatStore = (
activeChatIDRef.current = null;
};
}, [
aiGatewayDisabled,
chatID,
initialDataLoaded,
queryClient,
@@ -177,6 +177,7 @@ interface ChatPageInputProps {
providerCount?: number;
modelCount?: number;
unsupportedProviderNames?: readonly string[];
aiGatewayDisabled?: boolean;
planModeEnabled?: boolean;
onPlanModeToggle?: (enabled: boolean) => void;
isModelCatalogLoading?: boolean;
@@ -247,6 +248,7 @@ export const ChatPageInput: FC<ChatPageInputProps> = ({
providerCount,
modelCount,
unsupportedProviderNames,
aiGatewayDisabled,
planModeEnabled,
onPlanModeToggle,
isModelCatalogLoading = false,
@@ -524,6 +526,7 @@ export const ChatPageInput: FC<ChatPageInputProps> = ({
providerCount={providerCount}
modelCount={modelCount}
unsupportedProviderNames={unsupportedProviderNames}
aiGatewayDisabled={aiGatewayDisabled}
/>
);
+1
View File
@@ -620,6 +620,7 @@ export const MockUserSecrets: TypesGen.UserSecret[] = [
];
export const MockTasksTabVisible: boolean = false;
export const MockAIGatewayEnabled: boolean = true;
export const MockOrganizationMember: TypesGen.OrganizationMemberWithUserData = {
organization_id: MockOrganization.id,