mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
fix: hide and reject chat models from disabled AI providers (#27070)
This commit is contained in:
+89
-11
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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"`
|
||||
|
||||
@@ -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<AIProvider> =>
|
||||
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<AIProvider> =>
|
||||
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),
|
||||
]);
|
||||
},
|
||||
});
|
||||
|
||||
@@ -1830,6 +1830,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,
|
||||
@@ -1876,6 +1877,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),
|
||||
|
||||
Generated
+1
@@ -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;
|
||||
|
||||
@@ -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);
|
||||
});
|
||||
|
||||
@@ -195,6 +195,7 @@ export const canManageProviderModels = (
|
||||
): boolean => {
|
||||
return Boolean(
|
||||
providerState?.providerConfig &&
|
||||
providerState.providerConfig.enabled !== false &&
|
||||
(providerState.hasEffectiveAPIKey ||
|
||||
providerState.providerConfig.allow_user_api_key),
|
||||
);
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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<CoderAgentsPageViewProps> = ({
|
||||
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<CoderAgentsPageViewProps> = ({
|
||||
isAdvisorConfigLoading={isAdvisorConfigLoading}
|
||||
isAdvisorConfigFetching={isAdvisorConfigFetching}
|
||||
isAdvisorConfigLoadError={isAdvisorConfigLoadError}
|
||||
modelConfigs={modelConfigsData ?? []}
|
||||
enabledModelConfigs={enabledModelConfigs}
|
||||
providerInfoByID={providerInfoByID}
|
||||
modelConfigsError={modelConfigsError}
|
||||
isLoadingModelConfigs={isLoadingModelConfigs}
|
||||
|
||||
@@ -6,11 +6,13 @@ import ModelsPageView from "./ModelsPageView";
|
||||
import {
|
||||
MockAnthropicProviderState,
|
||||
MockBedrockProviderState,
|
||||
MockDisabledProviderState,
|
||||
MockOpenAIProviderState,
|
||||
mockBedrockClaude,
|
||||
mockClaude,
|
||||
mockDisabledModel,
|
||||
mockGPT5,
|
||||
mockProviderDisabledModel,
|
||||
} from "./testFixtures";
|
||||
|
||||
const meta: Meta<typeof ModelsPageView> = {
|
||||
@@ -113,6 +115,22 @@ export const NoMatchingModels: Story = {
|
||||
},
|
||||
};
|
||||
|
||||
export const DisabledProviderModelsStillListed: Story = {
|
||||
args: {
|
||||
models: [mockGPT5, mockProviderDisabledModel],
|
||||
providerStates: [MockOpenAIProviderState, MockDisabledProviderState],
|
||||
providerTypeByID: new Map<string, string>([
|
||||
["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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -257,12 +257,15 @@ export const ModelForm: FC<ModelFormProps> = ({
|
||||
selectedProviderKey={selectedProviderKey}
|
||||
onProviderChange={onProviderChange}
|
||||
disabled={isDuplicating || providerStates.length === 0}
|
||||
isEditing={isEditing}
|
||||
/>
|
||||
{selectedProviderState && (
|
||||
<p className="text-sm text-content-secondary m-0">
|
||||
{!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."}
|
||||
</p>
|
||||
)}
|
||||
</div>
|
||||
|
||||
@@ -149,6 +149,7 @@ export const ModelFormFields: FC<{
|
||||
<ModelFormProviderSelect
|
||||
providerStates={providerStates}
|
||||
selectedProviderKey={selectedProviderKey}
|
||||
isEditing={mode === "edit"}
|
||||
onProviderChange={onProviderChange}
|
||||
disabled={isDuplicating || providerStates.length === 0}
|
||||
/>
|
||||
|
||||
@@ -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 (
|
||||
<div className="grid gap-1.5">
|
||||
<Label
|
||||
@@ -40,7 +55,7 @@ export const ModelFormProviderSelect: FC<{
|
||||
<SelectValue placeholder="Select provider" />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
{providerStates.map((ps) => (
|
||||
{selectableProviderStates.map((ps) => (
|
||||
<SelectItem key={ps.key} value={ps.key}>
|
||||
<span className="flex items-center gap-2">
|
||||
<ProviderIcon provider={ps.provider} />
|
||||
|
||||
@@ -106,6 +106,31 @@ export const MockBedrockProviderState: ProviderState = {
|
||||
modelConfigs: [mockBedrockClaude],
|
||||
};
|
||||
|
||||
const MockDisabledProviderConfig: ChatProviderConfig = {
|
||||
...MockOpenAIProviderConfig,
|
||||
id: "prov-openai-disabled",
|
||||
display_name: "OpenAI Secondary",
|
||||
enabled: false,
|
||||
};
|
||||
|
||||
export const mockProviderDisabledModel: ChatModelConfig = {
|
||||
...mockGPT5,
|
||||
id: "model-provider-disabled",
|
||||
ai_provider_id: "prov-openai-disabled",
|
||||
model: "gpt-4o-secondary",
|
||||
display_name: "GPT-4o Secondary",
|
||||
is_default: false,
|
||||
};
|
||||
|
||||
export const MockDisabledProviderState: ProviderState = {
|
||||
...MockOpenAIProviderState,
|
||||
key: "prov-openai-disabled",
|
||||
provider: "openai",
|
||||
label: "OpenAI Secondary",
|
||||
providerConfig: MockDisabledProviderConfig,
|
||||
modelConfigs: [mockProviderDisabledModel],
|
||||
};
|
||||
|
||||
export const MockCopilotProviderState: ProviderState = {
|
||||
...MockOpenAIProviderState,
|
||||
key: "prov-copilot",
|
||||
|
||||
@@ -18,6 +18,7 @@ const createProvider = (
|
||||
provider: overrides.provider,
|
||||
display_name: overrides.display_name ?? overrides.provider,
|
||||
icon: overrides.icon ?? "",
|
||||
enabled: overrides.enabled ?? true,
|
||||
has_user_api_key: overrides.has_user_api_key ?? false,
|
||||
has_central_api_key_fallback: overrides.has_central_api_key_fallback ?? false,
|
||||
byok_enabled: overrides.byok_enabled ?? true,
|
||||
|
||||
@@ -26,7 +26,7 @@ interface AdvisorSettingsProps {
|
||||
isAdvisorConfigLoading: boolean;
|
||||
isAdvisorConfigFetching: boolean;
|
||||
isAdvisorConfigLoadError: boolean;
|
||||
modelConfigs: readonly ChatModelConfig[];
|
||||
enabledModelConfigs: readonly ChatModelConfig[];
|
||||
providerInfoByID: ReadonlyMap<string, ProviderInfo>;
|
||||
modelConfigsError: unknown;
|
||||
isLoadingModelConfigs: boolean;
|
||||
@@ -120,7 +120,7 @@ export const AdvisorSettings: FC<AdvisorSettingsProps> = ({
|
||||
isAdvisorConfigLoading,
|
||||
isAdvisorConfigFetching,
|
||||
isAdvisorConfigLoadError,
|
||||
modelConfigs,
|
||||
enabledModelConfigs,
|
||||
providerInfoByID,
|
||||
modelConfigsError,
|
||||
isLoadingModelConfigs,
|
||||
@@ -134,27 +134,25 @@ export const AdvisorSettings: FC<AdvisorSettingsProps> = ({
|
||||
const maxOutputTokensId = useId();
|
||||
const { isSavedVisible, showSavedState } = useTemporarySavedState();
|
||||
const hasLoadedAdvisorConfig = advisorConfigData !== undefined;
|
||||
const enabledModelOptions = modelConfigs
|
||||
.filter((config) => config.enabled)
|
||||
.map((config) => {
|
||||
const providerInfo = providerInfoByID.get(config.ai_provider_id);
|
||||
const reasoningEffort = config.model_config?.reasoning_effort;
|
||||
const reasoningEfforts = config.reasoning_efforts ?? [];
|
||||
return {
|
||||
id: config.id,
|
||||
provider: providerInfo?.provider ?? "",
|
||||
providerId: config.ai_provider_id,
|
||||
providerLabel: providerInfo?.displayName,
|
||||
providerIcon: providerInfo?.icon,
|
||||
model: config.model,
|
||||
displayName: config.display_name.trim() || config.model,
|
||||
contextLimit: config.context_limit,
|
||||
...(reasoningEffort?.default
|
||||
? { reasoningEffortDefault: reasoningEffort.default }
|
||||
: {}),
|
||||
...(reasoningEfforts.length > 0 ? { reasoningEfforts } : {}),
|
||||
};
|
||||
});
|
||||
const enabledModelOptions = enabledModelConfigs.map((config) => {
|
||||
const providerInfo = providerInfoByID.get(config.ai_provider_id);
|
||||
const reasoningEffort = config.model_config?.reasoning_effort;
|
||||
const reasoningEfforts = config.reasoning_efforts ?? [];
|
||||
return {
|
||||
id: config.id,
|
||||
provider: providerInfo?.provider ?? "",
|
||||
providerId: config.ai_provider_id,
|
||||
providerLabel: providerInfo?.displayName,
|
||||
providerIcon: providerInfo?.icon,
|
||||
model: config.model,
|
||||
displayName: config.display_name.trim() || config.model,
|
||||
contextLimit: config.context_limit,
|
||||
...(reasoningEffort?.default
|
||||
? { reasoningEffortDefault: reasoningEffort.default }
|
||||
: {}),
|
||||
...(reasoningEfforts.length > 0 ? { reasoningEfforts } : {}),
|
||||
};
|
||||
});
|
||||
|
||||
const form = useFormik<AdvisorSettingsFormValues>({
|
||||
enableReinitialize: true,
|
||||
@@ -163,17 +161,19 @@ export const AdvisorSettings: FC<AdvisorSettingsProps> = ({
|
||||
validate: validateAdvisorConfig,
|
||||
onSubmit: (values, { resetForm }) => {
|
||||
// If the last committed model override references a model config
|
||||
// that no longer exists, the backend rejects the stale ID with a
|
||||
// 400. Clear the override so a save stays reliable in that edge
|
||||
// case. Only scrub when model configs have loaded successfully and
|
||||
// no refetch is in flight.
|
||||
// that is no longer available, the backend rejects the stale ID
|
||||
// with a 400. Clear the override so a save stays reliable in
|
||||
// that edge case. Only scrub when model configs have loaded
|
||||
// successfully and no refetch is in flight.
|
||||
let source = values;
|
||||
if (
|
||||
!isUnsetModelConfigId(source.model_config_id) &&
|
||||
!isLoadingModelConfigs &&
|
||||
!isFetchingModelConfigs &&
|
||||
!modelConfigsError &&
|
||||
!modelConfigs.some((config) => config.id === source.model_config_id)
|
||||
!enabledModelConfigs.some(
|
||||
(config) => config.id === source.model_config_id,
|
||||
)
|
||||
) {
|
||||
source = { ...source, model_config_id: "", reasoning_effort: "" };
|
||||
}
|
||||
|
||||
@@ -10,6 +10,7 @@ import {
|
||||
} from "#/testHelpers/chatModels";
|
||||
import {
|
||||
countConfiguredProviderConfigs,
|
||||
filterConfigsWithEnabledProvider,
|
||||
formatProviderLabel,
|
||||
getModelOptionsFromConfigs,
|
||||
getModelSelectorPlaceholder,
|
||||
@@ -645,6 +646,165 @@ describe("getModelOptionsFromConfigs", () => {
|
||||
}),
|
||||
]);
|
||||
});
|
||||
|
||||
it("drops configs whose provider row is disabled", () => {
|
||||
const configs = [
|
||||
createConfig({
|
||||
id: "config-enabled",
|
||||
ai_provider_id: "prov-enabled",
|
||||
model: "gpt-4o",
|
||||
}),
|
||||
createConfig({
|
||||
id: "config-disabled",
|
||||
ai_provider_id: "prov-disabled",
|
||||
model: "gpt-4o-mini",
|
||||
}),
|
||||
];
|
||||
const catalog = createCatalog([
|
||||
{ provider: "openai", available: true, models: [] },
|
||||
]);
|
||||
const providers = new Map([
|
||||
[
|
||||
"prov-enabled",
|
||||
{ provider: "openai", displayName: "OpenAI", icon: "", enabled: true },
|
||||
],
|
||||
[
|
||||
"prov-disabled",
|
||||
{
|
||||
provider: "openai",
|
||||
displayName: "OpenAI Disabled",
|
||||
icon: "",
|
||||
enabled: false,
|
||||
},
|
||||
],
|
||||
]);
|
||||
|
||||
expect(
|
||||
getModelOptionsFromConfigs(configs, catalog, providers).map(
|
||||
(option) => option.id,
|
||||
),
|
||||
).toEqual(["config-enabled"]);
|
||||
});
|
||||
|
||||
it("keeps configs when the provider enabled flag is undefined", () => {
|
||||
const configs = [
|
||||
createConfig({
|
||||
id: "config-openai",
|
||||
ai_provider_id: "prov-openai",
|
||||
model: "gpt-4o",
|
||||
}),
|
||||
];
|
||||
const catalog = createCatalog([
|
||||
{ provider: "openai", available: true, models: [] },
|
||||
]);
|
||||
|
||||
expect(
|
||||
getModelOptionsFromConfigs(configs, catalog, providerInfoByID).map(
|
||||
(option) => option.id,
|
||||
),
|
||||
).toEqual(["config-openai"]);
|
||||
});
|
||||
|
||||
it("excludes only the disabled instance for same-type providers", () => {
|
||||
// The catalog marks the type as available because of the enabled
|
||||
// instance, so only the per-row flag can exclude the disabled one.
|
||||
const configs = [
|
||||
createConfig({
|
||||
id: "config-primary",
|
||||
ai_provider_id: "prov-anthropic-primary",
|
||||
model: "claude-sonnet-4-20250514",
|
||||
}),
|
||||
createConfig({
|
||||
id: "config-secondary",
|
||||
ai_provider_id: "prov-anthropic-secondary",
|
||||
model: "claude-opus-4-20250514",
|
||||
}),
|
||||
];
|
||||
const catalog = createCatalog([
|
||||
{ provider: "anthropic", available: true, models: [] },
|
||||
]);
|
||||
const sameTypeProviders = new Map([
|
||||
[
|
||||
"prov-anthropic-primary",
|
||||
{
|
||||
provider: "anthropic",
|
||||
displayName: "Anthropic",
|
||||
icon: "",
|
||||
enabled: true,
|
||||
},
|
||||
],
|
||||
[
|
||||
"prov-anthropic-secondary",
|
||||
{
|
||||
provider: "anthropic",
|
||||
displayName: "Anthropic Secondary",
|
||||
icon: "",
|
||||
enabled: false,
|
||||
},
|
||||
],
|
||||
]);
|
||||
|
||||
expect(
|
||||
getModelOptionsFromConfigs(configs, catalog, sameTypeProviders).map(
|
||||
(option) => option.id,
|
||||
),
|
||||
).toEqual(["config-primary"]);
|
||||
});
|
||||
});
|
||||
|
||||
describe("filterConfigsWithEnabledProvider", () => {
|
||||
const configs = [
|
||||
createConfig({
|
||||
id: "config-enabled",
|
||||
ai_provider_id: "prov-enabled",
|
||||
model: "gpt-4o",
|
||||
}),
|
||||
createConfig({
|
||||
id: "config-disabled",
|
||||
ai_provider_id: "prov-disabled",
|
||||
model: "gpt-4o-mini",
|
||||
}),
|
||||
createConfig({
|
||||
id: "config-unknown",
|
||||
ai_provider_id: "prov-unknown",
|
||||
model: "claude-sonnet-4-20250514",
|
||||
}),
|
||||
];
|
||||
const providers = new Map([
|
||||
[
|
||||
"prov-enabled",
|
||||
{ provider: "openai", displayName: "OpenAI", icon: "", enabled: true },
|
||||
],
|
||||
[
|
||||
"prov-disabled",
|
||||
{
|
||||
provider: "openai",
|
||||
displayName: "OpenAI Disabled",
|
||||
icon: "",
|
||||
enabled: false,
|
||||
},
|
||||
],
|
||||
]);
|
||||
|
||||
it("drops configs of disabled or unknown providers", () => {
|
||||
expect(
|
||||
filterConfigsWithEnabledProvider(configs, providers).map(
|
||||
(config) => config.id,
|
||||
),
|
||||
).toEqual(["config-enabled"]);
|
||||
});
|
||||
|
||||
it("keeps configs whose provider rows lack enabled flags", () => {
|
||||
const flaglessProviders = new Map(
|
||||
["prov-enabled", "prov-disabled", "prov-unknown"].map((id) => [
|
||||
id,
|
||||
{ provider: "openai", displayName: "OpenAI", icon: "" },
|
||||
]),
|
||||
);
|
||||
expect(
|
||||
filterConfigsWithEnabledProvider(configs, flaglessProviders),
|
||||
).toEqual(configs);
|
||||
});
|
||||
});
|
||||
|
||||
describe("providerInfoByIDFromConfigs", () => {
|
||||
@@ -663,6 +823,7 @@ describe("providerInfoByIDFromConfigs", () => {
|
||||
provider: "openai",
|
||||
displayName: "Primary OpenAI",
|
||||
icon: "/icon/openai.svg",
|
||||
enabled: true,
|
||||
});
|
||||
expect(map.size).toBe(1);
|
||||
});
|
||||
@@ -676,6 +837,7 @@ describe("providerInfoByIDFromUserConfigs", () => {
|
||||
provider: "openai",
|
||||
display_name: "Primary OpenAI",
|
||||
icon: "/icon/openai.svg",
|
||||
enabled: true,
|
||||
has_user_api_key: false,
|
||||
has_central_api_key_fallback: true,
|
||||
byok_enabled: true,
|
||||
@@ -686,6 +848,7 @@ describe("providerInfoByIDFromUserConfigs", () => {
|
||||
provider: "openai",
|
||||
displayName: "Primary OpenAI",
|
||||
icon: "/icon/openai.svg",
|
||||
enabled: true,
|
||||
});
|
||||
expect(map.size).toBe(1);
|
||||
});
|
||||
@@ -721,6 +884,7 @@ describe("providerTypeByIDFromUserConfigs", () => {
|
||||
provider: "openai",
|
||||
display_name: "OpenAI",
|
||||
icon: "",
|
||||
enabled: true,
|
||||
has_user_api_key: false,
|
||||
has_central_api_key_fallback: true,
|
||||
byok_enabled: true,
|
||||
@@ -798,6 +962,7 @@ describe("resolveModelSelector", () => {
|
||||
provider: "openai",
|
||||
display_name: "OpenAI",
|
||||
icon: "",
|
||||
enabled: true,
|
||||
has_user_api_key: false,
|
||||
has_central_api_key_fallback: true,
|
||||
byok_enabled: true,
|
||||
|
||||
@@ -162,6 +162,8 @@ export type ProviderInfo = {
|
||||
readonly provider: string;
|
||||
readonly displayName: string;
|
||||
readonly icon: string;
|
||||
// Absent is treated as enabled.
|
||||
readonly enabled?: boolean;
|
||||
};
|
||||
|
||||
// providerInfoByIDFromConfigs and providerInfoByIDFromUserConfigs build
|
||||
@@ -179,6 +181,7 @@ export const providerInfoByIDFromConfigs = (
|
||||
provider: providerConfig.provider,
|
||||
displayName: providerConfig.display_name,
|
||||
icon: providerConfig.icon,
|
||||
enabled: providerConfig.enabled,
|
||||
},
|
||||
]),
|
||||
);
|
||||
@@ -196,6 +199,7 @@ export const providerInfoByIDFromUserConfigs = (
|
||||
provider: providerConfig.provider,
|
||||
displayName: providerConfig.display_name,
|
||||
icon: providerConfig.icon,
|
||||
enabled: providerConfig.enabled,
|
||||
},
|
||||
]),
|
||||
);
|
||||
@@ -223,6 +227,20 @@ export const providerTypeByIDFromUserConfigs = (
|
||||
),
|
||||
);
|
||||
|
||||
/**
|
||||
* Drops model configs whose provider row is disabled or missing. Both
|
||||
* provider-info sources include every enabled provider, so a missing row
|
||||
* means the provider is disabled or deleted.
|
||||
*/
|
||||
export const filterConfigsWithEnabledProvider = (
|
||||
configs: readonly TypesGen.ChatModelConfig[],
|
||||
providerInfoByID: ReadonlyMap<string, ProviderInfo>,
|
||||
): readonly TypesGen.ChatModelConfig[] =>
|
||||
configs.filter((config) => {
|
||||
const info = providerInfoByID.get(config.ai_provider_id);
|
||||
return info !== undefined && info.enabled !== false;
|
||||
});
|
||||
|
||||
export const getModelOptionsFromConfigs = (
|
||||
configs: readonly TypesGen.ChatModelConfig[] | null | undefined,
|
||||
catalog: TypesGen.ChatModelsResponse | null | undefined,
|
||||
@@ -235,7 +253,12 @@ export const getModelOptionsFromConfigs = (
|
||||
const availableProviders = getAvailableProviders(catalog);
|
||||
const options: ModelSelectorOption[] = [];
|
||||
|
||||
for (const config of configs) {
|
||||
// The catalog check below is keyed by provider type, so it cannot
|
||||
// exclude a disabled provider when another of the same type is enabled.
|
||||
for (const config of filterConfigsWithEnabledProvider(
|
||||
configs,
|
||||
providerInfoByID,
|
||||
)) {
|
||||
if (!config.enabled) {
|
||||
continue;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user