feat: add general subagent model override (#24610)

Adds a deployment-wide admin override for general delegated subagents.

## What changed
- store the general override in `site_configs` and expose it through the
shared `agent-model-override/{context}` API
- apply the general override when spawning delegated general subagents,
while preserving the existing Explore override behavior
- reuse a shared Agents settings form for the general and Explore
override sections

## Validation
- `make gen`
- `go test ./coderd -run 'TestChatModelOverrides'`
- `go test ./coderd/x/chatd -run
'TestSpawnAgent_(GeneralUsesConfiguredModelOverride|GeneralOverrideLogsAndFallsBackWhenCredentialsUnavailable|GeneralOverrideLogsAndFallsBackWhenProviderDisabled)'`
- `pnpm -C site lint:types`
- `pnpm -C site test:storybook --
AgentSettingsAgentsPageView.stories.tsx`
- `make lint`
- `make pre-commit`

> Mux is acting on Mike's behalf.
This commit is contained in:
Michael Suchacz
2026-04-24 12:37:20 +02:00
committed by GitHub
parent 4505278a9f
commit 3d90546aae
23 changed files with 1780 additions and 677 deletions
+187 -42
View File
@@ -4,6 +4,7 @@ import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"slices"
"sort"
@@ -27,6 +28,18 @@ import (
var ErrSubagentNotDescendant = xerrors.New("target chat is not a descendant of current chat")
var errInvalidModelOverrideMetadata = xerrors.New("invalid model override metadata")
type modelOverrideConfigResolver func(
context.Context,
uuid.UUID,
) (database.ChatModelConfig, string, error)
type modelOverrideProviderKeysResolver func(
context.Context,
uuid.UUID,
) (chatprovider.ProviderAPIKeys, error)
const (
subagentAwaitPollInterval = 200 * time.Millisecond
subagentAwaitFallbackPoll = 5 * time.Second
@@ -90,66 +103,199 @@ func (p *Server) isDesktopEnabled(ctx context.Context) bool {
return enabled
}
func (p *Server) resolveExploreSubagentModelConfigID(
func subagentModelOverrideLogLabel(
overrideContext codersdk.ChatAgentModelOverrideContext,
) string {
switch overrideContext {
case codersdk.ChatAgentModelOverrideContextGeneral:
return "general delegated child"
case codersdk.ChatAgentModelOverrideContextExplore:
return "explore"
default:
return string(overrideContext)
}
}
func readSubagentModelOverride(
ctx context.Context,
ownerID uuid.UUID,
fallback uuid.UUID,
) (uuid.UUID, error) {
//nolint:gocritic // Chatd needs its scoped deployment-config read access here.
chatdCtx := dbauthz.AsChatd(ctx)
raw, err := p.db.GetChatExploreModelOverride(chatdCtx)
if err != nil {
return uuid.Nil, xerrors.Errorf("get Explore model override: %w", err)
}
trimmed := strings.TrimSpace(raw)
if trimmed == "" {
return fallback, nil
}
configuredModelConfigID, err := uuid.Parse(trimmed)
if err != nil {
p.logger.Warn(ctx,
"invalid Explore model override, falling back to current turn model",
slog.F("raw_model_config_id", trimmed),
slog.Error(err),
db database.Store,
overrideContext codersdk.ChatAgentModelOverrideContext,
) (string, error) {
switch overrideContext {
case codersdk.ChatAgentModelOverrideContextGeneral:
return db.GetChatGeneralModelOverride(ctx)
case codersdk.ChatAgentModelOverrideContextExplore:
return db.GetChatExploreModelOverride(ctx)
default:
return "", xerrors.Errorf(
"unknown subagent model override context %q",
overrideContext,
)
return fallback, nil
}
modelConfig, err := p.db.GetEnabledChatModelConfigByID(
chatdCtx,
configuredModelConfigID,
)
if err != nil {
if xerrors.Is(err, sql.ErrNoRows) {
p.logger.Warn(ctx,
"explore model override is unavailable, falling back to current turn model",
slog.F("model_config_id", configuredModelConfigID),
)
return fallback, nil
}
return uuid.Nil, xerrors.Errorf("get enabled chat model config by id: %w", err)
}
func validateModelConfigAndResolveProvider(
modelConfig database.ChatModelConfig,
) (database.ChatModelConfig, string, error) {
if !modelConfig.Enabled {
return database.ChatModelConfig{}, "", sql.ErrNoRows
}
providerName, _, err := chatprovider.ResolveModelWithProviderHint(
modelConfig.Model,
modelConfig.Provider,
)
if err != nil {
return uuid.Nil, xerrors.Errorf("resolve Explore model provider: %w", err)
return database.ChatModelConfig{}, "", xerrors.Errorf(
"%w: %v",
errInvalidModelOverrideMetadata,
err,
)
}
providerKeys, err := p.resolveUserProviderAPIKeys(ctx, ownerID)
return modelConfig, providerName, nil
}
func enabledProviderContainsName(
providers []database.ChatProvider,
providerName string,
) bool {
normalizedProviderName := chatprovider.NormalizeProvider(providerName)
for _, provider := range providers {
if chatprovider.NormalizeProvider(provider.Provider) == normalizedProviderName {
return true
}
}
return false
}
func (p *Server) resolveConfiguredModelOverride(
ctx context.Context,
overrideContext string,
raw string,
ownerID uuid.UUID,
resolveModelConfig modelOverrideConfigResolver,
resolveProviderKeys modelOverrideProviderKeysResolver,
) (database.ChatModelConfig, bool, error) {
trimmed := strings.TrimSpace(raw)
if trimmed == "" {
return database.ChatModelConfig{}, false, nil
}
configuredModelConfigID, err := uuid.Parse(trimmed)
if err != nil {
return uuid.Nil, xerrors.Errorf("resolve provider API keys: %w", err)
p.logger.Info(ctx,
"invalid model override, ignoring",
slog.F("override_context", overrideContext),
slog.F("raw_model_config_id", trimmed),
slog.Error(err),
)
return database.ChatModelConfig{}, false, nil
}
if providerKeys.APIKey(providerName) == "" {
p.logger.Warn(ctx,
"explore model override credentials are unavailable, falling back to current turn model",
modelConfig, providerName, err := resolveModelConfig(
ctx,
configuredModelConfigID,
)
if err != nil {
switch {
case xerrors.Is(err, sql.ErrNoRows):
p.logger.Info(ctx,
"model override is unavailable, ignoring",
slog.F("override_context", overrideContext),
slog.F("model_config_id", configuredModelConfigID),
)
case errors.Is(err, errInvalidModelOverrideMetadata):
p.logger.Info(ctx,
"model override metadata is invalid, ignoring",
slog.F("override_context", overrideContext),
slog.F("model_config_id", configuredModelConfigID),
slog.Error(err),
)
default:
p.logger.Warn(ctx,
"failed to resolve model override, ignoring",
slog.F("override_context", overrideContext),
slog.F("model_config_id", configuredModelConfigID),
slog.Error(err),
)
}
return database.ChatModelConfig{}, false, nil
}
providerKeys, err := resolveProviderKeys(ctx, ownerID)
if err != nil {
return database.ChatModelConfig{}, false, xerrors.Errorf(
"resolve provider API keys: %w",
err,
)
}
if providerKeys.APIKey(providerName) == "" &&
!(chatprovider.ProviderAllowsAmbientCredentials(providerName) &&
providerKeys.HasProvider(providerName)) {
p.logger.Info(ctx,
"model override credentials are unavailable, ignoring",
slog.F("override_context", overrideContext),
slog.F("model_config_id", configuredModelConfigID),
slog.F("provider", providerName),
)
return fallback, nil
return database.ChatModelConfig{}, false, nil
}
return modelConfig, true, nil
}
func (p *Server) resolveSubagentModelConfigID(
ctx context.Context,
ownerID uuid.UUID,
overrideContext codersdk.ChatAgentModelOverrideContext,
) (uuid.UUID, error) {
//nolint:gocritic // Chatd needs its scoped deployment-config read access here.
chatdCtx := dbauthz.AsChatd(ctx)
raw, err := readSubagentModelOverride(chatdCtx, p.db, overrideContext)
if err != nil {
return uuid.Nil, xerrors.Errorf(
"get %s model override: %w",
subagentModelOverrideLogLabel(overrideContext),
err,
)
}
modelConfig, ok, err := p.resolveConfiguredModelOverride(
ctx,
string(overrideContext),
raw,
ownerID,
p.resolveModelConfigAndNormalizedProvider,
p.resolveUserProviderAPIKeys,
)
if err != nil {
return uuid.Nil, err
}
if !ok {
return uuid.Nil, nil
}
return modelConfig.ID, nil
}
func (p *Server) resolveModelConfigAndNormalizedProvider(
ctx context.Context,
modelConfigID uuid.UUID,
) (database.ChatModelConfig, string, error) {
if modelConfigID == uuid.Nil {
return database.ChatModelConfig{}, "", sql.ErrNoRows
}
modelConfig, err := p.configCache.ModelConfigByID(ctx, modelConfigID)
if err != nil {
return database.ChatModelConfig{}, "", err
}
modelConfig, providerName, err := validateModelConfigAndResolveProvider(modelConfig)
if err != nil {
return database.ChatModelConfig{}, "", err
}
enabledProviders, err := p.configCache.EnabledProviders(ctx)
if err != nil {
return database.ChatModelConfig{}, "", err
}
if !enabledProviderContainsName(enabledProviders, providerName) {
return database.ChatModelConfig{}, "", sql.ErrNoRows
}
return modelConfig, providerName, nil
}
func (p *Server) subagentTools(
ctx context.Context,
currentChat func() database.Chat,
@@ -444,7 +590,6 @@ func (p *Server) loadSubagentSpawnParentChat(
if err := validateSubagentSpawnParent(parent); err != nil {
return database.Chat{}, err
}
reloadedParent, err := p.db.GetChatByID(ctx, parent.ID)
if err != nil {
p.logger.Warn(ctx, "failed to load parent chat for spawn_agent",
+20 -4
View File
@@ -9,6 +9,7 @@ import (
"golang.org/x/xerrors"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/codersdk"
)
const (
@@ -43,22 +44,37 @@ func allSubagentDefinitions() []subagentDefinition {
{
id: subagentTypeGeneral,
description: "delegated work that may inspect or modify workspace files",
buildOptions: func(_ context.Context, _ *Server, _ database.Chat, _ database.Chat, _ uuid.UUID, _ string) (childSubagentChatOptions, error) {
return childSubagentChatOptions{}, nil
buildOptions: func(ctx context.Context, p *Server, parent database.Chat, _ database.Chat, _ uuid.UUID, _ string) (childSubagentChatOptions, error) {
modelConfigID, err := p.resolveSubagentModelConfigID(
ctx,
parent.OwnerID,
codersdk.ChatAgentModelOverrideContextGeneral,
)
if err != nil {
return childSubagentChatOptions{}, err
}
options := childSubagentChatOptions{}
if modelConfigID != uuid.Nil {
options.modelConfigIDOverride = &modelConfigID
}
return options, nil
},
},
{
id: subagentTypeExplore,
description: "read-only discovery, code tracing, and system understanding",
buildOptions: func(ctx context.Context, p *Server, _ database.Chat, turnParent database.Chat, currentModelConfigID uuid.UUID, _ string) (childSubagentChatOptions, error) {
modelConfigID, err := p.resolveExploreSubagentModelConfigID(
modelConfigID, err := p.resolveSubagentModelConfigID(
ctx,
turnParent.OwnerID,
currentModelConfigID,
codersdk.ChatAgentModelOverrideContextExplore,
)
if err != nil {
return childSubagentChatOptions{}, err
}
if modelConfigID == uuid.Nil {
modelConfigID = currentModelConfigID
}
inheritedMCPServerIDs, err := p.resolveExploreToolSnapshot(
ctx,
turnParent,
+325 -5
View File
@@ -4,6 +4,7 @@ import (
"context"
"database/sql"
"encoding/json"
"sync"
"testing"
"time"
@@ -12,6 +13,7 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"cdr.dev/slog/v3"
"cdr.dev/slog/v3/sloggers/slogtest"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/database/dbauthz"
@@ -71,7 +73,14 @@ func newInternalTestServer(
ps pubsub.Pubsub,
keys chatprovider.ProviderAPIKeys,
) *Server {
return newInternalTestServerWithClock(t, db, ps, keys, nil)
return newInternalTestServerWithLoggerAndClock(
t,
db,
ps,
keys,
slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}),
nil,
)
}
func newInternalTestServerWithClock(
@@ -80,10 +89,37 @@ func newInternalTestServerWithClock(
ps pubsub.Pubsub,
keys chatprovider.ProviderAPIKeys,
clk quartz.Clock,
) *Server {
return newInternalTestServerWithLoggerAndClock(
t,
db,
ps,
keys,
slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}),
clk,
)
}
func newInternalTestServerWithLogger(
t *testing.T,
db database.Store,
ps pubsub.Pubsub,
keys chatprovider.ProviderAPIKeys,
logger slog.Logger,
) *Server {
return newInternalTestServerWithLoggerAndClock(t, db, ps, keys, logger, nil)
}
func newInternalTestServerWithLoggerAndClock(
t *testing.T,
db database.Store,
ps pubsub.Pubsub,
keys chatprovider.ProviderAPIKeys,
logger slog.Logger,
clk quartz.Clock,
) *Server {
t.Helper()
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
server := New(Config{
Logger: logger,
Database: db,
@@ -101,6 +137,35 @@ func newInternalTestServerWithClock(
return server
}
type subagentTestLogSink struct {
mu sync.Mutex
entries []slog.SinkEntry
}
func (s *subagentTestLogSink) LogEntry(_ context.Context, entry slog.SinkEntry) {
s.mu.Lock()
defer s.mu.Unlock()
s.entries = append(s.entries, entry)
}
func (*subagentTestLogSink) Sync() {}
func (s *subagentTestLogSink) entriesAtLevelWithMessage(
level slog.Level,
message string,
) []slog.SinkEntry {
s.mu.Lock()
defer s.mu.Unlock()
entries := make([]slog.SinkEntry, 0, len(s.entries))
for _, entry := range s.entries {
if entry.Level == level && entry.Message == message {
entries = append(entries, entry)
}
}
return entries
}
// seedInternalChatDeps inserts an OpenAI provider and model config
// into the database and returns the created user, organization,
// and model. This deliberately does NOT create an Anthropic
@@ -218,6 +283,54 @@ func insertInternalChatModelConfig(
userID uuid.UUID,
model string,
enabled bool,
) database.ChatModelConfig {
return insertInternalChatModelConfigForProvider(
ctx,
t,
db,
userID,
"openai",
model,
enabled,
)
}
func insertInternalChatProvider(
ctx context.Context,
t *testing.T,
db database.Store,
userID uuid.UUID,
provider string,
apiKey string,
centralAPIKeyEnabled bool,
allowUserAPIKey bool,
allowCentralAPIKeyFallback bool,
) database.ChatProvider {
t.Helper()
providerConfig, err := db.InsertChatProvider(ctx, database.InsertChatProviderParams{
Provider: provider,
DisplayName: provider,
APIKey: apiKey,
CreatedBy: uuid.NullUUID{UUID: userID, Valid: true},
Enabled: true,
CentralApiKeyEnabled: centralAPIKeyEnabled,
AllowUserApiKey: allowUserAPIKey,
AllowCentralApiKeyFallback: allowCentralAPIKeyFallback,
})
require.NoError(t, err)
return providerConfig
}
func insertInternalChatModelConfigForProvider(
ctx context.Context,
t *testing.T,
db database.Store,
userID uuid.UUID,
provider string,
model string,
enabled bool,
) database.ChatModelConfig {
t.Helper()
return insertInternalChatModelConfigWithOptions(
@@ -225,6 +338,7 @@ func insertInternalChatModelConfig(
t,
db,
userID,
provider,
model,
enabled,
json.RawMessage(`{}`),
@@ -236,6 +350,7 @@ func insertInternalChatModelConfigWithOptions(
t *testing.T,
db database.Store,
userID uuid.UUID,
provider string,
model string,
enabled bool,
options json.RawMessage,
@@ -243,7 +358,7 @@ func insertInternalChatModelConfigWithOptions(
t.Helper()
modelConfig, err := db.InsertChatModelConfig(ctx, database.InsertChatModelConfigParams{
Provider: "openai",
Provider: provider,
Model: model,
DisplayName: model,
CreatedBy: uuid.NullUUID{UUID: userID, Valid: true},
@@ -574,6 +689,213 @@ func TestSpawnAgent_GeneralInheritsParentModelWhenOmitted(t *testing.T) {
require.Equal(t, parentChat.LastModelConfigID, childChat.LastModelConfigID)
}
func TestSpawnAgent_GeneralUsesConfiguredModelOverride(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{})
ctx := chatdTestContext(t)
user, org, model := seedInternalChatDeps(ctx, t, db)
overrideModel := insertInternalChatModelConfig(
ctx, t, db, user.ID, "general-override-"+uuid.NewString(), true,
)
require.NoError(t, db.UpsertChatGeneralModelOverride(ctx, overrideModel.ID.String()))
parentChat := createInternalParentChat(
ctx, t, server, db, org.ID, user.ID, model.ID, "parent-general-override",
)
resp := runSpawnAgentTool(ctx, t, server, parentChat, spawnAgentArgs{
Type: subagentTypeGeneral,
Prompt: "delegate general work",
})
childID := requireSpawnAgentChildChatID(t, resp)
childChat, err := db.GetChatByID(ctx, childID)
require.NoError(t, err)
require.Equal(t, overrideModel.ID, childChat.LastModelConfigID)
require.False(t, childChat.PlanMode.Valid)
}
func TestSpawnAgent_GeneralOverrideLogsAndFallsBackWhenCredentialsUnavailable(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
logSink := &subagentTestLogSink{}
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).AppendSinks(logSink)
server := newInternalTestServerWithLogger(t, db, ps, chatprovider.ProviderAPIKeys{}, logger)
ctx := chatdTestContext(t)
user, org, model := seedInternalChatDeps(ctx, t, db)
insertInternalChatProvider(
ctx,
t,
db,
user.ID,
"openai-compat",
"",
false,
true,
false,
)
overrideModel := insertInternalChatModelConfigForProvider(
ctx,
t,
db,
user.ID,
"openai-compat",
"gpt-4o-mini",
true,
)
require.NoError(t, db.UpsertChatGeneralModelOverride(ctx, overrideModel.ID.String()))
parent, err := server.CreateChat(ctx, CreateOptions{
OrganizationID: org.ID,
OwnerID: user.ID,
Title: "parent-general-credentials-fallback",
ModelConfigID: model.ID,
InitialUserContent: []codersdk.ChatMessagePart{
codersdk.ChatMessageText("delegate work"),
},
})
require.NoError(t, err)
parentChat, err := db.GetChatByID(ctx, parent.ID)
require.NoError(t, err)
resp := runSpawnAgentTool(ctx, t, server, parentChat, spawnAgentArgs{
Type: subagentTypeGeneral,
Prompt: "inspect provider credentials",
})
childID := requireSpawnAgentChildChatID(t, resp)
childChat, err := db.GetChatByID(ctx, childID)
require.NoError(t, err)
require.Equal(t, model.ID, childChat.LastModelConfigID)
require.False(t, childChat.PlanMode.Valid)
require.Len(t, logSink.entriesAtLevelWithMessage(
slog.LevelInfo,
"model override credentials are unavailable, ignoring",
), 1)
}
func TestSpawnAgent_GeneralOverrideLogsAndFallsBackWhenProviderDisabled(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
logSink := &subagentTestLogSink{}
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).AppendSinks(logSink)
server := newInternalTestServerWithLogger(
t,
db,
ps,
chatprovider.ProviderAPIKeys{
ByProvider: map[string]string{
"openai-compat": "fallback-key",
},
},
logger,
)
ctx := chatdTestContext(t)
user, org, model := seedInternalChatDeps(ctx, t, db)
_, err := db.InsertChatProvider(ctx, database.InsertChatProviderParams{
Provider: "openai-compat",
DisplayName: "openai-compat",
APIKey: "",
CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true},
Enabled: false,
CentralApiKeyEnabled: false,
AllowUserApiKey: true,
AllowCentralApiKeyFallback: false,
})
require.NoError(t, err)
overrideModel := insertInternalChatModelConfigForProvider(
ctx,
t,
db,
user.ID,
"openai-compat",
"gpt-4o-mini",
true,
)
require.NoError(t, db.UpsertChatGeneralModelOverride(ctx, overrideModel.ID.String()))
parent, err := server.CreateChat(ctx, CreateOptions{
OrganizationID: org.ID,
OwnerID: user.ID,
Title: "parent-general-disabled-provider-fallback",
ModelConfigID: model.ID,
InitialUserContent: []codersdk.ChatMessagePart{
codersdk.ChatMessageText("delegate work"),
},
})
require.NoError(t, err)
parentChat, err := db.GetChatByID(ctx, parent.ID)
require.NoError(t, err)
resp := runSpawnAgentTool(ctx, t, server, parentChat, spawnAgentArgs{
Type: subagentTypeGeneral,
Prompt: "inspect disabled providers",
})
childID := requireSpawnAgentChildChatID(t, resp)
childChat, err := db.GetChatByID(ctx, childID)
require.NoError(t, err)
require.Equal(t, model.ID, childChat.LastModelConfigID)
require.False(t, childChat.PlanMode.Valid)
require.Len(t, logSink.entriesAtLevelWithMessage(
slog.LevelInfo,
"model override is unavailable, ignoring",
), 1)
}
func TestResolveConfiguredModelOverride_AcceptsAmbientCredentialsProvider(
t *testing.T,
) {
t.Parallel()
logSink := &subagentTestLogSink{}
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).AppendSinks(logSink)
server := &Server{logger: logger}
ctx := chatdTestContext(t)
ownerID := uuid.New()
modelConfig := database.ChatModelConfig{
ID: uuid.New(),
Provider: "bedrock",
Model: "anthropic.claude-haiku-4-5-20251001-v1:0",
DisplayName: "Ambient Bedrock Override",
Enabled: true,
}
resolvedModelConfig, ok, err := server.resolveConfiguredModelOverride(
ctx,
"plan",
modelConfig.ID.String(),
ownerID,
func(
_ context.Context,
configuredModelConfigID uuid.UUID,
) (database.ChatModelConfig, string, error) {
require.Equal(t, modelConfig.ID, configuredModelConfigID)
return modelConfig, "bedrock", nil
},
func(
_ context.Context,
resolvedOwnerID uuid.UUID,
) (chatprovider.ProviderAPIKeys, error) {
require.Equal(t, ownerID, resolvedOwnerID)
return chatprovider.ProviderAPIKeys{
ByProvider: map[string]string{"bedrock": ""},
}, nil
},
)
require.NoError(t, err)
require.True(t, ok)
require.Equal(t, modelConfig, resolvedModelConfig)
require.Empty(t, logSink.entriesAtLevelWithMessage(
slog.LevelInfo,
"model override credentials are unavailable, ignoring",
))
}
func TestCreateChildSubagentChat_OverrideWorksWhenParentHasNoModel(t *testing.T) {
t.Parallel()
@@ -1328,7 +1650,6 @@ func TestSubagentLifecycleToolsIncludePersistedSubagentTypeAcrossVariants(t *tes
}
for _, tt := range tests {
tt := tt
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
@@ -1450,7 +1771,6 @@ func TestSubagentLifecycleToolErrorsIncludePersistedSubagentType(t *testing.T) {
}
for _, tt := range tests {
tt := tt
t.Run(tt.name, func(t *testing.T) {
t.Parallel()