diff --git a/coderd/exp_chats.go b/coderd/exp_chats.go index 4632384683495..56d0337da0d15 100644 --- a/coderd/exp_chats.go +++ b/coderd/exp_chats.go @@ -1193,6 +1193,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 @@ -1431,6 +1446,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, @@ -3336,6 +3358,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) { @@ -3385,6 +3411,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 @@ -3501,6 +3533,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) { @@ -3540,6 +3576,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): @@ -4818,6 +4858,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 } @@ -4912,6 +4955,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 } @@ -5703,16 +5761,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 } @@ -7625,7 +7678,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 { @@ -7633,12 +7699,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 { diff --git a/coderd/exp_chats_test.go b/coderd/exp_chats_test.go index 209e51d1e7131..d8beed573103c 100644 --- a/coderd/exp_chats_test.go +++ b/coderd/exp_chats_test.go @@ -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) }) } @@ -11889,6 +12199,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, @@ -12704,6 +13033,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() @@ -14201,7 +14544,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) { diff --git a/coderd/x/chatd/chatd.go b/coderd/x/chatd/chatd.go index fa2012b2adece..1bb1209c1ba5c 100644 --- a/coderd/x/chatd/chatd.go +++ b/coderd/x/chatd/chatd.go @@ -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. @@ -1341,6 +1342,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, opts.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, @@ -1564,22 +1571,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( @@ -1589,7 +1609,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( @@ -1607,6 +1627,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 } @@ -1674,24 +1709,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 diff --git a/coderd/x/chatd/chatd_internal_test.go b/coderd/x/chatd/chatd_internal_test.go index 916df0f80fb88..a2d72ac069a3c 100644 --- a/coderd/x/chatd/chatd_internal_test.go +++ b/coderd/x/chatd/chatd_internal_test.go @@ -3483,3 +3483,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) + }) +} diff --git a/codersdk/chats.go b/codersdk/chats.go index f28ef011eea9b..dc361d1e9ecd7 100644 --- a/codersdk/chats.go +++ b/codersdk/chats.go @@ -1272,6 +1272,7 @@ type UserChatProviderConfig struct { Provider string `json:"provider"` DisplayName string `json:"display_name"` Icon string `json:"icon"` + Enabled bool `json:"enabled"` HasUserAPIKey bool `json:"has_user_api_key"` HasCentralAPIKeyFallback bool `json:"has_central_api_key_fallback"` BYOKEnabled bool `json:"byok_enabled"` diff --git a/site/src/api/queries/aiProviders.ts b/site/src/api/queries/aiProviders.ts index 7a3f01cf53528..3b5b279314730 100644 --- a/site/src/api/queries/aiProviders.ts +++ b/site/src/api/queries/aiProviders.ts @@ -1,5 +1,6 @@ import type { QueryClient } from "react-query"; import { API } from "#/api/api"; +import { invalidateChatProviderDependentQueries } from "#/api/queries/chats"; import type { AIProvider, CreateAIProviderRequest, @@ -25,7 +26,10 @@ export const createAIProviderMutation = (queryClient: QueryClient) => ({ mutationFn: (request: CreateAIProviderRequest): Promise => API.createAIProvider(request), onSuccess: async () => { - await queryClient.invalidateQueries({ queryKey: aiProvidersListKey }); + await Promise.all([ + queryClient.invalidateQueries({ queryKey: aiProvidersListKey }), + invalidateChatProviderDependentQueries(queryClient), + ]); }, }); @@ -36,10 +40,13 @@ export const updateAIProviderMutation = ( mutationFn: (request: UpdateAIProviderRequest): Promise => API.updateAIProvider(idOrName, request), onSuccess: async () => { - await queryClient.invalidateQueries({ queryKey: aiProvidersListKey }); - await queryClient.invalidateQueries({ - queryKey: aiProviderKeyFor(idOrName), - }); + await Promise.all([ + queryClient.invalidateQueries({ queryKey: aiProvidersListKey }), + queryClient.invalidateQueries({ + queryKey: aiProviderKeyFor(idOrName), + }), + invalidateChatProviderDependentQueries(queryClient), + ]); }, }); @@ -49,7 +56,10 @@ export const deleteAIProviderMutation = ( ) => ({ mutationFn: () => API.deleteAIProvider(idOrName), onSuccess: async () => { - await queryClient.invalidateQueries({ queryKey: aiProvidersListKey }); queryClient.removeQueries({ queryKey: aiProviderKeyFor(idOrName) }); + await Promise.all([ + queryClient.invalidateQueries({ queryKey: aiProvidersListKey }), + invalidateChatProviderDependentQueries(queryClient), + ]); }, }); diff --git a/site/src/api/queries/chats.ts b/site/src/api/queries/chats.ts index 281f04234b18a..f126cbbb6cf04 100644 --- a/site/src/api/queries/chats.ts +++ b/site/src/api/queries/chats.ts @@ -1826,6 +1826,7 @@ export const userChatProviderConfigs = () => ({ provider: config.provider.type, display_name: config.provider.display_name || config.provider.type, icon: config.provider.icon, + enabled: config.provider.enabled, has_user_api_key: config.has_user_api_key, byok_enabled: config.byok_enabled, has_central_api_key_fallback: config.has_provider_api_key, @@ -1872,6 +1873,16 @@ const invalidateChatConfigurationQueries = async (queryClient: QueryClient) => { ]); }; +// Called after AI provider mutations so open model pickers refresh. +export const invalidateChatProviderDependentQueries = async ( + queryClient: QueryClient, +) => { + await Promise.all([ + invalidateChatConfigurationQueries(queryClient), + queryClient.invalidateQueries({ queryKey: userChatProviderConfigsKey }), + ]); +}; + export const createChatModelConfig = (queryClient: QueryClient) => ({ mutationFn: (req: TypesGen.CreateChatModelConfigRequest) => API.experimental.createChatModelConfig(req), diff --git a/site/src/api/typesGenerated.ts b/site/src/api/typesGenerated.ts index c21a3338a6e2c..3aebc9edcaac6 100644 --- a/site/src/api/typesGenerated.ts +++ b/site/src/api/typesGenerated.ts @@ -9947,6 +9947,7 @@ export interface UserChatProviderConfig { readonly provider: string; readonly display_name: string; readonly icon: string; + readonly enabled: boolean; readonly has_user_api_key: boolean; readonly has_central_api_key_fallback: boolean; readonly byok_enabled: boolean; diff --git a/site/src/modules/aiModels/providerStates.test.ts b/site/src/modules/aiModels/providerStates.test.ts index 60d165cbc9eab..d0da0d1008d4d 100644 --- a/site/src/modules/aiModels/providerStates.test.ts +++ b/site/src/modules/aiModels/providerStates.test.ts @@ -225,6 +225,15 @@ describe("canManageProviderModels", () => { ).toBe(false); }); + it("returns false when the provider is disabled", () => { + expect( + canManageProviderModels({ + ...baseState, + providerConfig: { ...MockChatProviderConfig, enabled: false }, + }), + ).toBe(false); + }); + it("returns false for undefined provider state", () => { expect(canManageProviderModels(undefined)).toBe(false); }); diff --git a/site/src/modules/aiModels/providerStates.ts b/site/src/modules/aiModels/providerStates.ts index a94fc6b9f39d5..ce06c9063fc12 100644 --- a/site/src/modules/aiModels/providerStates.ts +++ b/site/src/modules/aiModels/providerStates.ts @@ -195,6 +195,7 @@ export const canManageProviderModels = ( ): boolean => { return Boolean( providerState?.providerConfig && + providerState.providerConfig.enabled !== false && (providerState.hasEffectiveAPIKey || providerState.providerConfig.allow_user_api_key), ); diff --git a/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPage.tsx b/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPage.tsx index d794dd74f3b99..f825aab183be3 100644 --- a/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPage.tsx +++ b/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPage.tsx @@ -154,9 +154,15 @@ const CoderAgentsPage: FC = () => { exploreModelOverrideData={exploreModelOverrideQuery.data} modelConfigsData={modelConfigsQuery.data} providerInfoByID={providerInfoByID} - modelConfigsError={modelConfigsQuery.error} - isLoadingModelConfigs={modelConfigsQuery.isLoading} - isFetchingModelConfigs={modelConfigsQuery.isFetching} + modelConfigsError={ + modelConfigsQuery.error ?? providerConfigsQuery.error + } + isLoadingModelConfigs={ + modelConfigsQuery.isLoading || providerConfigsQuery.isLoading + } + isFetchingModelConfigs={ + modelConfigsQuery.isFetching || providerConfigsQuery.isFetching + } onSaveGeneralModelOverride={saveGeneralModelOverrideMutation.mutate} isSavingGeneralModelOverride={ saveGeneralModelOverrideMutation.isPending diff --git a/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.stories.tsx b/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.stories.tsx index a67df1eaf97b0..c70db40d48a2b 100644 --- a/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.stories.tsx +++ b/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.stories.tsx @@ -113,6 +113,14 @@ const compactionDisabledModelConfig = buildModelConfig({ context_limit: 128_000, }); +const providerDisabledModelConfig = buildModelConfig({ + id: "model-provider-disabled", + ai_provider_id: "provider-openai-disabled", + model: "gpt-4o-secondary", + display_name: "GPT 4o Secondary", + context_limit: 128_000, +}); + const allModelConfigs: TypesGen.ChatModelConfig[] = [ generalModelConfig, claudeSonnetModelConfig, @@ -123,13 +131,31 @@ const allModelConfigs: TypesGen.ChatModelConfig[] = [ titleDisabledModelConfig, exploreDisabledModelConfig, compactionDisabledModelConfig, + providerDisabledModelConfig, ]; const providerInfoByID = new Map([ - ["provider-1", { provider: "openai", displayName: "OpenAI", icon: "" }], + [ + "provider-1", + { provider: "openai", displayName: "OpenAI", icon: "", enabled: true }, + ], [ "provider-anthropic", - { provider: "anthropic", displayName: "Anthropic", icon: "" }, + { + provider: "anthropic", + displayName: "Anthropic", + icon: "", + enabled: true, + }, + ], + [ + "provider-openai-disabled", + { + provider: "openai", + displayName: "OpenAI Secondary", + icon: "", + enabled: false, + }, ], ]); @@ -668,6 +694,47 @@ export const AdvisorReasoningEffort: Story = { }, }; +export const DisabledProviderModelsHidden: Story = { + args: buildArgs({ + showAdvisorSettings: true, + advisorConfigData: { + enabled: true, + max_uses_per_run: 3, + max_output_tokens: 16384, + model_config_id: "00000000-0000-0000-0000-000000000000", + }, + }), + play: async ({ canvasElement }) => { + const body = within(canvasElement.ownerDocument.body); + + const generalSection = await getSection(canvasElement, "General model"); + const generalTrigger = within(generalSection).getByRole("combobox", { + name: "Use chat default", + }); + await userEvent.click(generalTrigger); + expect( + await body.findByRole("option", { name: /GPT 4\.1 Mini/ }), + ).toBeInTheDocument(); + expect( + body.queryByRole("option", { name: /GPT 4o Secondary/ }), + ).not.toBeInTheDocument(); + await userEvent.keyboard("{Escape}"); + + const advisorSection = await getSection(canvasElement, "Advisor"); + const advisorTrigger = within(advisorSection).getByRole("combobox", { + name: "Use chat model", + }); + await userEvent.click(advisorTrigger); + expect( + await body.findByRole("option", { name: /GPT 4\.1 Mini/ }), + ).toBeInTheDocument(); + expect( + body.queryByRole("option", { name: /GPT 4o Secondary/ }), + ).not.toBeInTheDocument(); + await userEvent.keyboard("{Escape}"); + }, +}; + export const AdvisorClearButton: Story = { args: buildArgs({ showAdvisorSettings: true, diff --git a/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.tsx b/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.tsx index 0022c2dc3b86d..c9e781a8d7c82 100644 --- a/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.tsx +++ b/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.tsx @@ -8,7 +8,10 @@ import { } from "#/components/SettingsHeader/SettingsHeader"; import { AdvisorSettings } from "#/pages/AgentsPage/components/AdvisorSettings"; import { VirtualDesktopSettings } from "#/pages/AgentsPage/components/VirtualDesktopSettings"; -import type { ProviderInfo } from "#/pages/AgentsPage/utils/modelOptions"; +import { + filterConfigsWithEnabledProvider, + type ProviderInfo, +} from "#/pages/AgentsPage/utils/modelOptions"; import { AdminPersonalModelOverridesSettings, type SavePersonalModelOverridesAdminSetting, @@ -122,8 +125,9 @@ export const CoderAgentsPageView: FC = ({ isSavingComputerUseProvider, computerUseProviderSaveError, }) => { - const enabledModelConfigs = (modelConfigsData ?? []).filter( - (modelConfig) => modelConfig.enabled, + const enabledModelConfigs = filterConfigsWithEnabledProvider( + (modelConfigsData ?? []).filter((modelConfig) => modelConfig.enabled), + providerInfoByID, ); const showGeneralModelSection = onSaveGeneralModelOverride !== undefined || @@ -222,7 +226,7 @@ export const CoderAgentsPageView: FC = ({ isAdvisorConfigLoading={isAdvisorConfigLoading} isAdvisorConfigFetching={isAdvisorConfigFetching} isAdvisorConfigLoadError={isAdvisorConfigLoadError} - modelConfigs={modelConfigsData ?? []} + enabledModelConfigs={enabledModelConfigs} providerInfoByID={providerInfoByID} modelConfigsError={modelConfigsError} isLoadingModelConfigs={isLoadingModelConfigs} diff --git a/site/src/pages/AISettingsPage/ModelsPage/ModelsPageView.stories.tsx b/site/src/pages/AISettingsPage/ModelsPage/ModelsPageView.stories.tsx index c4de1c4ffda59..86e388dd2d6e6 100644 --- a/site/src/pages/AISettingsPage/ModelsPage/ModelsPageView.stories.tsx +++ b/site/src/pages/AISettingsPage/ModelsPage/ModelsPageView.stories.tsx @@ -6,11 +6,13 @@ import ModelsPageView from "./ModelsPageView"; import { MockAnthropicProviderState, MockBedrockProviderState, + MockDisabledProviderState, MockOpenAIProviderState, mockBedrockClaude, mockClaude, mockDisabledModel, mockGPT5, + mockProviderDisabledModel, } from "./testFixtures"; const meta: Meta = { @@ -113,6 +115,22 @@ export const NoMatchingModels: Story = { }, }; +export const DisabledProviderModelsStillListed: Story = { + args: { + models: [mockGPT5, mockProviderDisabledModel], + providerStates: [MockOpenAIProviderState, MockDisabledProviderState], + providerTypeByID: new Map([ + ["prov-openai", "openai"], + ["prov-openai-disabled", "openai"], + ]), + }, + play: async ({ canvasElement }) => { + const canvas = within(canvasElement); + await expect(canvas.getByText("GPT-4o Secondary")).toBeInTheDocument(); + await expect(canvas.getByText("OpenAI Secondary")).toBeInTheDocument(); + }, +}; + export const Loading: Story = { args: { isLoading: true, diff --git a/site/src/pages/AISettingsPage/ModelsPage/components/ModelForm.stories.tsx b/site/src/pages/AISettingsPage/ModelsPage/components/ModelForm.stories.tsx index 91b9de7462080..5ae23484ad44b 100644 --- a/site/src/pages/AISettingsPage/ModelsPage/components/ModelForm.stories.tsx +++ b/site/src/pages/AISettingsPage/ModelsPage/components/ModelForm.stories.tsx @@ -4,8 +4,10 @@ import { reactRouterParameters } from "storybook-addon-remix-react-router"; import { withToaster } from "#/testHelpers/storybook"; import { MockAnthropicProviderState, + MockDisabledProviderState, MockOpenAIProviderState, mockGPT5, + mockProviderDisabledModel, } from "../testFixtures"; import { ModelForm } from "./ModelForm"; @@ -126,6 +128,66 @@ export const ReplaceDefaultWarning: Story = { }, }; +export const AddHidesDisabledProviders: Story = { + args: { + providerStates: [ + MockOpenAIProviderState, + MockAnthropicProviderState, + MockDisabledProviderState, + ], + }, + play: async ({ canvasElement }) => { + const canvas = within(canvasElement); + await userEvent.click(canvas.getByRole("combobox", { name: /provider/i })); + // Option names include the provider icon alt text, so match loosely. + const optionNames = screen + .getAllByRole("option") + .map((option) => option.textContent?.trim()); + await expect(optionNames).toEqual(["OpenAI", "Anthropic"]); + await expect( + screen.queryByRole("option", { name: /Secondary/ }), + ).not.toBeInTheDocument(); + }, +}; + +export const AddBlocksDisabledSelectedProvider: Story = { + args: { + providerStates: [MockOpenAIProviderState, MockDisabledProviderState], + selectedProviderState: MockDisabledProviderState, + }, + play: async ({ canvasElement }) => { + const canvas = within(canvasElement); + // A ?provider= query param can preselect a disabled provider on + // the add page. + await expect( + canvas.getByText(/OpenAI Secondary is disabled/), + ).toBeInTheDocument(); + await expect( + canvas.queryByRole("button", { name: /add model/i }), + ).not.toBeInTheDocument(); + await userEvent.click(canvas.getByRole("combobox", { name: /provider/i })); + await expect( + screen.queryByRole("option", { name: /Secondary/ }), + ).not.toBeInTheDocument(); + }, +}; + +export const EditKeepsDisabledProviderVisible: Story = { + args: { + providerStates: [MockOpenAIProviderState, MockDisabledProviderState], + selectedProviderState: MockDisabledProviderState, + editingModel: mockProviderDisabledModel, + onDeleteModel: fn(async () => undefined), + onDuplicate: fn(), + }, + play: async ({ canvasElement }) => { + const canvas = within(canvasElement); + await expect( + canvas.getByRole("combobox", { name: /provider/i }), + ).toHaveTextContent("OpenAI Secondary"); + }, +}; + export const Edit: Story = { args: { editingModel: mockGPT5, diff --git a/site/src/pages/AISettingsPage/ModelsPage/components/ModelForm.tsx b/site/src/pages/AISettingsPage/ModelsPage/components/ModelForm.tsx index e6062e15d1630..72317fb82e54c 100644 --- a/site/src/pages/AISettingsPage/ModelsPage/components/ModelForm.tsx +++ b/site/src/pages/AISettingsPage/ModelsPage/components/ModelForm.tsx @@ -257,12 +257,15 @@ export const ModelForm: FC = ({ selectedProviderKey={selectedProviderKey} onProviderChange={onProviderChange} disabled={isDuplicating || providerStates.length === 0} + isEditing={isEditing} /> {selectedProviderState && (

{!selectedProviderState.providerConfig ? "Create a managed provider before adding models." - : "Set an API key for this provider before adding models."} + : selectedProviderState.providerConfig.enabled === false + ? `${selectedProviderState.label} is disabled. Enable it before adding models.` + : "Set an API key for this provider before adding models."}

)} diff --git a/site/src/pages/AISettingsPage/ModelsPage/components/ModelFormFields.tsx b/site/src/pages/AISettingsPage/ModelsPage/components/ModelFormFields.tsx index cb464de7e8f2a..aff9db74be82a 100644 --- a/site/src/pages/AISettingsPage/ModelsPage/components/ModelFormFields.tsx +++ b/site/src/pages/AISettingsPage/ModelsPage/components/ModelFormFields.tsx @@ -149,6 +149,7 @@ export const ModelFormFields: FC<{ diff --git a/site/src/pages/AISettingsPage/ModelsPage/components/ModelFormProviderSelect.tsx b/site/src/pages/AISettingsPage/ModelsPage/components/ModelFormProviderSelect.tsx index 9ae0b7d5c1687..a5375567dbea0 100644 --- a/site/src/pages/AISettingsPage/ModelsPage/components/ModelFormProviderSelect.tsx +++ b/site/src/pages/AISettingsPage/ModelsPage/components/ModelFormProviderSelect.tsx @@ -15,7 +15,22 @@ export const ModelFormProviderSelect: FC<{ selectedProviderKey: string; onProviderChange: (providerKey: string) => void; disabled: boolean; -}> = ({ providerStates, selectedProviderKey, onProviderChange, disabled }) => { + isEditing: boolean; +}> = ({ + providerStates, + selectedProviderKey, + onProviderChange, + disabled, + isEditing, +}) => { + // Hide disabled providers; the backend rejects new model configs under + // them. When editing, keep the selected provider visible so a config + // whose provider was disabled afterwards still renders. + const selectableProviderStates = providerStates.filter( + (ps) => + ps.providerConfig?.enabled !== false || + (isEditing && ps.key === selectedProviderKey), + ); return (