mirror of
https://github.com/coder/coder.git
synced 2026-09-01 14:53:15 +08:00
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:
+2
-3
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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
@@ -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(),
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
Generated
+4
@@ -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,
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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}
|
||||
/>
|
||||
);
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user