feat: route chatd provider traffic through aibridge (#25629)

## Summary

Routes chatd model calls backed by concrete AI Provider rows through the
in-process aibridge transport by default, with deployment options to use
direct provider routing when AI Gateway is disabled or chat AI Gateway
routing is disabled.

- Splits model routing into common, direct provider, and AI Gateway
paths behind a single deployment-mode entry point.
- Builds chatd models through explicit request, route, and options data.
Active API key attribution is passed explicitly instead of being hidden
inside generic model construction.
- For AI Gateway BYOK routes, resolves the user's provider key in chatd,
forwards it through provider-specific auth headers, and sets
`X-Coder-AI-Governance-Token` to the `delegated` marker so aibridge
preserves those headers while still stripping Coder-specific metadata.
- Keeps central provider credentials and deployment fallback credentials
out of forwarded provider auth headers, so AI Gateway central policy
remains authoritative.
- Redacts delegated provider auth from default string formatting to
avoid accidental plaintext logging of user BYOK credentials.
- Covers selected chat models, advisor overrides, title and quickgen
paths, subagent overrides, computer use model selection, and an
integration-style chat turn through the aibridge transport path.
- Persists initiating API key IDs on chat and queued user messages,
including subagent child messages, and fails closed for AI
Gateway-routed model builds without an active key.
- Removes unused `api_key_id` indexes while keeping the persistence
columns and foreign keys.
- Keeps the deployment option available through config and env parsing,
but hides it from CLI help and generated docs.
- Stabilizes the subagent poll fallback test so background CreateChat
processing cannot win the state transition under slower CI environments.

## Tests

- `go test ./coderd/x/chatd -run
'TestAIGatewayProviderAuthForUser|TestAIGatewayProviderAuthRedactsFormatting|TestResolveModelRouteForConfigAIGatewayProviderAuth|TestAIGatewayModelForwardsProviderAuth|TestProcessChat_AIGatewayRoutingUsesDelegatedAPIKey|TestAwaitSubagentCompletion'
-count=1`
- `go test ./coderd/aibridged -run
'TestServeHTTP_DelegatedAPIKey|TestServeHTTP_StripCoderToken' -count=1`
- `git diff --check HEAD~1..HEAD`
- `make lint`

> Mux working on behalf of Mike.
This commit is contained in:
Michael Suchacz
2026-05-26 19:31:52 +00:00
committed by GitHub
parent a56c88a0cc
commit 8b1705eb65
31 changed files with 2463 additions and 377 deletions
+239 -156
View File
@@ -29,6 +29,7 @@ import (
"golang.org/x/xerrors"
"cdr.dev/slog/v3"
"github.com/coder/coder/v2/coderd/aibridge"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/database/db2sdk"
"github.com/coder/coder/v2/coderd/database/dbauthz"
@@ -254,6 +255,9 @@ type Server struct {
metrics *chatloop.Metrics
recordingSem chan struct{}
aibridgeTransportFactory *atomic.Pointer[aibridge.TransportFactory]
aiGatewayRoutingEnabled bool
// Configuration
pendingChatAcquireInterval time.Duration
maxChatsPerAcquire int32
@@ -344,10 +348,11 @@ func (p *Server) resolveAdvisorModelOverride(
fallbackModel fantasy.LanguageModel,
fallbackCallConfig codersdk.ChatModelCallConfig,
providerKeys chatprovider.ProviderAPIKeys,
modelOpts modelBuildOptions,
logger slog.Logger,
) (fantasy.LanguageModel, codersdk.ChatModelCallConfig) {
) (fantasy.LanguageModel, codersdk.ChatModelCallConfig, error) {
if advisorCfg.ModelConfigID == uuid.Nil {
return fallbackModel, fallbackCallConfig
return fallbackModel, fallbackCallConfig, nil
}
// Re-read the override instead of using the cache so disabled models
@@ -363,7 +368,7 @@ func (p *Server) resolveAdvisorModelOverride(
"advisor model config is disabled or unavailable, continuing with chat model",
slog.F("model_config_id", advisorCfg.ModelConfigID),
)
return fallbackModel, fallbackCallConfig
return fallbackModel, fallbackCallConfig, nil
}
logger.Warn(
ctx,
@@ -371,7 +376,7 @@ func (p *Server) resolveAdvisorModelOverride(
slog.F("model_config_id", advisorCfg.ModelConfigID),
slog.Error(err),
)
return fallbackModel, fallbackCallConfig
return fallbackModel, fallbackCallConfig, nil
}
overrideCallConfig := codersdk.ChatModelCallConfig{}
@@ -383,29 +388,48 @@ func (p *Server) resolveAdvisorModelOverride(
slog.F("model_config_id", advisorCfg.ModelConfigID),
slog.Error(err),
)
return fallbackModel, fallbackCallConfig
return fallbackModel, fallbackCallConfig, nil
}
}
overrideModel, err := chatprovider.ModelFromConfig(
overrideConfig.Provider,
overrideConfig.Model,
route, err := p.resolveModelRouteForConfig(
ctx,
chat.OwnerID,
overrideConfig,
providerKeys,
chatprovider.UserAgent(),
chatprovider.CoderHeaders(chat),
nil,
)
if err != nil {
if p.shouldUseAIGatewayRouting() && overrideConfig.AIProviderID.Valid {
return nil, codersdk.ChatModelCallConfig{}, xerrors.Errorf("resolve advisor override route: %w", err)
}
logger.Warn(
ctx,
"failed to resolve advisor override route, continuing with chat model",
slog.F("model_config_id", advisorCfg.ModelConfigID),
slog.Error(err),
)
return fallbackModel, fallbackCallConfig, nil
}
overrideModel, err := p.newModel(ctx, modelClientRequest{
Chat: chat,
ModelName: overrideConfig.Model,
UserAgent: chatprovider.UserAgent(),
ExtraHeaders: chatprovider.CoderHeaders(chat),
}, route, modelOpts)
if err != nil {
if p.shouldUseAIGatewayRouting() && overrideConfig.AIProviderID.Valid {
return nil, codersdk.ChatModelCallConfig{}, xerrors.Errorf("create advisor override model: %w", err)
}
logger.Warn(
ctx,
"failed to create advisor override model, continuing with chat model",
slog.F("model_config_id", advisorCfg.ModelConfigID),
slog.Error(err),
)
return fallbackModel, fallbackCallConfig
return fallbackModel, fallbackCallConfig, nil
}
return overrideModel, overrideCallConfig
return overrideModel, overrideCallConfig, nil
}
func (p *Server) newAdvisorRuntime(
@@ -415,17 +439,22 @@ func (p *Server) newAdvisorRuntime(
fallbackModel fantasy.LanguageModel,
fallbackCallConfig codersdk.ChatModelCallConfig,
providerKeys chatprovider.ProviderAPIKeys,
modelOpts modelBuildOptions,
logger slog.Logger,
) *chatadvisor.Runtime {
advisorModel, advisorCallConfig := p.resolveAdvisorModelOverride(
) (*chatadvisor.Runtime, error) {
advisorModel, advisorCallConfig, err := p.resolveAdvisorModelOverride(
ctx,
chat,
advisorCfg,
fallbackModel,
fallbackCallConfig,
providerKeys,
modelOpts,
logger,
)
if err != nil {
return nil, err
}
maxUsesPerRun := advisorCfg.MaxUsesPerRun
switch {
@@ -441,7 +470,7 @@ func (p *Server) newAdvisorRuntime(
"invalid advisor max uses per run, continuing without advisor",
slog.F("max_uses_per_run", maxUsesPerRun),
)
return nil
return nil, nil //nolint:nilnil // Nil runtime with nil error means advisor is skipped for this turn.
}
maxOutputTokens := advisorCfg.MaxOutputTokens
@@ -468,9 +497,9 @@ func (p *Server) newAdvisorRuntime(
"failed to create advisor runtime, continuing without advisor",
slog.Error(err),
)
return nil
return nil, nil //nolint:nilnil // Nil runtime with nil error means advisor is skipped for this turn.
}
return rt
return rt, nil
}
// cachedWorkspaceMCPTools stores workspace MCP tools discovered
@@ -1436,6 +1465,7 @@ type CreateOptions struct {
ClientType database.ChatClientType
SystemPrompt string
InitialUserContent []codersdk.ChatMessagePart
APIKeyID string
MCPServerIDs []uuid.UUID
Labels database.StringMap
DynamicTools json.RawMessage
@@ -1460,6 +1490,7 @@ type SendMessageOptions struct {
CreatedBy uuid.UUID
Content []codersdk.ChatMessagePart
ModelConfigID uuid.UUID
APIKeyID string
BusyBehavior SendMessageBusyBehavior
PlanMode *database.NullChatPlanMode
MCPServerIDs *[]uuid.UUID
@@ -1479,6 +1510,7 @@ type EditMessageOptions struct {
CreatedBy uuid.UUID
EditedMessageID int64
Content []codersdk.ChatMessagePart
APIKeyID string
// ModelConfigID, when non-zero, overrides the model used for
// the replacement user message. When set to uuid.Nil the
// original message's model is preserved.
@@ -1647,7 +1679,7 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C
database.ChatMessageVisibilityBoth,
opts.ModelConfigID,
chatprompt.CurrentContentVersion,
).withCreatedBy(opts.OwnerID))
).withCreatedBy(opts.OwnerID).withAPIKeyID(opts.APIKeyID))
_, err = tx.InsertChatMessages(ctx, msgParams)
if err != nil {
@@ -1786,6 +1818,10 @@ func (p *Server) SendMessage(
UUID: modelConfigID,
Valid: modelConfigID != uuid.Nil,
},
APIKeyID: sql.NullString{
String: opts.APIKeyID,
Valid: opts.APIKeyID != "",
},
})
if err != nil {
return xerrors.Errorf("insert queued message: %w", err)
@@ -1810,6 +1846,7 @@ func (p *Server) SendMessage(
modelConfigID,
content,
opts.CreatedBy,
opts.APIKeyID,
)
if err != nil {
return err
@@ -2083,7 +2120,7 @@ func (p *Server) EditMessage(
editedMsg.Visibility,
messageModelConfigID,
chatprompt.CurrentContentVersion,
).withCreatedBy(opts.CreatedBy))
).withCreatedBy(opts.CreatedBy).withAPIKeyID(opts.APIKeyID))
newMessages, err := insertChatMessageWithStore(ctx, tx, msgParams)
if err != nil {
return xerrors.Errorf("insert replacement message: %w", err)
@@ -2416,12 +2453,14 @@ func (p *Server) PromoteQueued(
var (
targetContent json.RawMessage
targetModelConfigID uuid.NullUUID
targetAPIKeyID sql.NullString
found bool
)
for _, qm := range queuedMessages {
if qm.ID == opts.QueuedMessageID {
targetContent = qm.Content
targetModelConfigID = qm.ModelConfigID
targetAPIKeyID = qm.APIKeyID
found = true
break
}
@@ -2511,6 +2550,7 @@ func (p *Server) PromoteQueued(
Valid: len(targetContent) > 0,
},
opts.CreatedBy,
targetAPIKeyID.String,
)
if err != nil {
return err
@@ -2754,6 +2794,7 @@ func (p *Server) SubmitToolResults(
params := database.InsertChatMessagesParams{
ChatID: opts.ChatID,
CreatedBy: make([]uuid.UUID, n),
APIKeyID: make([]string, n),
ModelConfigID: make([]uuid.UUID, n),
Role: make([]database.ChatMessageRole, n),
Content: make([]string, n),
@@ -2862,16 +2903,18 @@ var ErrManualTitleRegenerationInProgress = xerrors.New(
)
type manualTitleCandidateResult struct {
title string
modelConfig database.ChatModelConfig
usage fantasy.Usage
hasMessages bool
title string
modelConfig database.ChatModelConfig
usage fantasy.Usage
activeAPIKeyID string
hasMessages bool
}
type manualTitleGenerationError struct {
cause error
modelConfig database.ChatModelConfig
usage fantasy.Usage
cause error
modelConfig database.ChatModelConfig
usage fantasy.Usage
activeAPIKeyID string
}
func (e *manualTitleGenerationError) Error() string {
@@ -3105,6 +3148,7 @@ func (p *Server) recordManualTitleGenerationFailure(
chat,
generationErr.modelConfig,
generationErr.usage,
generationErr.activeAPIKeyID,
"",
); recordErr != nil {
return errors.Join(
@@ -3154,11 +3198,13 @@ func (p *Server) generateManualTitleCandidate(
if len(messages) == 0 {
return manualTitleCandidateResult{}, nil
}
modelOpts := modelBuildOptionsFromMessages(messages)
model, modelConfig, modelKeys, err := p.resolveManualTitleModel(ctx, store, chat, keys)
model, modelConfig, modelKeys, err := p.resolveManualTitleModel(ctx, store, chat, keys, modelOpts)
result := manualTitleCandidateResult{
modelConfig: modelConfig,
hasMessages: true,
modelConfig: modelConfig,
activeAPIKeyID: modelOpts.ActiveAPIKeyID,
hasMessages: true,
}
if err != nil {
return result, err
@@ -3174,6 +3220,7 @@ func (p *Server) generateManualTitleCandidate(
chat,
modelConfig,
modelKeys,
modelOpts,
messages,
model,
)
@@ -3189,9 +3236,10 @@ func (p *Server) generateManualTitleCandidate(
return result, wrappedErr
}
return result, &manualTitleGenerationError{
cause: wrappedErr,
modelConfig: modelConfig,
usage: usage,
cause: wrappedErr,
modelConfig: modelConfig,
usage: usage,
activeAPIKeyID: modelOpts.ActiveAPIKeyID,
}
}
@@ -3220,6 +3268,7 @@ func (p *Server) proposeChatTitleWithStore(
chat,
result.modelConfig,
result.usage,
result.activeAPIKeyID,
"",
); recordErr != nil {
return "", xerrors.Errorf("record manual title usage: %w", recordErr)
@@ -3250,6 +3299,7 @@ func (p *Server) regenerateChatTitleWithStore(
chat,
result.modelConfig,
result.usage,
result.activeAPIKeyID,
result.title,
)
if recordErr != nil {
@@ -3272,6 +3322,7 @@ func (p *Server) prepareManualTitleDebugRun(
chat database.Chat,
modelConfig database.ChatModelConfig,
keys chatprovider.ProviderAPIKeys,
modelOpts modelBuildOptions,
messages []database.ChatMessage,
fallbackModel fantasy.LanguageModel,
) (context.Context, fantasy.LanguageModel, func(error)) {
@@ -3279,15 +3330,21 @@ func (p *Server) prepareManualTitleDebugRun(
titleModel := fallbackModel
finishDebugRun := func(error) {}
httpClient := &http.Client{Transport: &chatdebug.RecordingTransport{}}
debugModel, debugModelErr := chatprovider.ModelFromConfig(
modelConfig.Provider,
modelConfig.Model,
keys,
chatprovider.UserAgent(),
chatprovider.CoderHeaders(chat),
httpClient,
)
route, routeErr := p.resolveModelRouteForConfig(ctx, chat.OwnerID, modelConfig, keys)
debugOpts := modelOpts
debugOpts.RecordHTTP = true
var debugModelErr error
var debugModel fantasy.LanguageModel
if routeErr != nil {
debugModelErr = routeErr
} else {
debugModel, debugModelErr = p.newModel(ctx, modelClientRequest{
Chat: chat,
ModelName: modelConfig.Model,
UserAgent: chatprovider.UserAgent(),
ExtraHeaders: chatprovider.CoderHeaders(chat),
}, route, debugOpts)
}
switch {
case debugModelErr != nil:
p.logger.Warn(ctx, "failed to create debug-aware manual title model",
@@ -3535,11 +3592,13 @@ func (p *Server) resolveManualTitleModel(
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(
overrideConfig, overrideModel, overrideKeys, _, overrideSet, overrideErr := p.resolveTitleGenerationModelOverride(
ctx,
chat,
keys,
modelOpts,
)
if overrideErr != nil {
if overrideSet {
@@ -3562,15 +3621,15 @@ func (p *Server) resolveManualTitleModel(
slog.F("chat_id", chat.ID),
slog.Error(err),
)
return p.resolveFallbackManualTitleModel(ctx, chat, keys)
return p.resolveFallbackManualTitleModel(ctx, chat, keys, modelOpts)
}
config, ok := selectPreferredConfiguredShortTextModelConfig(configs)
if !ok {
return p.resolveFallbackManualTitleModel(ctx, chat, keys)
return p.resolveFallbackManualTitleModel(ctx, chat, keys, modelOpts)
}
providerHint, modelKeys, err := p.resolveModelConfigProviderHintAndKeys(ctx, chat.OwnerID, config, keys)
route, err := p.resolveModelRouteForConfig(ctx, chat.OwnerID, config, keys)
if err != nil {
p.logger.Debug(ctx, "manual title preferred model unavailable",
slog.F("chat_id", chat.ID),
@@ -3578,33 +3637,32 @@ func (p *Server) resolveManualTitleModel(
slog.F("model", config.Model),
slog.Error(err),
)
return p.resolveFallbackManualTitleModel(ctx, chat, keys)
return p.resolveFallbackManualTitleModel(ctx, chat, keys, modelOpts)
}
model, err := chatprovider.ModelFromConfig(
providerHint,
config.Model,
modelKeys,
chatprovider.UserAgent(),
chatprovider.CoderHeaders(chat),
nil,
)
model, err := p.newModel(ctx, modelClientRequest{
Chat: chat,
ModelName: config.Model,
UserAgent: chatprovider.UserAgent(),
ExtraHeaders: chatprovider.CoderHeaders(chat),
}, route, modelOpts)
if err != nil {
p.logger.Debug(ctx, "manual title preferred model unavailable",
slog.F("chat_id", chat.ID),
slog.F("provider", providerHint),
slog.F("provider", config.Provider),
slog.F("model", config.Model),
slog.Error(err),
)
return p.resolveFallbackManualTitleModel(ctx, chat, keys)
return p.resolveFallbackManualTitleModel(ctx, chat, keys, modelOpts)
}
return model, config, modelKeys, nil
return model, config, route.directProviderKeys(), nil
}
func (p *Server) resolveFallbackManualTitleModel(
ctx context.Context,
chat database.Chat,
keys chatprovider.ProviderAPIKeys,
modelOpts modelBuildOptions,
) (fantasy.LanguageModel, database.ChatModelConfig, chatprovider.ProviderAPIKeys, error) {
config, err := p.resolveModelConfig(ctx, chat)
if err != nil {
@@ -3613,25 +3671,23 @@ func (p *Server) resolveFallbackManualTitleModel(
err,
)
}
providerHint, modelKeys, err := p.resolveModelConfigProviderHintAndKeys(ctx, chat.OwnerID, config, keys)
route, err := p.resolveModelRouteForConfig(ctx, chat.OwnerID, config, keys)
if err != nil {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, err
}
model, err := chatprovider.ModelFromConfig(
providerHint,
config.Model,
modelKeys,
chatprovider.UserAgent(),
chatprovider.CoderHeaders(chat),
nil,
)
model, err := p.newModel(ctx, modelClientRequest{
Chat: chat,
ModelName: config.Model,
UserAgent: chatprovider.UserAgent(),
ExtraHeaders: chatprovider.CoderHeaders(chat),
}, route, modelOpts)
if err != nil {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, xerrors.Errorf(
"create fallback manual title model: %w",
err,
)
}
return model, config, modelKeys, nil
return model, config, route.directProviderKeys(), nil
}
func mergeManualTitleMessages(
@@ -3682,6 +3738,7 @@ func recordManualTitleUsage(
chat database.Chat,
modelConfig database.ChatModelConfig,
usage fantasy.Usage,
activeAPIKeyID string,
newTitle string,
) (database.Chat, error) {
hasUsage := usage != (fantasy.Usage{})
@@ -3720,6 +3777,7 @@ func recordManualTitleUsage(
messages, err := tx.InsertChatMessages(ctx, database.InsertChatMessagesParams{
ChatID: chat.ID,
CreatedBy: []uuid.UUID{chat.OwnerID},
APIKeyID: []string{activeAPIKeyID},
ModelConfigID: []uuid.UUID{modelConfig.ID},
Role: []database.ChatMessageRole{database.ChatMessageRoleAssistant},
Content: []string{content},
@@ -3845,6 +3903,7 @@ type chatMessage struct {
visibility database.ChatMessageVisibility
modelConfigID uuid.UUID
createdBy uuid.UUID
apiKeyID string
contentVersion int16
compressed bool
inputTokens int64
@@ -3880,6 +3939,11 @@ func (m chatMessage) withCreatedBy(id uuid.UUID) chatMessage {
return m
}
func (m chatMessage) withAPIKeyID(id string) chatMessage {
m.apiKeyID = id
return m
}
func (m chatMessage) withCompressed() chatMessage {
m.compressed = true
return m
@@ -3924,6 +3988,7 @@ func appendChatMessage(
msg chatMessage,
) {
params.CreatedBy = append(params.CreatedBy, msg.createdBy)
params.APIKeyID = append(params.APIKeyID, msg.apiKeyID)
params.ModelConfigID = append(params.ModelConfigID, msg.modelConfigID)
params.Role = append(params.Role, msg.role)
params.Content = append(params.Content, string(msg.content.RawMessage))
@@ -3973,6 +4038,7 @@ func insertUserMessageAndSetPending(
modelConfigID uuid.UUID,
content pqtype.NullRawMessage,
createdBy uuid.UUID,
apiKeyID string,
) (database.ChatMessage, database.Chat, error) {
msgParams := database.InsertChatMessagesParams{ //nolint:exhaustruct // Fields populated by appendChatMessage.
ChatID: lockedChat.ID,
@@ -3983,7 +4049,7 @@ func insertUserMessageAndSetPending(
database.ChatMessageVisibilityBoth,
modelConfigID,
chatprompt.CurrentContentVersion,
).withCreatedBy(createdBy))
).withCreatedBy(createdBy).withAPIKeyID(apiKeyID))
messages, err := insertChatMessageWithStore(ctx, store, msgParams)
if err != nil {
return database.ChatMessage{}, database.Chat{}, err
@@ -4052,7 +4118,10 @@ type Config struct {
WebpushDispatcher webpush.Dispatcher
UsageTracker *workspacestats.UsageTracker
Clock quartz.Clock
PrometheusRegistry prometheus.Registerer
AIBridgeTransportFactory *atomic.Pointer[aibridge.TransportFactory]
AIGatewayRoutingEnabled bool
PrometheusRegistry prometheus.Registerer
// OIDCTokenSource resolves the calling user's OIDC access
// token for MCP servers configured with auth_type=user_oidc.
@@ -4139,6 +4208,8 @@ func New(cfg Config) *Server {
debugSvc.SetStaleAfter(inFlightChatStaleAfter * 3)
return debugSvc
},
aibridgeTransportFactory: cfg.AIBridgeTransportFactory,
aiGatewayRoutingEnabled: cfg.AIGatewayRoutingEnabled,
pendingChatAcquireInterval: pendingChatAcquireInterval,
maxChatsPerAcquire: maxChatsPerAcquire,
inFlightChatStaleAfter: inFlightChatStaleAfter,
@@ -5803,7 +5874,7 @@ func (p *Server) tryAutoPromoteQueuedMessage(
database.ChatMessageVisibilityBoth,
effectiveModelConfigID,
chatprompt.CurrentContentVersion,
).withCreatedBy(chat.OwnerID))
).withCreatedBy(chat.OwnerID).withAPIKeyID(nextQueued.APIKeyID.String))
msgs, err := insertChatMessageWithStore(ctx, tx, msgParams)
if err != nil {
return nil, nil, false, xerrors.Errorf("insert promoted message: %w", err)
@@ -6322,11 +6393,39 @@ type runChatResult struct {
ProviderKeys chatprovider.ProviderAPIKeys
PendingDynamicToolCalls []chatloop.PendingToolCall
FallbackProvider string
FallbackRoute resolvedModelRoute
FallbackModel string
ModelBuildOptions modelBuildOptions
TriggerMessageID int64
HistoryTipMessageID int64
}
func contextWithActiveTurnAPIKeyID(ctx context.Context, messages []database.ChatMessage) context.Context {
apiKeyID, ok := activeTurnAPIKeyIDFromMessages(messages)
if !ok {
return ctx
}
return aibridge.WithDelegatedAPIKeyID(ctx, apiKeyID)
}
func activeTurnAPIKeyIDFromMessages(messages []database.ChatMessage) (string, bool) {
for i := len(messages) - 1; i >= 0; i-- {
message := messages[i]
if message.Role != database.ChatMessageRoleUser {
continue
}
if message.Visibility != database.ChatMessageVisibilityBoth &&
message.Visibility != database.ChatMessageVisibilityUser {
continue
}
if !message.APIKeyID.Valid || message.APIKeyID.String == "" {
return "", false
}
return message.APIKeyID.String, true
}
return "", false
}
func allToolNames(allTools []fantasy.AgentTool) []string {
toolNames := make([]string, 0, len(allTools))
for _, tool := range allTools {
@@ -6948,12 +7047,20 @@ func (p *Server) runChat(
err error
debugEnabled bool
debugProvider string
modelRoute resolvedModelRoute
debugModel string
)
// Load MCP server configs and user tokens in parallel with
// model resolution and message loading. These queries have
// no dependencies on each other and all hit different tables.
messages, err = p.db.GetChatMessagesForPromptByChatID(ctx, chat.ID)
if err != nil {
return result, xerrors.Errorf("get chat messages: %w", err)
}
modelOpts := modelBuildOptionsFromMessages(messages)
ctx = contextWithActiveTurnAPIKeyID(ctx, messages)
// Load MCP server configs and user tokens in parallel with model
// resolution. These queries have no dependencies on each other and all
// hit different tables.
var (
mcpConfigs []database.MCPServerConfig
mcpTokens []database.MCPServerUserToken
@@ -6961,7 +7068,7 @@ func (p *Server) runChat(
var g errgroup.Group
g.Go(func() error {
var err error
model, modelConfig, providerKeys, debugEnabled, debugProvider, debugModel, err = p.resolveChatModel(ctx, chat)
model, modelConfig, providerKeys, modelRoute, debugEnabled, debugProvider, debugModel, err = p.resolveChatModel(ctx, chat, modelOpts)
if err != nil {
return err
}
@@ -6972,14 +7079,6 @@ func (p *Server) runChat(
}
return nil
})
g.Go(func() error {
var err error
messages, err = p.db.GetChatMessagesForPromptByChatID(ctx, chat.ID)
if err != nil {
return xerrors.Errorf("get chat messages: %w", err)
}
return nil
})
if len(chat.MCPServerIDs) > 0 {
g.Go(func() error {
var err error
@@ -7055,15 +7154,20 @@ func (p *Server) runChat(
// registering the runtime there would inject guidance for a tool
// that is never exposed to the model.
if advisorCfg.Enabled && isRootChat && !isPlanModeTurn && !isExploreSubagent {
advisorRuntime = p.newAdvisorRuntime(
var advisorErr error
advisorRuntime, advisorErr = p.newAdvisorRuntime(
ctx,
chat,
advisorCfg,
model,
callConfig,
providerKeys,
modelOpts,
logger,
)
if advisorErr != nil {
return result, advisorErr
}
}
var advisorPromptSnapshot []fantasy.Message
@@ -7090,7 +7194,9 @@ func (p *Server) runChat(
result.StatusLabelModel = model
result.ProviderKeys = providerKeys
result.FallbackProvider = modelConfig.Provider
result.FallbackRoute = modelRoute
result.FallbackModel = modelConfig.Model
result.ModelBuildOptions = modelOpts
debugSvc := p.existingDebugService()
// Fire title generation asynchronously so it doesn't block the
// chat response. It uses a detached context so it can finish
@@ -7111,7 +7217,9 @@ func (p *Server) runChat(
modelConfig.Provider,
modelConfig.Model,
titleModel,
modelRoute,
titleProviderKeys,
modelOpts,
generatedTitle,
titleLogger,
debugSvc,
@@ -7681,20 +7789,21 @@ func (p *Server) runChat(
}
if isComputerUse {
computerUseProviderKeys, keyErr := p.resolveUserProviderAPIKeysForProviderType(ctx, chat.OwnerID, computerUseModelProvider)
computerUseRoute, keyErr := p.resolveModelRouteForProviderType(ctx, chat.OwnerID, computerUseModelProvider)
if keyErr != nil {
return result, xerrors.Errorf("resolve computer use provider API keys: %w", keyErr)
return result, xerrors.Errorf("resolve computer use provider route: %w", keyErr)
}
providerKeys = computerUseProviderKeys
providerKeys = computerUseRoute.directProviderKeys()
// Override model for computer use subagent.
cuModel, cuDebugEnabled, resolvedProvider, resolvedModel, cuErr := p.resolveComputerUseModel(
ctx,
chat,
providerKeys,
computerUseRoute,
computerUseProvider,
computerUseModelProvider,
computerUseModelName,
modelOpts,
)
if cuErr != nil {
return result, cuErr
@@ -8394,45 +8503,15 @@ func (p *Server) persistChatContextSummary(
return nil
}
func (p *Server) resolveModelConfigProviderHintAndKeys(
ctx context.Context,
ownerID uuid.UUID,
modelConfig database.ChatModelConfig,
fallbackKeys chatprovider.ProviderAPIKeys,
) (string, chatprovider.ProviderAPIKeys, error) {
providerHint := modelConfig.Provider
if !modelConfig.AIProviderID.Valid {
if !fallbackKeys.Empty() && userCanUseProviderKeys(fallbackKeys, providerHint) {
return providerHint, fallbackKeys, nil
}
keys, err := p.resolveUserProviderAPIKeys(ctx, ownerID, uuid.Nil)
if err != nil {
return "", chatprovider.ProviderAPIKeys{}, xerrors.Errorf("resolve provider API keys: %w", err)
}
return providerHint, keys, nil
}
//nolint:gocritic // Manual title generation needs chatd-scoped provider reads for user-owned chats.
provider, err := p.db.GetAIProviderByID(dbauthz.AsChatd(ctx), modelConfig.AIProviderID.UUID)
if err != nil {
return "", chatprovider.ProviderAPIKeys{}, xerrors.Errorf("get AI provider: %w", err)
}
if !provider.Enabled {
return "", chatprovider.ProviderAPIKeys{}, xerrors.Errorf("AI provider %s is disabled", provider.ID)
}
providerKeys, err := p.resolveUserProviderAPIKeysForProvider(ctx, ownerID, provider)
if err != nil {
return "", chatprovider.ProviderAPIKeys{}, xerrors.Errorf("resolve provider API keys: %w", err)
}
return string(provider.Type), providerKeys, nil
}
func (p *Server) resolveChatModel(
ctx context.Context,
chat database.Chat,
modelOpts modelBuildOptions,
) (
model fantasy.LanguageModel,
dbConfig database.ChatModelConfig,
keys chatprovider.ProviderAPIKeys,
route resolvedModelRoute,
debugEnabled bool,
resolvedProvider string,
resolvedModel string,
@@ -8440,57 +8519,45 @@ func (p *Server) resolveChatModel(
) {
dbConfig, err = p.resolveModelConfig(ctx, chat)
if err != nil {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, false, "", "", xerrors.Errorf("resolve model config: %w", err)
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, resolvedModelRoute{}, false, "", "", xerrors.Errorf("resolve model config: %w", err)
}
if !dbConfig.Enabled {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, false, "", "", xerrors.Errorf("chat model config %s is disabled", dbConfig.ID)
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, resolvedModelRoute{}, false, "", "", xerrors.Errorf("chat model config %s is disabled", dbConfig.ID)
}
providerHint := dbConfig.Provider
var keyErr error
if dbConfig.AIProviderID.Valid {
provider, err := p.db.GetAIProviderByID(ctx, dbConfig.AIProviderID.UUID)
if err != nil {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, false, "", "", xerrors.Errorf("get AI provider: %w", err)
}
if !provider.Enabled {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, false, "", "", xerrors.Errorf("AI provider %s is disabled", provider.ID)
}
providerHint = string(provider.Type)
keys, keyErr = p.resolveUserProviderAPIKeysForProvider(ctx, chat.OwnerID, provider)
} else {
keys, keyErr = p.resolveUserProviderAPIKeys(ctx, chat.OwnerID, uuid.Nil)
}
if keyErr != nil {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, false, "", "", xerrors.Errorf("resolve provider API keys: %w", keyErr)
route, err = p.resolveModelRouteForConfig(ctx, chat.OwnerID, dbConfig, chatprovider.ProviderAPIKeys{})
if err != nil {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, resolvedModelRoute{}, false, "", "", err
}
keys = route.directProviderKeys()
providerHint, err := route.providerHint()
if err != nil {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, resolvedModelRoute{}, false, "", "", err
}
resolvedProvider, resolvedModel, err = chatprovider.ResolveModelWithProviderHint(
dbConfig.Model,
providerHint,
)
if err != nil {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, false, "", "", xerrors.Errorf(
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, resolvedModelRoute{}, false, "", "", xerrors.Errorf(
"resolve model metadata: %w", err,
)
}
model, debugEnabled, err = p.newDebugAwareModelFromConfig(
ctx,
chat,
providerHint,
dbConfig.Model,
keys,
chatprovider.UserAgent(),
chatprovider.CoderHeaders(chat),
)
model, debugEnabled, err = p.newDebugAwareModel(ctx, modelClientRequest{
Chat: chat,
ModelName: dbConfig.Model,
UserAgent: chatprovider.UserAgent(),
ExtraHeaders: chatprovider.CoderHeaders(chat),
}, route, modelOpts)
if err != nil {
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, false, "", "", xerrors.Errorf(
return nil, database.ChatModelConfig{}, chatprovider.ProviderAPIKeys{}, resolvedModelRoute{}, false, "", "", xerrors.Errorf(
"create model: %w", err,
)
}
return model, dbConfig, keys, debugEnabled, resolvedProvider, resolvedModel, nil
return model, dbConfig, keys, route, debugEnabled, resolvedProvider, resolvedModel, nil
}
func (p *Server) aiProviderConfig(ctx context.Context, provider database.AIProvider) (chatprovider.ConfiguredProvider, error) {
@@ -8605,9 +8672,18 @@ func (p *Server) resolveUserProviderAPIKeysForProviderType(
ownerID uuid.UUID,
providerType string,
) (chatprovider.ProviderAPIKeys, error) {
keys, _, err := p.resolveUserProviderAPIKeysAndProviderForProviderType(ctx, ownerID, providerType)
return keys, err
}
func (p *Server) resolveUserProviderAPIKeysAndProviderForProviderType(
ctx context.Context,
ownerID uuid.UUID,
providerType string,
) (chatprovider.ProviderAPIKeys, *database.AIProvider, error) {
providers, err := p.db.GetAIProviders(ctx, database.GetAIProvidersParams{})
if err != nil {
return chatprovider.ProviderAPIKeys{}, xerrors.Errorf("get enabled AI providers: %w", err)
return chatprovider.ProviderAPIKeys{}, nil, xerrors.Errorf("get enabled AI providers: %w", err)
}
normalizedProviderType := chatprovider.NormalizeProvider(providerType)
for _, provider := range providers {
@@ -8616,13 +8692,17 @@ func (p *Server) resolveUserProviderAPIKeysForProviderType(
}
keys, err := p.resolveUserProviderAPIKeysForProvider(ctx, ownerID, provider)
if err != nil {
return chatprovider.ProviderAPIKeys{}, err
return chatprovider.ProviderAPIKeys{}, nil, err
}
if userCanUseProviderKeys(keys, normalizedProviderType) {
return keys, nil
return keys, &provider, nil
}
}
return p.resolveUserProviderAPIKeys(ctx, ownerID, uuid.Nil)
keys, err := p.resolveUserProviderAPIKeys(ctx, ownerID, uuid.Nil)
if err != nil {
return chatprovider.ProviderAPIKeys{}, nil, err
}
return keys, nil, nil
}
func (p *Server) resolveUserProviderAPIKeys(
@@ -9391,6 +9471,7 @@ func insertSyntheticToolResultsTx(
params := database.InsertChatMessagesParams{
ChatID: chat.ID,
CreatedBy: make([]uuid.UUID, n),
APIKeyID: make([]string, n),
ModelConfigID: make([]uuid.UUID, n),
Role: make([]database.ChatMessageRole, n),
Content: make([]string, n),
@@ -9537,7 +9618,7 @@ func (p *Server) generateFinalTurnStatusLabel(
return fallbackTurnStatusLabel(status)
}
statusLabel := generateTurnStatusLabel(
statusLabel := p.generateTurnStatusLabel(
ctx,
chat,
status,
@@ -9545,7 +9626,9 @@ func (p *Server) generateFinalTurnStatusLabel(
runResult.FallbackProvider,
runResult.FallbackModel,
runResult.StatusLabelModel,
runResult.FallbackRoute,
runResult.ProviderKeys,
runResult.ModelBuildOptions,
logger,
p.existingDebugService(),
runResult.TriggerMessageID,