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)
|
||||
})
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user