fix: hide and reject chat models from disabled AI providers (#27070)

This commit is contained in:
Michael Suchacz
2026-07-19 08:33:40 +02:00
committed by GitHub
parent 34ed124478
commit 9b3af629cd
23 changed files with 1127 additions and 84 deletions
+60 -24
View File
@@ -1120,7 +1120,8 @@ func (c *turnWorkspaceContext) getWorkspaceConn(ctx context.Context) (workspaces
type AgentConnFunc func(ctx context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error)
var (
// ErrInvalidModelConfigID indicates the requested model config does not exist.
// ErrInvalidModelConfigID indicates the requested model config does not
// exist, is disabled, or its provider is disabled.
ErrInvalidModelConfigID = xerrors.New("invalid model config ID")
// ErrEditedMessageNotFound indicates the edited message does not exist
// in the target chat.
@@ -1333,6 +1334,12 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C
initialMessages = append(initialMessages, systemMessage(workspaceAwarenessContent, opts.ModelConfigID))
initialMessages = append(initialMessages, userMessageWithAPIKeyID(userContent, opts.ModelConfigID, opts.OwnerID, apiKeyID, opts.ReasoningEffort))
if opts.ModelConfigID != uuid.Nil {
if err := requireEnabledChatModelConfig(ctx, p.db, opts.ModelConfigID); err != nil {
return database.Chat{}, err
}
}
result, err := chatstate.CreateChat(ctx, p.db, p.pubsub, chatstate.CreateChatInput{
OrganizationID: opts.OrganizationID,
OwnerID: opts.OwnerID,
@@ -1562,22 +1569,35 @@ func resolveSendMessageModelConfigID(
return resolveFallbackModelConfigID(ctx, store, chat.LastModelConfigID)
}
if err := requireEnabledChatModelConfig(ctx, store, requested); err != nil {
return uuid.Nil, err
}
return requested, nil
}
// requireEnabledChatModelConfig rechecks enabled state inside the daemon:
// the coderd preflight can race an admin disabling the model or provider.
func requireEnabledChatModelConfig(
ctx context.Context,
store database.Store,
modelConfigID uuid.UUID,
) error {
chatdCtx := chatdModelConfigLookupContext(ctx)
if _, err := store.GetChatModelConfigByID(chatdCtx, requested); err != nil {
if _, err := store.GetEnabledChatModelConfigByID(chatdCtx, modelConfigID); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return uuid.Nil, xerrors.Errorf(
return xerrors.Errorf(
"%w: %s",
ErrInvalidModelConfigID,
requested,
modelConfigID,
)
}
return uuid.Nil, xerrors.Errorf(
return xerrors.Errorf(
"get requested model config %s: %w",
requested,
modelConfigID,
err,
)
}
return requested, nil
return nil
}
func resolveFallbackModelConfigID(
@@ -1587,7 +1607,7 @@ func resolveFallbackModelConfigID(
) (uuid.UUID, error) {
chatdCtx := chatdModelConfigLookupContext(ctx)
if modelConfigID != uuid.Nil {
if _, err := store.GetChatModelConfigByID(chatdCtx, modelConfigID); err == nil {
if _, err := store.GetEnabledChatModelConfigByID(chatdCtx, modelConfigID); err == nil {
return modelConfigID, nil
} else if !errors.Is(err, sql.ErrNoRows) {
return uuid.Nil, xerrors.Errorf(
@@ -1605,6 +1625,21 @@ func resolveFallbackModelConfigID(
}
return uuid.Nil, xerrors.Errorf("get default chat model config: %w", err)
}
// The default may itself be disabled or under a disabled provider.
if _, err := store.GetEnabledChatModelConfigByID(chatdCtx, defaultConfig.ID); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return uuid.Nil, xerrors.Errorf(
"%w: default model config %s or its provider is disabled",
ErrNoDefaultChatModelConfig,
defaultConfig.ID,
)
}
return uuid.Nil, xerrors.Errorf(
"get default chat model config %s: %w",
defaultConfig.ID,
err,
)
}
return defaultConfig.ID, nil
}
@@ -1677,24 +1712,25 @@ func (p *Server) EditMessage(
// foreign-key error from the message-insert path.
var modelOverride uuid.NullUUID
if opts.ModelConfigID != uuid.Nil {
if _, err := store.GetChatModelConfigByID(
chatdModelConfigLookupContext(ctx),
opts.ModelConfigID,
); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return xerrors.Errorf(
"%w: %s",
ErrInvalidModelConfigID,
opts.ModelConfigID,
)
}
return xerrors.Errorf(
"get requested model config %s: %w",
opts.ModelConfigID,
err,
)
if err := requireEnabledChatModelConfig(ctx, store, opts.ModelConfigID); err != nil {
return err
}
modelOverride = uuid.NullUUID{UUID: opts.ModelConfigID, Valid: true}
} else {
// Without an explicit override the transition preserves
// the edited message's original model, which may have been
// disabled since; resolve it like a normal message send.
preserved := uuid.Nil
if target.ModelConfigID.Valid {
preserved = target.ModelConfigID.UUID
}
resolved, err := resolveFallbackModelConfigID(ctx, store, preserved)
if err != nil {
return err
}
if resolved != preserved {
modelOverride = uuid.NullUUID{UUID: resolved, Valid: true}
}
}
var reasoningEffortOverride database.NullChatReasoningEffort
+123
View File
@@ -3485,3 +3485,126 @@ func TestPrepareManualTitleDebugRun_RouteFailureDerivesProviderFromConfig(t *tes
require.True(t, gotProvider.Valid, "debug run provider should be populated from the linked config")
require.Equal(t, "anthropic", gotProvider.String)
}
// TestResolveFallbackModelConfigID verifies that admission does not reuse
// a disabled last model and rejects a disabled default.
func TestResolveFallbackModelConfigID(t *testing.T) {
t.Parallel()
newProvider := func(t *testing.T, db database.Store, enabled bool) database.AIProvider {
return dbgen.AIProvider(t, db, database.AIProvider{}, func(p *database.InsertAIProviderParams) {
p.Enabled = enabled
})
}
newModelConfig := func(t *testing.T, db database.Store, providerID uuid.UUID, isDefault bool) database.ChatModelConfig {
return dbgen.ChatModelConfig(t, db, database.ChatModelConfig{
AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true},
IsDefault: isDefault,
})
}
t.Run("EnabledLastModel", func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitShort)
provider := newProvider(t, db, true)
lastModel := newModelConfig(t, db, provider.ID, false)
resolved, err := resolveFallbackModelConfigID(ctx, db, lastModel.ID)
require.NoError(t, err)
require.Equal(t, lastModel.ID, resolved)
})
t.Run("ProviderDisabledLastModelFallsBackToDefault", func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitShort)
disabledProvider := newProvider(t, db, false)
lastModel := newModelConfig(t, db, disabledProvider.ID, false)
enabledProvider := newProvider(t, db, true)
defaultModel := newModelConfig(t, db, enabledProvider.ID, true)
resolved, err := resolveFallbackModelConfigID(ctx, db, lastModel.ID)
require.NoError(t, err)
require.Equal(t, defaultModel.ID, resolved)
})
t.Run("NilLastModelUsesDefault", func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitShort)
provider := newProvider(t, db, true)
defaultModel := newModelConfig(t, db, provider.ID, true)
resolved, err := resolveFallbackModelConfigID(ctx, db, uuid.Nil)
require.NoError(t, err)
require.Equal(t, defaultModel.ID, resolved)
})
t.Run("ProviderDisabledDefaultRejected", func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitShort)
disabledProvider := newProvider(t, db, false)
lastModel := newModelConfig(t, db, disabledProvider.ID, false)
newModelConfig(t, db, disabledProvider.ID, true)
_, err := resolveFallbackModelConfigID(ctx, db, lastModel.ID)
require.ErrorIs(t, err, ErrNoDefaultChatModelConfig)
})
t.Run("ExplicitEnabledModel", func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitShort)
provider := newProvider(t, db, true)
model := newModelConfig(t, db, provider.ID, false)
resolved, err := resolveSendMessageModelConfigID(ctx, db, database.Chat{}, model.ID)
require.NoError(t, err)
require.Equal(t, model.ID, resolved)
})
// An explicit model whose provider was disabled after the coderd
// preflight must still be rejected inside the daemon.
t.Run("ExplicitProviderDisabledRejected", func(t *testing.T) {
t.Parallel()
db, _ := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitShort)
disabledProvider := newProvider(t, db, false)
model := newModelConfig(t, db, disabledProvider.ID, false)
_, err := resolveSendMessageModelConfigID(ctx, db, database.Chat{}, model.ID)
require.ErrorIs(t, err, ErrInvalidModelConfigID)
})
// The create path performs the same daemon-side recheck before
// inserting the chat and its initial messages.
t.Run("CreateChatProviderDisabledRejected", func(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
ctx := testutil.Context(t, testutil.WaitShort)
disabledProvider := newProvider(t, db, false)
model := newModelConfig(t, db, disabledProvider.ID, false)
server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{})
_, err := server.CreateChat(ctx, CreateOptions{
OrganizationID: uuid.New(),
OwnerID: uuid.New(),
Title: "provider disabled create",
ModelConfigID: model.ID,
APIKeyID: "test-api-key-id",
InitialUserContent: []codersdk.ChatMessagePart{
codersdk.ChatMessageText("hello"),
},
})
require.ErrorIs(t, err, ErrInvalidModelConfigID)
})
}