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
+89 -11
View File
@@ -1192,6 +1192,21 @@ func (api *API) validateUserChatModelConfigAvailable(
}
}
// validateExplicitChatModelConfigAvailable validates a caller-supplied
// model config ID. A nil ID keeps the chat's current model and is
// validated by the daemon's fallback resolution instead.
func (api *API) validateExplicitChatModelConfigAvailable(
ctx context.Context,
userID uuid.UUID,
modelConfigID uuid.UUID,
) (int, *codersdk.Response) {
if modelConfigID == uuid.Nil {
return 0, nil
}
_, status, resp := api.validateUserChatModelConfigAvailable(ctx, userID, modelConfigID)
return status, resp
}
// EXPERIMENTAL: this endpoint is experimental and is subject to change.
//
// @Summary Create chat
@@ -1429,6 +1444,13 @@ func (api *API) postChats(rw http.ResponseWriter, r *http.Request) {
if maybeWriteLimitErr(ctx, rw, err) {
return
}
if xerrors.Is(err, chatd.ErrInvalidModelConfigID) {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
Message: "Invalid model config ID.",
Detail: err.Error(),
})
return
}
if database.IsForeignKeyViolation(
err,
database.ForeignKeyChatsLastModelConfigID,
@@ -3334,6 +3356,10 @@ func (api *API) postChatMessages(rw http.ResponseWriter, r *http.Request) {
if req.ModelConfigID != nil {
modelConfigID = *req.ModelConfigID
}
if status, resp := api.validateExplicitChatModelConfigAvailable(ctx, apiKey.UserID, modelConfigID); resp != nil {
httpapi.Write(ctx, rw, status, *resp)
return
}
reasoningEffort := req.ReasoningEffort
if reasoningEffort != nil && !chatprovider.IsValidReasoningEffort(*reasoningEffort) {
@@ -3382,6 +3408,12 @@ func (api *API) postChatMessages(rw http.ResponseWriter, r *http.Request) {
})
return
}
if xerrors.Is(sendErr, chatd.ErrNoDefaultChatModelConfig) {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
Message: "No default chat model config is configured.",
})
return
}
if errors.Is(sendErr, chatstate.ErrChatNotFound) {
httpapi.ResourceNotFound(rw)
return
@@ -3498,6 +3530,10 @@ func (api *API) patchChatMessage(rw http.ResponseWriter, r *http.Request) {
if req.ModelConfigID != nil {
editModelConfigID = *req.ModelConfigID
}
if status, resp := api.validateExplicitChatModelConfigAvailable(ctx, apiKey.UserID, editModelConfigID); resp != nil {
httpapi.Write(ctx, rw, status, *resp)
return
}
editReasoningEffort := req.ReasoningEffort
if editReasoningEffort != nil && !chatprovider.IsValidReasoningEffort(*editReasoningEffort) {
@@ -3536,6 +3572,10 @@ func (api *API) patchChatMessage(rw http.ResponseWriter, r *http.Request) {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
Message: "Invalid model config ID.",
})
case xerrors.Is(editErr, chatd.ErrNoDefaultChatModelConfig):
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
Message: "No default chat model config is configured.",
})
case errors.Is(editErr, chatstate.ErrChatNotFound):
httpapi.ResourceNotFound(rw)
case writeChatInvalidState(ctx, rw, editErr):
@@ -4812,6 +4852,9 @@ func (api *API) resolveCreateChatModelConfigID(
Message: "Invalid model config ID.",
}
}
if _, status, resp := api.validateUserChatModelConfigAvailable(ctx, userID, *req.ModelConfigID); resp != nil {
return uuid.Nil, nil, status, resp
}
return *req.ModelConfigID, nil, 0, nil
}
@@ -4906,6 +4949,21 @@ func (api *API) defaultCreateChatModelConfigID(
}
}
// The resolved default may itself be disabled or under a disabled
// provider.
if _, err := lookupEnabledChatModelConfigByID(ctx, api.Database, defaultModelConfig.ID); err != nil {
if xerrors.Is(err, sql.ErrNoRows) {
return uuid.Nil, http.StatusBadRequest, &codersdk.Response{
Message: "No default chat model config is configured.",
Detail: "The default chat model or its provider is disabled.",
}
}
return uuid.Nil, http.StatusInternalServerError, &codersdk.Response{
Message: "Failed to resolve chat model config.",
Detail: err.Error(),
}
}
return defaultModelConfig.ID, 0, nil
}
@@ -5697,16 +5755,11 @@ func (api *API) putChatAdvisorConfig(rw http.ResponseWriter, r *http.Request) {
return
}
} else {
// Use system context because GetChatModelConfigByID requires
// deployment-config read access, which can be broader than the
// handler's explicit update check. The lookup validates the model and
// any selected reasoning effort before persisting deployment config.
//nolint:gocritic // This admin-authorized validation lookup intentionally bypasses read authz.
modelConfig, err := api.Database.GetChatModelConfigByID(dbauthz.AsSystemRestricted(ctx), req.ModelConfigID)
modelConfig, err := lookupEnabledChatModelConfigByID(ctx, api.Database, req.ModelConfigID)
if err != nil {
if errors.Is(err, sql.ErrNoRows) || httpapi.Is404Error(err) {
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
Message: fmt.Sprintf("model_config_id %q does not match any existing model config.", req.ModelConfigID),
Message: fmt.Sprintf("model_config_id %q does not match any enabled model config.", req.ModelConfigID),
})
return
}
@@ -7619,7 +7672,20 @@ func ensureDefaultChatModelConfig(
return nil
}
candidateConfig := modelConfigs[0]
// Prefer a config that can actually serve requests (enabled, under an
// enabled provider) so the promoted default does not reject
// omitted-model chat creation. Fall back to any non-excluded config
// when no usable candidate exists.
//nolint:gocritic // Candidate usability depends on deployment-wide provider state, not the caller's permissions.
enabledRows, err := tx.GetEnabledChatModelConfigs(dbauthz.AsChatd(ctx))
if err != nil {
return xerrors.Errorf("list enabled chat model configs: %w", err)
}
usable := make(map[uuid.UUID]struct{}, len(enabledRows))
for _, row := range enabledRows {
usable[row.ChatModelConfig.ID] = struct{}{}
}
excluded := make(map[uuid.UUID]struct{}, len(excludedConfigIDs))
for _, configID := range excludedConfigIDs {
if configID == uuid.Nil {
@@ -7627,12 +7693,24 @@ func ensureDefaultChatModelConfig(
}
excluded[configID] = struct{}{}
}
for _, config := range modelConfigs {
candidateConfig := modelConfigs[0]
var selected *database.ChatModelConfig
for i := range modelConfigs {
config := &modelConfigs[i]
if _, skip := excluded[config.ID]; skip {
continue
}
candidateConfig = config
break
if selected == nil {
selected = config
}
if _, ok := usable[config.ID]; ok {
selected = config
break
}
}
if selected != nil {
candidateConfig = *selected
}
if err := tx.UnsetDefaultChatModelConfigs(ctx); err != nil {
+385 -2
View File
@@ -517,6 +517,84 @@ func TestPostChats(t *testing.T) {
}
})
t.Run("DisabledModelConfigRejected", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client := newChatClient(t)
firstUser := coderdtest.CreateFirstUser(t, client.Client)
_ = createChatModelConfig(t, client)
disabledConfig := createDisabledChatModelConfig(
t,
client,
coderdtest.TestChatProviderOpenAICompat,
"gpt-4o-create-disabled-"+uuid.NewString(),
)
_, err := client.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: firstUser.OrganizationID,
Content: []codersdk.ChatInputPart{{
Type: codersdk.ChatInputPartTypeText,
Text: "hello",
}},
ModelConfigID: ptr.Ref(disabledConfig.ID),
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "Invalid model_config_id: model config not found or disabled.", sdkErr.Message)
})
t.Run("ProviderDisabledModelConfigRejected", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client := newChatClient(t)
firstUser := coderdtest.CreateFirstUser(t, client.Client)
_ = createChatModelConfig(t, client)
providerDisabledConfig := createProviderDisabledChatModelConfig(
t,
client,
"openai",
"gpt-4o-create-provider-disabled-"+uuid.NewString(),
)
_, err := client.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: firstUser.OrganizationID,
Content: []codersdk.ChatInputPart{{
Type: codersdk.ChatInputPartTypeText,
Text: "hello",
}},
ModelConfigID: ptr.Ref(providerDisabledConfig.ID),
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "Invalid model_config_id: provider is not enabled for this model.", sdkErr.Message)
})
t.Run("ProviderDisabledDefaultModelRejected", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client := newChatClient(t)
firstUser := coderdtest.CreateFirstUser(t, client.Client)
defaultConfig := createChatModelConfig(t, client)
_, err := client.UpdateAIProvider(ctx, defaultConfig.AIProviderID.String(), codersdk.UpdateAIProviderRequest{
Enabled: ptr.Ref(false),
})
require.NoError(t, err)
// Omitting model_config_id resolves the default model, whose
// provider is now disabled.
_, err = client.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: firstUser.OrganizationID,
Content: []codersdk.ChatInputPart{{
Type: codersdk.ChatInputPartTypeText,
Text: "hello",
}},
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "No default chat model config is configured.", sdkErr.Message)
require.Equal(t, "The default chat model or its provider is disabled.", sdkErr.Detail)
})
t.Run("WithPerChatSystemPrompt", func(t *testing.T) {
t.Parallel()
@@ -3847,6 +3925,39 @@ func TestListChatModelConfigs(t *testing.T) {
require.True(t, configs[0].Enabled)
})
// An enabled config under a disabled provider must stay visible to
// admins (management view) while being hidden from non-admins (usage
// view).
t.Run("ProviderDisabled", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
adminClient := newChatClient(t)
firstUser := coderdtest.CreateFirstUser(t, adminClient.Client)
enabledConfig := createChatModelConfig(t, adminClient)
providerDisabledConfig := createProviderDisabledChatModelConfig(
t,
adminClient,
"openai",
"gpt-4o-provider-disabled-"+uuid.NewString(),
)
memberClientRaw, _ := coderdtest.CreateAnotherUser(t, adminClient.Client, firstUser.OrganizationID)
memberClient := codersdk.NewExperimentalClient(memberClientRaw)
adminConfigs, err := adminClient.ListChatModelConfigs(ctx)
require.NoError(t, err)
adminIDs := make([]uuid.UUID, 0, len(adminConfigs))
for _, config := range adminConfigs {
adminIDs = append(adminIDs, config.ID)
}
require.Contains(t, adminIDs, providerDisabledConfig.ID)
memberConfigs, err := memberClient.ListChatModelConfigs(ctx)
require.NoError(t, err)
require.Len(t, memberConfigs, 1)
require.Equal(t, enabledConfig.ID, memberConfigs[0].ID)
})
t.Run("DeserializesLegacyPricingJSON", func(t *testing.T) {
t.Parallel()
@@ -4846,6 +4957,46 @@ func TestDeleteChatModelConfig(t *testing.T) {
}
})
// Deleting the default must not promote a config whose provider is
// disabled while a usable candidate exists.
t.Run("PromotesUsableConfigOverDisabledProvider", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client := newChatClient(t)
_ = coderdtest.CreateFirstUser(t, client.Client)
defaultConfig := createChatModelConfig(t, client)
// Same provider type as the enabled candidate with an
// alphabetically earlier model, so it sorts first in the
// reselection order.
createProviderDisabledChatModelConfig(
t,
client,
coderdtest.TestChatProviderOpenAICompat,
"a-provider-disabled-model",
)
enabledConfig := createAdditionalChatModelConfig(
t,
client,
coderdtest.TestChatProviderOpenAICompat,
"z-enabled-model",
)
err := client.DeleteChatModelConfig(ctx, defaultConfig.ID)
require.NoError(t, err)
configs, err := client.ListChatModelConfigs(ctx)
require.NoError(t, err)
defaultID := uuid.Nil
for _, config := range configs {
if config.IsDefault {
defaultID = config.ID
}
}
require.Equal(t, enabledConfig.ID, defaultID)
})
t.Run("NotFound", func(t *testing.T) {
t.Parallel()
@@ -6881,6 +7032,75 @@ func TestPostChatMessages(t *testing.T) {
}
})
t.Run("ProviderDisabledModelConfigRejected", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client := newChatClient(t)
firstUser := coderdtest.CreateFirstUser(t, client.Client)
_ = createChatModelConfig(t, client)
chat, err := client.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: firstUser.OrganizationID,
Content: []codersdk.ChatInputPart{{
Type: codersdk.ChatInputPartTypeText,
Text: "initial message before disabled provider switch",
}},
})
require.NoError(t, err)
providerDisabledConfig := createProviderDisabledChatModelConfig(
t,
client,
"openai",
"gpt-4o-send-provider-disabled-"+uuid.NewString(),
)
_, err = client.CreateChatMessage(ctx, chat.ID, codersdk.CreateChatMessageRequest{
Content: []codersdk.ChatInputPart{{
Type: codersdk.ChatInputPartTypeText,
Text: "switch to a provider-disabled model",
}},
ModelConfigID: ptr.Ref(providerDisabledConfig.ID),
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "Invalid model_config_id: provider is not enabled for this model.", sdkErr.Message)
})
t.Run("ProviderDisabledDefaultFallbackRejected", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client := newChatClient(t)
firstUser := coderdtest.CreateFirstUser(t, client.Client)
defaultConfig := createChatModelConfig(t, client)
chat, err := client.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: firstUser.OrganizationID,
Content: []codersdk.ChatInputPart{{
Type: codersdk.ChatInputPartTypeText,
Text: "initial message before provider disable",
}},
})
require.NoError(t, err)
_, err = client.UpdateAIProvider(ctx, defaultConfig.AIProviderID.String(), codersdk.UpdateAIProviderRequest{
Enabled: ptr.Ref(false),
})
require.NoError(t, err)
// Without an explicit model the fallback walks last model ->
// default, both under the now-disabled provider.
_, err = client.CreateChatMessage(ctx, chat.ID, codersdk.CreateChatMessageRequest{
Content: []codersdk.ChatInputPart{{
Type: codersdk.ChatInputPartTypeText,
Text: "message after provider disable",
}},
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "No default chat model config is configured.", sdkErr.Message)
})
t.Run("MemberWithoutAgentsAccess", func(t *testing.T) {
t.Parallel()
@@ -8961,7 +9181,97 @@ func TestPatchChatMessage(t *testing.T) {
ModelConfigID: &unknownID,
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "Invalid model config ID.", sdkErr.Message)
require.Equal(t, "Invalid model_config_id: model config not found or disabled.", sdkErr.Message)
})
t.Run("ProviderDisabledModelConfigID", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client := newChatClient(t)
firstUser := coderdtest.CreateFirstUser(t, client.Client)
_ = createChatModelConfig(t, client)
chat, err := client.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: firstUser.OrganizationID,
Content: []codersdk.ChatInputPart{{
Type: codersdk.ChatInputPartTypeText,
Text: "hello",
}},
})
require.NoError(t, err)
messagesResult, err := client.GetChatMessages(ctx, chat.ID, nil)
require.NoError(t, err)
var userMessageID int64
for _, message := range messagesResult.Messages {
if message.Role == codersdk.ChatMessageRoleUser {
userMessageID = message.ID
break
}
}
require.NotZero(t, userMessageID)
providerDisabledConfig := createProviderDisabledChatModelConfig(
t,
client,
"openai",
"gpt-4o-edit-provider-disabled-"+uuid.NewString(),
)
_, err = client.EditChatMessage(ctx, chat.ID, userMessageID, codersdk.EditChatMessageRequest{
Content: []codersdk.ChatInputPart{{
Type: codersdk.ChatInputPartTypeText,
Text: "edited with provider-disabled model",
}},
ModelConfigID: &providerDisabledConfig.ID,
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "Invalid model_config_id: provider is not enabled for this model.", sdkErr.Message)
})
t.Run("ProviderDisabledPreservedModelRejected", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client := newChatClient(t)
firstUser := coderdtest.CreateFirstUser(t, client.Client)
defaultConfig := createChatModelConfig(t, client)
chat, err := client.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: firstUser.OrganizationID,
Content: []codersdk.ChatInputPart{{
Type: codersdk.ChatInputPartTypeText,
Text: "hello before provider disable",
}},
})
require.NoError(t, err)
messagesResult, err := client.GetChatMessages(ctx, chat.ID, nil)
require.NoError(t, err)
var userMessageID int64
for _, message := range messagesResult.Messages {
if message.Role == codersdk.ChatMessageRoleUser {
userMessageID = message.ID
break
}
}
require.NotZero(t, userMessageID)
_, err = client.UpdateAIProvider(ctx, defaultConfig.AIProviderID.String(), codersdk.UpdateAIProviderRequest{
Enabled: ptr.Ref(false),
})
require.NoError(t, err)
// Editing without model_config_id preserves the edited message's
// original model; its provider and the default's are now disabled.
_, err = client.EditChatMessage(ctx, chat.ID, userMessageID, codersdk.EditChatMessageRequest{
Content: []codersdk.ChatInputPart{{
Type: codersdk.ChatInputPartTypeText,
Text: "edited after provider disable",
}},
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "No default chat model config is configured.", sdkErr.Message)
})
}
@@ -11892,6 +12202,25 @@ func createDisabledChatModelConfig(
return updated
}
// createProviderDisabledChatModelConfig creates an enabled model config,
// then disables its parent AI provider.
func createProviderDisabledChatModelConfig(
t *testing.T,
client *codersdk.ExperimentalClient,
provider string,
model string,
) codersdk.ChatModelConfig {
t.Helper()
modelConfig := createAdditionalChatModelConfig(t, client, provider, model)
ctx := testutil.Context(t, testutil.WaitLong)
_, err := client.UpdateAIProvider(ctx, modelConfig.AIProviderID.String(), codersdk.UpdateAIProviderRequest{
Enabled: ptr.Ref(false),
})
require.NoError(t, err)
return modelConfig
}
func enableUserChatProviderKey(
t testing.TB,
adminClient *codersdk.ExperimentalClient,
@@ -12707,6 +13036,20 @@ func TestChatModelOverrides(t *testing.T) {
require.Equal(t, "Invalid model_config_id.", sdkErr.Message)
})
t.Run("ProviderDisabledModelReturns400", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitLong)
providerDisabledModel := createProviderDisabledChatModelConfig(
t,
adminClient,
"openai",
"gpt-4.1-provider-disabled-"+string(setting.context),
)
err := putOverride(ctx, adminClient, setting.context, providerDisabledModel.ID.String())
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Equal(t, "Invalid model_config_id.", sdkErr.Message)
})
t.Run("UnknownModelReturns400", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitLong)
unknownModelID := uuid.New()
@@ -14204,7 +14547,47 @@ func TestChatAdvisorConfig_InvalidModelConfigID(t *testing.T) {
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Contains(t, sdkErr.Message, unknownID.String())
require.Contains(t, sdkErr.Message, "does not match any existing model config")
require.Contains(t, sdkErr.Message, "does not match any enabled model config")
}
func TestChatAdvisorConfig_DisabledModelConfigID(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
adminClient := newChatClient(t)
coderdtest.CreateFirstUser(t, adminClient.Client)
disabledConfig := createDisabledChatModelConfig(
t,
adminClient,
coderdtest.TestChatProviderOpenAICompat,
"gpt-4o-advisor-disabled-"+uuid.NewString(),
)
err := adminClient.UpdateChatAdvisorConfig(ctx, codersdk.UpdateAdvisorConfigRequest{
ModelConfigID: disabledConfig.ID,
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Contains(t, sdkErr.Message, "does not match any enabled model config")
}
func TestChatAdvisorConfig_ProviderDisabledModelConfigID(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
adminClient := newChatClient(t)
coderdtest.CreateFirstUser(t, adminClient.Client)
providerDisabledConfig := createProviderDisabledChatModelConfig(
t,
adminClient,
"openai",
"gpt-4o-advisor-provider-disabled-"+uuid.NewString(),
)
err := adminClient.UpdateChatAdvisorConfig(ctx, codersdk.UpdateAdvisorConfigRequest{
ModelConfigID: providerDisabledConfig.ID,
})
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
require.Contains(t, sdkErr.Message, "does not match any enabled model config")
}
func TestChatAdvisorConfig_ReasoningEffortRequiresModelConfig(t *testing.T) {
+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)
})
}