mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
+187
-42
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user