diff --git a/cli/exp_scaletest_chat_test.go b/cli/exp_scaletest_chat_test.go index 9bbab931c4..0529258123 100644 --- a/cli/exp_scaletest_chat_test.go +++ b/cli/exp_scaletest_chat_test.go @@ -128,7 +128,7 @@ func chatMessageText(messages []codersdk.ChatMessage, role codersdk.ChatMessageR func scaletestModelConfigsForProvider(configs []codersdk.ChatModelConfig, providerID uuid.UUID) []codersdk.ChatModelConfig { matches := make([]codersdk.ChatModelConfig, 0, 1) for _, config := range configs { - if config.AIProviderID == nil || *config.AIProviderID != providerID { + if config.AIProviderID != providerID { continue } if config.Model != "scaletest-model" { diff --git a/cli/server.go b/cli/server.go index 505321d852..5b754b88c7 100644 --- a/cli/server.go +++ b/cli/server.go @@ -1117,9 +1117,6 @@ func (r *RootCmd) Server(newAPI func(context.Context, *coderd.Options) (*coderd. } // Must run after newAPI so options.Database is dbcrypt-wrapped. coderd.BackfillBedrockProviderType(aibridgeInitCtx, options.Database, logger.Named("aibridge.backfill")) - // Must run after BackfillBedrockProviderType; shares aibridgeInitCtx so - // a timeout on the first backfill will skip this one until next startup. - coderd.BackfillChatModelConfigProviderStrings(aibridgeInitCtx, options.Database, logger.Named("aibridge.backfill")) // In-memory aibridge daemon. Registered on coderd so chatd can // dispatch LLM requests via the in-process transport without diff --git a/coderd/ai_providers_backfill.go b/coderd/ai_providers_backfill.go index 9ab2e8c95e..bcb267ffc9 100644 --- a/coderd/ai_providers_backfill.go +++ b/coderd/ai_providers_backfill.go @@ -9,16 +9,12 @@ import ( "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/database/db2sdk" "github.com/coder/coder/v2/coderd/database/dbauthz" - "github.com/coder/coder/v2/codersdk" ) // BackfillBedrockProviderType promotes legacy ai_providers rows stored as // type=anthropic with Bedrock settings to type=bedrock. Must run after newAPI // so options.Database is dbcrypt-wrapped. Idempotent; errors are logged and // startup continues. -// -// BackfillChatModelConfigProviderStrings must run after this function so -// provider types are correct when its JOIN executes. func BackfillBedrockProviderType(ctx context.Context, db database.Store, logger slog.Logger) { //nolint:gocritic // Startup-only backfill; no user actor is present. sysCtx := dbauthz.AsSystemRestricted(ctx) @@ -70,25 +66,3 @@ func BackfillBedrockProviderType(ctx context.Context, db database.Store, logger logger.Info(ctx, "backfilled bedrock provider types", slog.F("count", promoted)) } } - -// BackfillChatModelConfigProviderStrings fixes stale chat_model_configs.provider -// strings left as "anthropic" when the linked provider was promoted from -// type=anthropic to type=bedrock by BackfillBedrockProviderType. Errors are -// logged and startup continues. -func BackfillChatModelConfigProviderStrings(ctx context.Context, db database.Store, logger slog.Logger) { - //nolint:gocritic // Startup-only backfill; no user actor is present. - sysCtx := dbauthz.AsSystemRestricted(ctx) - result, err := db.BackfillChatModelConfigProvider(sysCtx, database.BackfillChatModelConfigProviderParams{ - OldProvider: string(codersdk.AIProviderTypeAnthropic), - NewProvider: string(codersdk.AIProviderTypeBedrock), - }) - if err != nil { - logger.Error(ctx, "backfill chat model config provider strings", slog.Error(err)) - return - } - if result != nil { - if n, _ := result.RowsAffected(); n > 0 { - logger.Info(ctx, "backfilled chat model config provider strings", slog.F("count", n)) - } - } -} diff --git a/coderd/ai_providers_backfill_test.go b/coderd/ai_providers_backfill_test.go index a813078193..891b55b1e7 100644 --- a/coderd/ai_providers_backfill_test.go +++ b/coderd/ai_providers_backfill_test.go @@ -4,7 +4,6 @@ import ( "database/sql" "testing" - "github.com/google/uuid" "github.com/stretchr/testify/require" "go.uber.org/mock/gomock" @@ -193,112 +192,6 @@ func TestBackfillBedrockProviderType(t *testing.T) { require.NoError(t, err) require.Equal(t, database.AIProviderTypeBedrock, goodRow.Type, "valid row alongside unparsable one must still be promoted") }) - - // --- chat_model_configs.provider backfill --- - // These subtests rely on the DB already having type=bedrock providers - // from the provider backfill subtests above. - - t.Run("FixesStaleModelConfigProvider", func(t *testing.T) { - // Simulate a model config created when the linked provider was still - // type=anthropic. The stored provider string is "anthropic" but the - // linked provider row now has type=bedrock. - bedrockProvider := dbgen.AIProvider(t, db, database.AIProvider{ - Type: database.AIProviderTypeBedrock, - Settings: bedrockSettings, - }) - staleConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "anthropic", - AIProviderID: uuid.NullUUID{UUID: bedrockProvider.ID, Valid: true}, - }) - - coderd.BackfillChatModelConfigProviderStrings(ctx, db, logger) - - updated, err := db.GetChatModelConfigByID(ctx, staleConfig.ID) - require.NoError(t, err) - require.Equal(t, "bedrock", updated.Provider, "stale anthropic provider string must be fixed to bedrock") - - // Second run must be a no-op: the same config must still be "bedrock". - coderd.BackfillChatModelConfigProviderStrings(ctx, db, logger) - - updated, err = db.GetChatModelConfigByID(ctx, staleConfig.ID) - require.NoError(t, err) - require.Equal(t, "bedrock", updated.Provider, "provider must remain bedrock after second run") - }) - - t.Run("ModelConfigIdempotent", func(t *testing.T) { - before, err := db.GetChatModelConfigs(ctx) - require.NoError(t, err) - - coderd.BackfillChatModelConfigProviderStrings(ctx, db, logger) - - after, err := db.GetChatModelConfigs(ctx) - require.NoError(t, err) - require.Equal(t, len(before), len(after), "second run must not create or delete rows") - }) - - t.Run("PreservesNonAnthropicModelConfig", func(t *testing.T) { - // A model config with provider="openai" linked to a Bedrock provider - // must not be touched. Only "anthropic" → "bedrock" is in scope. - bedrockProvider := dbgen.AIProvider(t, db, database.AIProvider{ - Type: database.AIProviderTypeBedrock, - Settings: bedrockSettings, - }) - openAIConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", - AIProviderID: uuid.NullUUID{UUID: bedrockProvider.ID, Valid: true}, - }) - - coderd.BackfillChatModelConfigProviderStrings(ctx, db, logger) - - row, err := db.GetChatModelConfigByID(ctx, openAIConfig.ID) - require.NoError(t, err) - require.Equal(t, "openai", row.Provider, "non-anthropic provider string must not be changed") - }) - - t.Run("SkipsModelConfigWithDeletedProvider", func(t *testing.T) { - // Verifies the EXISTS subquery excludes soft-deleted providers. - // The model config provider string must stay "anthropic" because - // the linked provider is deleted and therefore excluded by the - // AND deleted = FALSE condition in the query. - deletedProvider := dbgen.AIProvider(t, db, database.AIProvider{ - Type: database.AIProviderTypeBedrock, - Settings: bedrockSettings, - }) - staleConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "anthropic", - AIProviderID: uuid.NullUUID{UUID: deletedProvider.ID, Valid: true}, - }) - require.NoError(t, db.DeleteAIProviderByID(ctx, deletedProvider.ID)) - - coderd.BackfillChatModelConfigProviderStrings(ctx, db, logger) - - row, err := db.GetChatModelConfigByID(ctx, staleConfig.ID) - require.NoError(t, err) - require.Equal(t, "anthropic", row.Provider, "config linked to deleted provider must not be updated") - }) - - t.Run("SkipsDeletedModelConfig", func(t *testing.T) { - // The SQL query guards on deleted = FALSE. Capture the config ID - // before deletion so we delete the right row regardless of ordering. - bedrockProvider := dbgen.AIProvider(t, db, database.AIProvider{ - Type: database.AIProviderTypeBedrock, - Settings: bedrockSettings, - }) - cfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "anthropic", - AIProviderID: uuid.NullUUID{UUID: bedrockProvider.ID, Valid: true}, - }) - - before, err := db.GetChatModelConfigs(ctx) - require.NoError(t, err) - require.NoError(t, db.DeleteChatModelConfigByID(ctx, cfg.ID)) - - coderd.BackfillChatModelConfigProviderStrings(ctx, db, logger) - - after, err := db.GetChatModelConfigs(ctx) - require.NoError(t, err) - require.Equal(t, len(before)-1, len(after), "deleted config must not reappear after backfill") - }) }) t.Run("ListFailure", func(t *testing.T) { @@ -352,17 +245,4 @@ func TestBackfillBedrockProviderType(t *testing.T) { // ErrNoRows is benign: provider was deleted between list and update. coderd.BackfillBedrockProviderType(ctx, db, testLogger(t)) }) - - t.Run("ModelConfigQueryFailure", func(t *testing.T) { - t.Parallel() - ctx := testutil.Context(t, testutil.WaitShort) - ctrl := gomock.NewController(t) - db := dbmock.NewMockStore(ctrl) - - db.EXPECT(). - BackfillChatModelConfigProvider(gomock.Any(), gomock.Any()). - Return(nil, sql.ErrConnDone) - - coderd.BackfillChatModelConfigProviderStrings(ctx, db, testLogger(t)) - }) } diff --git a/coderd/coderdtest/chat.go b/coderd/coderdtest/chat.go index bf460a5ff0..acaa7352e9 100644 --- a/coderd/coderdtest/chat.go +++ b/coderd/coderdtest/chat.go @@ -73,7 +73,6 @@ func CreateOpenAICompatChatModelConfig( contextLimit := int64(4096) isDefault := true modelConfig, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: TestChatProviderOpenAICompat, AIProviderID: &provider.ID, Model: TestChatModelOpenAICompat, ContextLimit: &contextLimit, diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index a735b85542..20a688c768 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -1724,13 +1724,6 @@ func (q *querier) AutoArchiveInactiveChats(ctx context.Context, arg database.Aut return q.db.AutoArchiveInactiveChats(ctx, arg) } -func (q *querier) BackfillChatModelConfigProvider(ctx context.Context, arg database.BackfillChatModelConfigProviderParams) (sql.Result, error) { - if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil { - return nil, err - } - return q.db.BackfillChatModelConfigProvider(ctx, arg) -} - func (q *querier) BackoffChatDiffStatus(ctx context.Context, arg database.BackoffChatDiffStatusParams) error { // This is a system-level operation used by the gitsync // background worker to reschedule failed refreshes. Same @@ -2113,13 +2106,6 @@ func (q *querier) DeleteChatModelConfigsByAIProviderID(ctx context.Context, aiPr return q.db.DeleteChatModelConfigsByAIProviderID(ctx, aiProviderID) } -func (q *querier) DeleteChatModelConfigsByProvider(ctx context.Context, provider string) error { - if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil { - return err - } - return q.db.DeleteChatModelConfigsByProvider(ctx, provider) -} - func (q *querier) DeleteChatQueuedMessage(ctx context.Context, arg database.DeleteChatQueuedMessageParams) error { chat, err := q.db.GetChatByID(ctx, arg.ChatID) if err != nil { @@ -3677,7 +3663,7 @@ func (q *querier) GetEnabledChatModelConfigByID(ctx context.Context, id uuid.UUI return q.db.GetEnabledChatModelConfigByID(ctx, id) } -func (q *querier) GetEnabledChatModelConfigs(ctx context.Context) ([]database.ChatModelConfig, error) { +func (q *querier) GetEnabledChatModelConfigs(ctx context.Context) ([]database.GetEnabledChatModelConfigsRow, error) { if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceDeploymentConfig); err != nil { return nil, err } diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index 865d075f0b..68d5d222bb 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -657,11 +657,6 @@ func (s *MethodTestSuite) TestChats() { dbm.EXPECT().DeleteChatModelConfigByID(gomock.Any(), id).Return(nil).AnyTimes() check.Args(id).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate) })) - s.Run("DeleteChatModelConfigsByProvider", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { - providerName := "test-provider" - dbm.EXPECT().DeleteChatModelConfigsByProvider(gomock.Any(), providerName).Return(nil).AnyTimes() - check.Args(providerName).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate) - })) s.Run("DeleteChatModelConfigsByAIProviderID", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { providerID := uuid.New() dbm.EXPECT().DeleteChatModelConfigsByAIProviderID(gomock.Any(), providerID).Return(nil).AnyTimes() @@ -1212,10 +1207,10 @@ func (s *MethodTestSuite) TestChats() { check.Args(config.ID).Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns(config) })) s.Run("GetEnabledChatModelConfigs", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { - configA := testutil.Fake(s.T(), faker, database.ChatModelConfig{}) - configB := testutil.Fake(s.T(), faker, database.ChatModelConfig{}) - dbm.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return([]database.ChatModelConfig{configA, configB}, nil).AnyTimes() - check.Args().Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns([]database.ChatModelConfig{configA, configB}) + rowA := testutil.Fake(s.T(), faker, database.GetEnabledChatModelConfigsRow{}) + rowB := testutil.Fake(s.T(), faker, database.GetEnabledChatModelConfigsRow{}) + dbm.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return([]database.GetEnabledChatModelConfigsRow{rowA, rowB}, nil).AnyTimes() + check.Args().Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead).Returns([]database.GetEnabledChatModelConfigsRow{rowA, rowB}) })) s.Run("GetStaleChats", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { @@ -1257,12 +1252,11 @@ func (s *MethodTestSuite) TestChats() { })) s.Run("InsertChatModelConfig", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { arg := database.InsertChatModelConfigParams{ - Provider: "test-provider", Model: "test-model", DisplayName: "Test Model", Enabled: true, } - config := testutil.Fake(s.T(), faker, database.ChatModelConfig{Provider: arg.Provider, Model: arg.Model, DisplayName: arg.DisplayName, Enabled: arg.Enabled}) + config := testutil.Fake(s.T(), faker, database.ChatModelConfig{Model: arg.Model, DisplayName: arg.DisplayName, Enabled: arg.Enabled}) dbm.EXPECT().InsertChatModelConfig(gomock.Any(), arg).Return(config, nil).AnyTimes() check.Args(arg).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate).Returns(config) })) @@ -1515,7 +1509,6 @@ func (s *MethodTestSuite) TestChats() { config := testutil.Fake(s.T(), faker, database.ChatModelConfig{}) arg := database.UpdateChatModelConfigParams{ ID: config.ID, - Provider: "updated-provider", Model: "updated-model", DisplayName: "Updated Model", Enabled: true, @@ -6859,14 +6852,6 @@ func (s *MethodTestSuite) TestAIBridge() { dbm.EXPECT().DeleteAIProviderByID(gomock.Any(), provider.ID).Return(nil).AnyTimes() check.Args(provider.ID).Asserts(rbac.ResourceAIProvider, policy.ActionDelete).Returns() })) - s.Run("BackfillChatModelConfigProvider", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { - arg := database.BackfillChatModelConfigProviderParams{ - OldProvider: "anthropic", - NewProvider: "bedrock", - } - dbm.EXPECT().BackfillChatModelConfigProvider(gomock.Any(), arg).Return(nil, nil).AnyTimes() - check.Args(arg).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate) - })) s.Run("UpdateEncryptedAIProviderSettings", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { provider := testutil.Fake(s.T(), faker, database.AIProvider{}) arg := database.UpdateEncryptedAIProviderSettingsParams{ diff --git a/coderd/database/dbgen/dbgen.go b/coderd/database/dbgen/dbgen.go index 43804ae84c..79b0245c6b 100644 --- a/coderd/database/dbgen/dbgen.go +++ b/coderd/database/dbgen/dbgen.go @@ -160,14 +160,15 @@ const ( func ChatModelConfig(t testing.TB, db database.Store, seed database.ChatModelConfig, munge ...func(*database.InsertChatModelConfigParams)) database.ChatModelConfig { t.Helper() - providerName := takeFirst(seed.Provider, "openai") aiProviderID := seed.AIProviderID if !aiProviderID.Valid { + // No AIProviderID supplied: reuse or create a default openai provider. + // Tests needing a specific provider type should pass seed.AIProviderID. providers, err := db.GetAIProviders(genCtx, database.GetAIProvidersParams{IncludeDisabled: true}) require.NoError(t, err, "get ai providers") var provider database.AIProvider for _, candidate := range providers { - if candidate.Type != database.AIProviderType(providerName) { + if candidate.Type != database.AIProviderTypeOpenai { continue } if provider.ID == uuid.Nil || candidate.CreatedAt.After(provider.CreatedAt) { @@ -176,13 +177,12 @@ func ChatModelConfig(t testing.TB, db database.Store, seed database.ChatModelCon } if provider.ID == uuid.Nil { provider = AIProvider(t, db, database.AIProvider{ - Type: database.AIProviderType(providerName), + Type: database.AIProviderTypeOpenai, }) } aiProviderID = uuid.NullUUID{UUID: provider.ID, Valid: true} } params := database.InsertChatModelConfigParams{ - Provider: providerName, Model: takeFirst(seed.Model, "gpt-4o-mini"), DisplayName: takeFirst(seed.DisplayName, "Test Model"), CreatedBy: seed.CreatedBy, diff --git a/coderd/database/dbgen/dbgen_test.go b/coderd/database/dbgen/dbgen_test.go index a07a9c5881..7765314847 100644 --- a/coderd/database/dbgen/dbgen_test.go +++ b/coderd/database/dbgen/dbgen_test.go @@ -296,7 +296,9 @@ func TestGenerator(t *testing.T) { // Defaults. cfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{}) require.NotEqual(t, uuid.Nil, cfg.ID) - require.Equal(t, "openai", cfg.Provider) + prov, err := db.GetAIProviderByID(context.Background(), cfg.AIProviderID.UUID) + require.NoError(t, err) + require.Equal(t, "openai", string(prov.Type)) require.Equal(t, "gpt-4o-mini", cfg.Model) require.Equal(t, "Test Model", cfg.DisplayName) require.True(t, cfg.Enabled) @@ -304,13 +306,15 @@ func TestGenerator(t *testing.T) { require.Equal(t, int32(70), cfg.CompressionThreshold) // Overrides. - _ = dbgen.ChatProvider(t, db, database.ChatProvider{Provider: "anthropic"}) + anthropicProvider := dbgen.ChatProvider(t, db, database.ChatProvider{Provider: "anthropic"}) cfg2 := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "anthropic", + AIProviderID: uuid.NullUUID{UUID: anthropicProvider.ID, Valid: true}, Model: "claude-4", ContextLimit: 200000, }) - require.Equal(t, "anthropic", cfg2.Provider) + prov2, err := db.GetAIProviderByID(context.Background(), cfg2.AIProviderID.UUID) + require.NoError(t, err) + require.Equal(t, "anthropic", string(prov2.Type)) require.Equal(t, "claude-4", cfg2.Model) require.Equal(t, int64(200000), cfg2.ContextLimit) }) @@ -325,7 +329,7 @@ func TestGenerator(t *testing.T) { OrganizationID: o.ID, }) p := dbgen.ChatProvider(t, db, database.ChatProvider{}) - m := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{Provider: p.Provider}) + m := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{AIProviderID: uuid.NullUUID{UUID: p.ID, Valid: true}}) // Defaults. chat := dbgen.Chat(t, db, database.Chat{ @@ -360,7 +364,7 @@ func TestGenerator(t *testing.T) { OrganizationID: o.ID, }) p := dbgen.ChatProvider(t, db, database.ChatProvider{}) - m := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{Provider: p.Provider}) + m := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{AIProviderID: uuid.NullUUID{UUID: p.ID, Valid: true}}) chat := dbgen.Chat(t, db, database.Chat{ OwnerID: u.ID, OrganizationID: o.ID, diff --git a/coderd/database/dbmetrics/querymetrics.go b/coderd/database/dbmetrics/querymetrics.go index 11e785f8ab..2095b41f5d 100644 --- a/coderd/database/dbmetrics/querymetrics.go +++ b/coderd/database/dbmetrics/querymetrics.go @@ -5,7 +5,6 @@ package dbmetrics import ( "context" - "database/sql" "encoding/json" "slices" "time" @@ -186,14 +185,6 @@ func (m queryMetricsStore) AutoArchiveInactiveChats(ctx context.Context, arg dat return r0, r1 } -func (m queryMetricsStore) BackfillChatModelConfigProvider(ctx context.Context, arg database.BackfillChatModelConfigProviderParams) (sql.Result, error) { - start := time.Now() - r0, r1 := m.s.BackfillChatModelConfigProvider(ctx, arg) - m.queryLatencies.WithLabelValues("BackfillChatModelConfigProvider").Observe(time.Since(start).Seconds()) - m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "BackfillChatModelConfigProvider").Inc() - return r0, r1 -} - func (m queryMetricsStore) BackoffChatDiffStatus(ctx context.Context, arg database.BackoffChatDiffStatusParams) error { start := time.Now() r0 := m.s.BackoffChatDiffStatus(ctx, arg) @@ -538,14 +529,6 @@ func (m queryMetricsStore) DeleteChatModelConfigsByAIProviderID(ctx context.Cont return r0 } -func (m queryMetricsStore) DeleteChatModelConfigsByProvider(ctx context.Context, provider string) error { - start := time.Now() - r0 := m.s.DeleteChatModelConfigsByProvider(ctx, provider) - m.queryLatencies.WithLabelValues("DeleteChatModelConfigsByProvider").Observe(time.Since(start).Seconds()) - m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "DeleteChatModelConfigsByProvider").Inc() - return r0 -} - func (m queryMetricsStore) DeleteChatQueuedMessage(ctx context.Context, arg database.DeleteChatQueuedMessageParams) error { start := time.Now() r0 := m.s.DeleteChatQueuedMessage(ctx, arg) @@ -2018,7 +2001,7 @@ func (m queryMetricsStore) GetEnabledChatModelConfigByID(ctx context.Context, id return r0, r1 } -func (m queryMetricsStore) GetEnabledChatModelConfigs(ctx context.Context) ([]database.ChatModelConfig, error) { +func (m queryMetricsStore) GetEnabledChatModelConfigs(ctx context.Context) ([]database.GetEnabledChatModelConfigsRow, error) { start := time.Now() r0, r1 := m.s.GetEnabledChatModelConfigs(ctx) m.queryLatencies.WithLabelValues("GetEnabledChatModelConfigs").Observe(time.Since(start).Seconds()) diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index 248c26d0a1..4b34989ae5 100644 --- a/coderd/database/dbmock/dbmock.go +++ b/coderd/database/dbmock/dbmock.go @@ -11,7 +11,6 @@ package dbmock import ( context "context" - sql "database/sql" json "encoding/json" reflect "reflect" time "time" @@ -194,21 +193,6 @@ func (mr *MockStoreMockRecorder) AutoArchiveInactiveChats(ctx, arg any) *gomock. return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "AutoArchiveInactiveChats", reflect.TypeOf((*MockStore)(nil).AutoArchiveInactiveChats), ctx, arg) } -// BackfillChatModelConfigProvider mocks base method. -func (m *MockStore) BackfillChatModelConfigProvider(ctx context.Context, arg database.BackfillChatModelConfigProviderParams) (sql.Result, error) { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "BackfillChatModelConfigProvider", ctx, arg) - ret0, _ := ret[0].(sql.Result) - ret1, _ := ret[1].(error) - return ret0, ret1 -} - -// BackfillChatModelConfigProvider indicates an expected call of BackfillChatModelConfigProvider. -func (mr *MockStoreMockRecorder) BackfillChatModelConfigProvider(ctx, arg any) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "BackfillChatModelConfigProvider", reflect.TypeOf((*MockStore)(nil).BackfillChatModelConfigProvider), ctx, arg) -} - // BackoffChatDiffStatus mocks base method. func (m *MockStore) BackoffChatDiffStatus(ctx context.Context, arg database.BackoffChatDiffStatusParams) error { m.ctrl.T.Helper() @@ -877,20 +861,6 @@ func (mr *MockStoreMockRecorder) DeleteChatModelConfigsByAIProviderID(ctx, aiPro return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteChatModelConfigsByAIProviderID", reflect.TypeOf((*MockStore)(nil).DeleteChatModelConfigsByAIProviderID), ctx, aiProviderID) } -// DeleteChatModelConfigsByProvider mocks base method. -func (m *MockStore) DeleteChatModelConfigsByProvider(ctx context.Context, provider string) error { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "DeleteChatModelConfigsByProvider", ctx, provider) - ret0, _ := ret[0].(error) - return ret0 -} - -// DeleteChatModelConfigsByProvider indicates an expected call of DeleteChatModelConfigsByProvider. -func (mr *MockStoreMockRecorder) DeleteChatModelConfigsByProvider(ctx, provider any) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteChatModelConfigsByProvider", reflect.TypeOf((*MockStore)(nil).DeleteChatModelConfigsByProvider), ctx, provider) -} - // DeleteChatQueuedMessage mocks base method. func (m *MockStore) DeleteChatQueuedMessage(ctx context.Context, arg database.DeleteChatQueuedMessageParams) error { m.ctrl.T.Helper() @@ -3733,10 +3703,10 @@ func (mr *MockStoreMockRecorder) GetEnabledChatModelConfigByID(ctx, id any) *gom } // GetEnabledChatModelConfigs mocks base method. -func (m *MockStore) GetEnabledChatModelConfigs(ctx context.Context) ([]database.ChatModelConfig, error) { +func (m *MockStore) GetEnabledChatModelConfigs(ctx context.Context) ([]database.GetEnabledChatModelConfigsRow, error) { m.ctrl.T.Helper() ret := m.ctrl.Call(m, "GetEnabledChatModelConfigs", ctx) - ret0, _ := ret[0].([]database.ChatModelConfig) + ret0, _ := ret[0].([]database.GetEnabledChatModelConfigsRow) ret1, _ := ret[1].(error) return ret0, ret1 } diff --git a/coderd/database/dbpurge/dbpurge_test.go b/coderd/database/dbpurge/dbpurge_test.go index 6d8024c497..f32e800bd3 100644 --- a/coderd/database/dbpurge/dbpurge_test.go +++ b/coderd/database/dbpurge/dbpurge_test.go @@ -2084,7 +2084,6 @@ func TestPurgeChatDebugRuns(t *testing.T) { DisplayName: "OpenAI", }) modelConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", Model: "test-model", ContextLimit: 8192, }) @@ -2311,7 +2310,6 @@ func TestDeleteOldChatFiles(t *testing.T) { DisplayName: "OpenAI", }) mc := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", Model: "test-model", ContextLimit: 8192, }) diff --git a/coderd/database/dump.sql b/coderd/database/dump.sql index a984421d82..35d63873d3 100644 --- a/coderd/database/dump.sql +++ b/coderd/database/dump.sql @@ -1909,7 +1909,6 @@ ALTER SEQUENCE chat_messages_id_seq OWNED BY chat_messages.id; CREATE TABLE chat_model_configs ( id uuid DEFAULT gen_random_uuid() NOT NULL, - provider text NOT NULL, model text NOT NULL, display_name text DEFAULT ''::text NOT NULL, created_by uuid, @@ -4625,10 +4624,6 @@ CREATE INDEX idx_chat_model_configs_ai_provider_id ON chat_model_configs USING b CREATE INDEX idx_chat_model_configs_enabled ON chat_model_configs USING btree (enabled); -CREATE INDEX idx_chat_model_configs_provider ON chat_model_configs USING btree (provider); - -CREATE INDEX idx_chat_model_configs_provider_model ON chat_model_configs USING btree (provider, model); - CREATE UNIQUE INDEX idx_chat_model_configs_single_default ON chat_model_configs USING btree ((1)) WHERE ((is_default = true) AND (deleted = false)); CREATE INDEX idx_chat_queued_messages_chat_id ON chat_queued_messages USING btree (chat_id); diff --git a/coderd/database/migrations/000534_drop_chat_model_configs_provider.down.sql b/coderd/database/migrations/000534_drop_chat_model_configs_provider.down.sql new file mode 100644 index 0000000000..a1fde819d0 --- /dev/null +++ b/coderd/database/migrations/000534_drop_chat_model_configs_provider.down.sql @@ -0,0 +1,13 @@ +ALTER TABLE chat_model_configs ADD COLUMN provider text; + +UPDATE chat_model_configs cmc +SET provider = ap.type::text +FROM ai_providers ap +WHERE ap.id = cmc.ai_provider_id; + +UPDATE chat_model_configs SET provider = '' WHERE provider IS NULL; + +ALTER TABLE chat_model_configs ALTER COLUMN provider SET NOT NULL; + +CREATE INDEX idx_chat_model_configs_provider ON chat_model_configs USING btree (provider); +CREATE INDEX idx_chat_model_configs_provider_model ON chat_model_configs USING btree (provider, model); diff --git a/coderd/database/migrations/000534_drop_chat_model_configs_provider.up.sql b/coderd/database/migrations/000534_drop_chat_model_configs_provider.up.sql new file mode 100644 index 0000000000..d73da7e27b --- /dev/null +++ b/coderd/database/migrations/000534_drop_chat_model_configs_provider.up.sql @@ -0,0 +1,4 @@ +DROP INDEX idx_chat_model_configs_provider; +DROP INDEX idx_chat_model_configs_provider_model; + +ALTER TABLE chat_model_configs DROP COLUMN provider; diff --git a/coderd/database/migrations/testdata/fixtures/000475_chat_model_config_soft_deleted.up.sql b/coderd/database/migrations/testdata/fixtures/000475_chat_model_config_soft_deleted.up.sql new file mode 100644 index 0000000000..bf6c4d26e1 --- /dev/null +++ b/coderd/database/migrations/testdata/fixtures/000475_chat_model_config_soft_deleted.up.sql @@ -0,0 +1,41 @@ +-- Soft-deleted chat model config whose provider never had an ai_providers +-- backfill match, so it reaches later migrations with ai_provider_id IS NULL. +-- +-- This row exercises the 000534 down migration's +-- `UPDATE ... SET provider = '' WHERE provider IS NULL` sweep: its NULL +-- ai_provider_id means the backfill join leaves provider NULL, and the sweep +-- must populate it before `ALTER COLUMN provider SET NOT NULL`. +-- +-- It is inserted at version 000475 (after 000474 dropped the provider foreign +-- key) so the provider value need not reference a chat_providers row, and the +-- 000504/000505 backfill (which matches on `cmc.provider = cp.provider`) skips +-- it. `deleted = TRUE` keeps it out of idx_chat_model_configs_single_default +-- and satisfies chat_model_configs_ai_provider_required_when_active (added in +-- 000505), which permits a NULL ai_provider_id only for deleted rows. +INSERT INTO chat_model_configs ( + id, + provider, + model, + display_name, + enabled, + is_default, + deleted, + deleted_at, + context_limit, + compression_threshold, + created_at, + updated_at +) VALUES ( + 'b3a1d2c4-5e6f-4a7b-8c9d-0e1f2a3b4c5d', + 'legacy-removed', + 'legacy-model', + 'Legacy Soft Deleted', + FALSE, + FALSE, + TRUE, + '2024-01-01 00:00:00+00', + 200000, + 70, + '2024-01-01 00:00:00+00', + '2024-01-01 00:00:00+00' +); diff --git a/coderd/database/models.go b/coderd/database/models.go index 93ae445581..c0dea30800 100644 --- a/coderd/database/models.go +++ b/coderd/database/models.go @@ -4953,7 +4953,6 @@ type ChatMessage struct { type ChatModelConfig struct { ID uuid.UUID `db:"id" json:"id"` - Provider string `db:"provider" json:"provider"` Model string `db:"model" json:"model"` DisplayName string `db:"display_name" json:"display_name"` CreatedBy uuid.NullUUID `db:"created_by" json:"created_by"` diff --git a/coderd/database/querier.go b/coderd/database/querier.go index d3b745e3f7..92dd60ebb5 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -6,7 +6,6 @@ package database import ( "context" - "database/sql" "encoding/json" "time" @@ -71,11 +70,6 @@ type sqlcQuerier interface { // created_at ASC flows through to dbpurge's digest truncation; see // buildDigestData in dbpurge.go for the tradeoff rationale. AutoArchiveInactiveChats(ctx context.Context, arg AutoArchiveInactiveChatsParams) ([]AutoArchiveInactiveChatsRow, error) - // old_provider is matched as text; new_provider is also cast to ai_provider_type - // for the EXISTS check against ai_providers.type. - // ai_provider_id IS NOT NULL is defensive; the check constraint already - // enforces that non-deleted rows always have a provider ID. - BackfillChatModelConfigProvider(ctx context.Context, arg BackfillChatModelConfigProviderParams) (sql.Result, error) BackoffChatDiffStatus(ctx context.Context, arg BackoffChatDiffStatusParams) error // Deletes heartbeat rows for the supplied (chat_id, runner_id) pairs. BatchDeleteChatHeartbeats(ctx context.Context, arg BatchDeleteChatHeartbeatsParams) (int64, error) @@ -150,7 +144,6 @@ type sqlcQuerier interface { DeleteChatDebugDataByChatID(ctx context.Context, arg DeleteChatDebugDataByChatIDParams) (int64, error) DeleteChatModelConfigByID(ctx context.Context, id uuid.UUID) error DeleteChatModelConfigsByAIProviderID(ctx context.Context, aiProviderID uuid.UUID) error - DeleteChatModelConfigsByProvider(ctx context.Context, provider string) error DeleteChatQueuedMessage(ctx context.Context, arg DeleteChatQueuedMessageParams) error // Deletes a queued message, scoped to the parent chat. Returns the // number of affected rows so callers can detect missing rows without @@ -449,6 +442,7 @@ type sqlcQuerier interface { GetChatModelConfigByID(ctx context.Context, id uuid.UUID) (ChatModelConfig, error) GetChatModelConfigs(ctx context.Context) ([]ChatModelConfig, error) // Returns all model configurations for telemetry snapshot collection. + // deleted = false guarantees ai_provider_id is non-null, so INNER JOIN is safe. GetChatModelConfigsForTelemetry(ctx context.Context) ([]GetChatModelConfigsForTelemetryRow, error) // GetChatPersonalModelOverridesEnabled returns whether users may configure // personal chat model overrides. It defaults to false when unset. @@ -539,7 +533,7 @@ type sqlcQuerier interface { // Providers can be disabled independently of their model configs. // Check both to ensure the selected config is actually usable. GetEnabledChatModelConfigByID(ctx context.Context, id uuid.UUID) (ChatModelConfig, error) - GetEnabledChatModelConfigs(ctx context.Context) ([]ChatModelConfig, error) + GetEnabledChatModelConfigs(ctx context.Context) ([]GetEnabledChatModelConfigsRow, error) GetEnabledMCPServerConfigs(ctx context.Context) ([]MCPServerConfig, error) // GetExternalAgentTokensByTemplateID returns the auth tokens for all // non-deleted external agents on the latest build of every running workspace diff --git a/coderd/database/querier_test.go b/coderd/database/querier_test.go index b6f1ca0ecc..2e8924e502 100644 --- a/coderd/database/querier_test.go +++ b/coderd/database/querier_test.go @@ -1250,7 +1250,6 @@ func TestChatContextHydration(t *testing.T) { owner := dbgen.User(t, db, database.User{}) _ = dbgen.ChatProvider(t, db, database.ChatProvider{Provider: "openai", DisplayName: "OpenAI"}) modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", Model: "test-model", CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true}, UpdatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true}, @@ -1381,7 +1380,6 @@ func TestGetAuthorizedChats(t *testing.T) { DisplayName: "OpenAI", }) modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", Model: "test-model", CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true}, UpdatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true}, @@ -1645,7 +1643,6 @@ func TestGetAuthorizedChatsACLSharing(t *testing.T) { dbgen.ChatProvider(t, db, database.ChatProvider{Provider: "openai", DisplayName: "OpenAI"}) modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", Model: "test-model", CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true}, UpdatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true}, @@ -1766,7 +1763,6 @@ func TestGetAuthorizedChatsACLSharingGroupACL(t *testing.T) { dbgen.ChatProvider(t, db, database.ChatProvider{Provider: "openai", DisplayName: "OpenAI"}) modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", Model: "test-model", CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true}, UpdatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true}, @@ -1869,7 +1865,6 @@ func TestGetAuthorizedChatsByChatFileIDACLSharing(t *testing.T) { dbgen.ChatProvider(t, db, database.ChatProvider{Provider: "openai", DisplayName: "OpenAI"}) modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", Model: "test-model", CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true}, UpdatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true}, @@ -11097,24 +11092,21 @@ func TestGetEnabledChatModelConfigsUsesAIProviders(t *testing.T) { params.Enabled = false }) enabledConfig := dbgen.ChatModelConfig(t, store, database.ChatModelConfig{ - Provider: string(enabledProvider.Type), - Model: "openrouter-model-" + uuid.NewString(), + Model: "openrouter-model-" + uuid.NewString(), AIProviderID: uuid.NullUUID{ UUID: enabledProvider.ID, Valid: true, }, }) disabledProviderConfig := dbgen.ChatModelConfig(t, store, database.ChatModelConfig{ - Provider: string(disabledProvider.Type), - Model: "vercel-model-" + uuid.NewString(), + Model: "vercel-model-" + uuid.NewString(), AIProviderID: uuid.NullUUID{ UUID: disabledProvider.ID, Valid: true, }, }) disabledModelConfig := dbgen.ChatModelConfig(t, store, database.ChatModelConfig{ - Provider: string(enabledProvider.Type), - Model: "disabled-model-" + uuid.NewString(), + Model: "disabled-model-" + uuid.NewString(), AIProviderID: uuid.NullUUID{ UUID: enabledProvider.ID, Valid: true, @@ -11125,14 +11117,14 @@ func TestGetEnabledChatModelConfigsUsesAIProviders(t *testing.T) { configs, err := store.GetEnabledChatModelConfigs(ctx) require.NoError(t, err) - require.True(t, slices.ContainsFunc(configs, func(config database.ChatModelConfig) bool { - return config.ID == enabledConfig.ID + require.True(t, slices.ContainsFunc(configs, func(row database.GetEnabledChatModelConfigsRow) bool { + return row.ChatModelConfig.ID == enabledConfig.ID })) - require.False(t, slices.ContainsFunc(configs, func(config database.ChatModelConfig) bool { - return config.ID == disabledProviderConfig.ID + require.False(t, slices.ContainsFunc(configs, func(row database.GetEnabledChatModelConfigsRow) bool { + return row.ChatModelConfig.ID == disabledProviderConfig.ID })) - require.False(t, slices.ContainsFunc(configs, func(config database.ChatModelConfig) bool { - return config.ID == disabledModelConfig.ID + require.False(t, slices.ContainsFunc(configs, func(row database.GetEnabledChatModelConfigsRow) bool { + return row.ChatModelConfig.ID == disabledModelConfig.ID })) config, err := store.GetEnabledChatModelConfigByID(ctx, enabledConfig.ID) @@ -11150,16 +11142,16 @@ func insertChatModelConfigForTest( ctx context.Context, t testing.TB, store database.Store, + providerType string, params database.InsertChatModelConfigParams, ) (database.ChatModelConfig, error) { t.Helper() if params.AIProviderID.Valid { return store.InsertChatModelConfig(ctx, params) } - providerName := params.Provider + providerName := providerType if providerName == "" { providerName = "openai" - params.Provider = providerName } providers, err := store.GetAIProviders(ctx, database.GetAIProvidersParams{IncludeDisabled: true}) if err != nil { @@ -11198,8 +11190,7 @@ func TestInsertChatMessages(t *testing.T) { ) database.ChatModelConfig { t.Helper() - modelConfig, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{ - Provider: provider, + modelConfig, err := insertChatModelConfigForTest(ctx, t, store, provider, database.InsertChatModelConfigParams{ Model: model, DisplayName: displayName, CreatedBy: uuid.NullUUID{UUID: userID, Valid: true}, @@ -11410,8 +11401,7 @@ func TestGetChatMessagesForPromptByChatID(t *testing.T) { APIKey: "test-key", }) - modelCfg, err := insertChatModelConfigForTest(ctx, t, db, database.InsertChatModelConfigParams{ - Provider: "openai", + modelCfg, err := insertChatModelConfigForTest(ctx, t, db, "openai", database.InsertChatModelConfigParams{ AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, Model: "test-model", DisplayName: "Test Model", @@ -11786,8 +11776,7 @@ func TestChatPinOrderQueries(t *testing.T) { CentralApiKeyEnabled: true, }) - modelCfg, err := insertChatModelConfigForTest(bg, t, db, database.InsertChatModelConfigParams{ - Provider: "openai", + modelCfg, err := insertChatModelConfigForTest(bg, t, db, "openai", database.InsertChatModelConfigParams{ Model: "test-model", DisplayName: "Test Model", CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true}, @@ -11966,8 +11955,7 @@ func TestChatPinOrderConstraints(t *testing.T) { CentralApiKeyEnabled: true, }) - modelCfg, err := insertChatModelConfigForTest(bg, t, db, database.InsertChatModelConfigParams{ - Provider: "openai", + modelCfg, err := insertChatModelConfigForTest(bg, t, db, "openai", database.InsertChatModelConfigParams{ Model: "test-model", DisplayName: "Test Model", CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true}, @@ -12058,8 +12046,7 @@ func TestChatLabels(t *testing.T) { CentralApiKeyEnabled: true, }) - modelCfg, err := insertChatModelConfigForTest(ctx, t, db, database.InsertChatModelConfigParams{ - Provider: "openai", + modelCfg, err := insertChatModelConfigForTest(ctx, t, db, "openai", database.InsertChatModelConfigParams{ Model: "test-model", DisplayName: "Test Model", CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true}, @@ -12357,8 +12344,7 @@ func TestUpdateChatLastTurnSummary(t *testing.T) { CentralApiKeyEnabled: true, }) - modelCfg, err := insertChatModelConfigForTest(ctx, t, db, database.InsertChatModelConfigParams{ - Provider: "openai", + modelCfg, err := insertChatModelConfigForTest(ctx, t, db, "openai", database.InsertChatModelConfigParams{ Model: "test-model", DisplayName: "Test Model", CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true}, @@ -12494,8 +12480,7 @@ func TestDeleteChatDebugDataAfterMessageIDIncludesTriggeredRuns(t *testing.T) { CentralApiKeyEnabled: true, }) - modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{ - Provider: providerName, + modelCfg, err := insertChatModelConfigForTest(ctx, t, store, providerName, database.InsertChatModelConfigParams{ Model: modelName, DisplayName: "Debug Model", CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, @@ -12686,8 +12671,7 @@ func TestDeleteChatDebugDataAfterMessageIDStepLevelFieldBoundariesAndNulls(t *te CentralApiKeyEnabled: true, }) - modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{ - Provider: providerName, + modelCfg, err := insertChatModelConfigForTest(ctx, t, store, providerName, database.InsertChatModelConfigParams{ Model: modelName, DisplayName: "Debug Model", CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, @@ -12943,8 +12927,7 @@ func TestFinalizeStaleChatDebugRows(t *testing.T) { CentralApiKeyEnabled: true, }) - modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{ - Provider: providerName, + modelCfg, err := insertChatModelConfigForTest(ctx, t, store, providerName, database.InsertChatModelConfigParams{ Model: modelName, DisplayName: "Debug Model", CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, @@ -13381,8 +13364,7 @@ func TestChatDebugSQLGuards(t *testing.T) { CentralApiKeyEnabled: true, }) - modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{ - Provider: providerName, + modelCfg, err := insertChatModelConfigForTest(ctx, t, store, providerName, database.InsertChatModelConfigParams{ Model: modelName, DisplayName: "Debug Model", CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, @@ -13514,8 +13496,7 @@ func TestChatDebugRunCOALESCEPreservation(t *testing.T) { CentralApiKeyEnabled: true, }) - modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{ - Provider: providerName, + modelCfg, err := insertChatModelConfigForTest(ctx, t, store, providerName, database.InsertChatModelConfigParams{ Model: modelName, DisplayName: "Debug Model", CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, @@ -13628,8 +13609,7 @@ func TestChatDebugStepCOALESCEPreservation(t *testing.T) { CentralApiKeyEnabled: true, }) - modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{ - Provider: providerName, + modelCfg, err := insertChatModelConfigForTest(ctx, t, store, providerName, database.InsertChatModelConfigParams{ Model: modelName, DisplayName: "Debug Model", CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, @@ -13752,8 +13732,7 @@ func TestDeleteChatDebugDataAfterMessageIDNullMessagesSurvive(t *testing.T) { CentralApiKeyEnabled: true, }) - modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{ - Provider: providerName, + modelCfg, err := insertChatModelConfigForTest(ctx, t, store, providerName, database.InsertChatModelConfigParams{ Model: modelName, DisplayName: "Debug Model", CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, @@ -13849,8 +13828,7 @@ func TestDeleteChatDebugDataAfterMessageIDStartedBeforeFiltersNewerRuns(t *testi CentralApiKeyEnabled: true, }) - modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{ - Provider: providerName, + modelCfg, err := insertChatModelConfigForTest(ctx, t, store, providerName, database.InsertChatModelConfigParams{ Model: modelName, DisplayName: "Debug Model", CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, @@ -13960,8 +13938,7 @@ func TestDeleteChatDebugDataByChatIDStartedBeforeFiltersNewerRuns(t *testing.T) CentralApiKeyEnabled: true, }) - modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{ - Provider: providerName, + modelCfg, err := insertChatModelConfigForTest(ctx, t, store, providerName, database.InsertChatModelConfigParams{ Model: modelName, DisplayName: "Debug Model", CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, @@ -14045,7 +14022,6 @@ func TestGetChatsFilter(t *testing.T) { }, "test-key") modelCfg, err := store.InsertChatModelConfig(ctx, database.InsertChatModelConfigParams{ - Provider: "openai", AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, Model: "test-model-" + uuid.NewString(), DisplayName: "Test Model", @@ -14343,8 +14319,7 @@ func TestChatHasUnread(t *testing.T) { CentralApiKeyEnabled: true, }) - modelCfg, err := insertChatModelConfigForTest(ctx, t, store, database.InsertChatModelConfigParams{ - Provider: "openai", + modelCfg, err := insertChatModelConfigForTest(ctx, t, store, "openai", database.InsertChatModelConfigParams{ Model: "test-model-" + uuid.NewString(), DisplayName: "Test Model", CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index 64fe9b53dd..6e41d273b4 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -5025,37 +5025,6 @@ func (q *sqlQuerier) InsertChatFile(ctx context.Context, arg InsertChatFileParam return i, err } -const backfillChatModelConfigProvider = `-- name: BackfillChatModelConfigProvider :execresult -UPDATE - chat_model_configs -SET - provider = $1::text, - updated_at = NOW() -WHERE - provider = $2::text - AND deleted = FALSE - AND ai_provider_id IS NOT NULL - AND EXISTS ( - SELECT 1 FROM ai_providers - WHERE id = chat_model_configs.ai_provider_id - AND type = $1::ai_provider_type - AND deleted = FALSE - ) -` - -type BackfillChatModelConfigProviderParams struct { - NewProvider string `db:"new_provider" json:"new_provider"` - OldProvider string `db:"old_provider" json:"old_provider"` -} - -// old_provider is matched as text; new_provider is also cast to ai_provider_type -// for the EXISTS check against ai_providers.type. -// ai_provider_id IS NOT NULL is defensive; the check constraint already -// enforces that non-deleted rows always have a provider ID. -func (q *sqlQuerier) BackfillChatModelConfigProvider(ctx context.Context, arg BackfillChatModelConfigProviderParams) (sql.Result, error) { - return q.db.ExecContext(ctx, backfillChatModelConfigProvider, arg.NewProvider, arg.OldProvider) -} - const deleteChatModelConfigByID = `-- name: DeleteChatModelConfigByID :exec UPDATE chat_model_configs @@ -5089,26 +5058,9 @@ func (q *sqlQuerier) DeleteChatModelConfigsByAIProviderID(ctx context.Context, a return err } -const deleteChatModelConfigsByProvider = `-- name: DeleteChatModelConfigsByProvider :exec -UPDATE - chat_model_configs -SET - deleted = TRUE, - deleted_at = NOW(), - updated_at = NOW() -WHERE - provider = $1::text - AND deleted = FALSE -` - -func (q *sqlQuerier) DeleteChatModelConfigsByProvider(ctx context.Context, provider string) error { - _, err := q.db.ExecContext(ctx, deleteChatModelConfigsByProvider, provider) - return err -} - const getChatModelConfigByID = `-- name: GetChatModelConfigByID :one SELECT - id, provider, model, display_name, created_by, updated_by, enabled, is_default, deleted, deleted_at, created_at, updated_at, context_limit, compression_threshold, options, ai_provider_id + id, model, display_name, created_by, updated_by, enabled, is_default, deleted, deleted_at, created_at, updated_at, context_limit, compression_threshold, options, ai_provider_id FROM chat_model_configs WHERE @@ -5121,7 +5073,6 @@ func (q *sqlQuerier) GetChatModelConfigByID(ctx context.Context, id uuid.UUID) ( var i ChatModelConfig err := row.Scan( &i.ID, - &i.Provider, &i.Model, &i.DisplayName, &i.CreatedBy, @@ -5142,16 +5093,18 @@ func (q *sqlQuerier) GetChatModelConfigByID(ctx context.Context, id uuid.UUID) ( const getChatModelConfigs = `-- name: GetChatModelConfigs :many SELECT - id, provider, model, display_name, created_by, updated_by, enabled, is_default, deleted, deleted_at, created_at, updated_at, context_limit, compression_threshold, options, ai_provider_id + cmc.id, cmc.model, cmc.display_name, cmc.created_by, cmc.updated_by, cmc.enabled, cmc.is_default, cmc.deleted, cmc.deleted_at, cmc.created_at, cmc.updated_at, cmc.context_limit, cmc.compression_threshold, cmc.options, cmc.ai_provider_id FROM - chat_model_configs + chat_model_configs cmc +LEFT JOIN + ai_providers ap ON ap.id = cmc.ai_provider_id WHERE - deleted = FALSE + cmc.deleted = FALSE ORDER BY - provider ASC, - model ASC, - updated_at DESC, - id DESC + ap.type::text ASC, + cmc.model ASC, + cmc.updated_at DESC, + cmc.id DESC ` func (q *sqlQuerier) GetChatModelConfigs(ctx context.Context) ([]ChatModelConfig, error) { @@ -5165,7 +5118,6 @@ func (q *sqlQuerier) GetChatModelConfigs(ctx context.Context) ([]ChatModelConfig var i ChatModelConfig if err := rows.Scan( &i.ID, - &i.Provider, &i.Model, &i.DisplayName, &i.CreatedBy, @@ -5196,7 +5148,7 @@ func (q *sqlQuerier) GetChatModelConfigs(ctx context.Context) ([]ChatModelConfig const getDefaultChatModelConfig = `-- name: GetDefaultChatModelConfig :one SELECT - id, provider, model, display_name, created_by, updated_by, enabled, is_default, deleted, deleted_at, created_at, updated_at, context_limit, compression_threshold, options, ai_provider_id + id, model, display_name, created_by, updated_by, enabled, is_default, deleted, deleted_at, created_at, updated_at, context_limit, compression_threshold, options, ai_provider_id FROM chat_model_configs WHERE @@ -5209,7 +5161,6 @@ func (q *sqlQuerier) GetDefaultChatModelConfig(ctx context.Context) (ChatModelCo var i ChatModelConfig err := row.Scan( &i.ID, - &i.Provider, &i.Model, &i.DisplayName, &i.CreatedBy, @@ -5230,7 +5181,7 @@ func (q *sqlQuerier) GetDefaultChatModelConfig(ctx context.Context) (ChatModelCo const getEnabledChatModelConfigByID = `-- name: GetEnabledChatModelConfigByID :one SELECT - cmc.id, cmc.provider, cmc.model, cmc.display_name, cmc.created_by, cmc.updated_by, cmc.enabled, cmc.is_default, cmc.deleted, cmc.deleted_at, cmc.created_at, cmc.updated_at, cmc.context_limit, cmc.compression_threshold, cmc.options, cmc.ai_provider_id + cmc.id, cmc.model, cmc.display_name, cmc.created_by, cmc.updated_by, cmc.enabled, cmc.is_default, cmc.deleted, cmc.deleted_at, cmc.created_at, cmc.updated_at, cmc.context_limit, cmc.compression_threshold, cmc.options, cmc.ai_provider_id FROM chat_model_configs cmc JOIN @@ -5250,7 +5201,6 @@ func (q *sqlQuerier) GetEnabledChatModelConfigByID(ctx context.Context, id uuid. var i ChatModelConfig err := row.Scan( &i.ID, - &i.Provider, &i.Model, &i.DisplayName, &i.CreatedBy, @@ -5271,7 +5221,8 @@ func (q *sqlQuerier) GetEnabledChatModelConfigByID(ctx context.Context, id uuid. const getEnabledChatModelConfigs = `-- name: GetEnabledChatModelConfigs :many SELECT - cmc.id, cmc.provider, cmc.model, cmc.display_name, cmc.created_by, cmc.updated_by, cmc.enabled, cmc.is_default, cmc.deleted, cmc.deleted_at, cmc.created_at, cmc.updated_at, cmc.context_limit, cmc.compression_threshold, cmc.options, cmc.ai_provider_id + cmc.id, cmc.model, cmc.display_name, cmc.created_by, cmc.updated_by, cmc.enabled, cmc.is_default, cmc.deleted, cmc.deleted_at, cmc.created_at, cmc.updated_at, cmc.context_limit, cmc.compression_threshold, cmc.options, cmc.ai_provider_id, + ap.type::text AS provider FROM chat_model_configs cmc JOIN @@ -5282,38 +5233,43 @@ WHERE AND ap.enabled = TRUE AND ap.deleted = FALSE ORDER BY - cmc.provider ASC, + ap.type::text ASC, cmc.model ASC, cmc.updated_at DESC, cmc.id DESC ` -func (q *sqlQuerier) GetEnabledChatModelConfigs(ctx context.Context) ([]ChatModelConfig, error) { +type GetEnabledChatModelConfigsRow struct { + ChatModelConfig ChatModelConfig `db:"chat_model_config" json:"chat_model_config"` + Provider string `db:"provider" json:"provider"` +} + +func (q *sqlQuerier) GetEnabledChatModelConfigs(ctx context.Context) ([]GetEnabledChatModelConfigsRow, error) { rows, err := q.db.QueryContext(ctx, getEnabledChatModelConfigs) if err != nil { return nil, err } defer rows.Close() - var items []ChatModelConfig + var items []GetEnabledChatModelConfigsRow for rows.Next() { - var i ChatModelConfig + var i GetEnabledChatModelConfigsRow if err := rows.Scan( - &i.ID, + &i.ChatModelConfig.ID, + &i.ChatModelConfig.Model, + &i.ChatModelConfig.DisplayName, + &i.ChatModelConfig.CreatedBy, + &i.ChatModelConfig.UpdatedBy, + &i.ChatModelConfig.Enabled, + &i.ChatModelConfig.IsDefault, + &i.ChatModelConfig.Deleted, + &i.ChatModelConfig.DeletedAt, + &i.ChatModelConfig.CreatedAt, + &i.ChatModelConfig.UpdatedAt, + &i.ChatModelConfig.ContextLimit, + &i.ChatModelConfig.CompressionThreshold, + &i.ChatModelConfig.Options, + &i.ChatModelConfig.AIProviderID, &i.Provider, - &i.Model, - &i.DisplayName, - &i.CreatedBy, - &i.UpdatedBy, - &i.Enabled, - &i.IsDefault, - &i.Deleted, - &i.DeletedAt, - &i.CreatedAt, - &i.UpdatedAt, - &i.ContextLimit, - &i.CompressionThreshold, - &i.Options, - &i.AIProviderID, ); err != nil { return nil, err } @@ -5330,7 +5286,6 @@ func (q *sqlQuerier) GetEnabledChatModelConfigs(ctx context.Context) ([]ChatMode const insertChatModelConfig = `-- name: InsertChatModelConfig :one INSERT INTO chat_model_configs ( - provider, model, display_name, created_by, @@ -5344,22 +5299,20 @@ INSERT INTO chat_model_configs ( ) VALUES ( $1::text, $2::text, - $3::text, + $3::uuid, $4::uuid, - $5::uuid, + $5::boolean, $6::boolean, - $7::boolean, - $8::bigint, - $9::integer, - $10::jsonb, - $11::uuid + $7::bigint, + $8::integer, + $9::jsonb, + $10::uuid ) RETURNING - id, provider, model, display_name, created_by, updated_by, enabled, is_default, deleted, deleted_at, created_at, updated_at, context_limit, compression_threshold, options, ai_provider_id + id, model, display_name, created_by, updated_by, enabled, is_default, deleted, deleted_at, created_at, updated_at, context_limit, compression_threshold, options, ai_provider_id ` type InsertChatModelConfigParams struct { - Provider string `db:"provider" json:"provider"` Model string `db:"model" json:"model"` DisplayName string `db:"display_name" json:"display_name"` CreatedBy uuid.NullUUID `db:"created_by" json:"created_by"` @@ -5374,7 +5327,6 @@ type InsertChatModelConfigParams struct { func (q *sqlQuerier) InsertChatModelConfig(ctx context.Context, arg InsertChatModelConfigParams) (ChatModelConfig, error) { row := q.db.QueryRowContext(ctx, insertChatModelConfig, - arg.Provider, arg.Model, arg.DisplayName, arg.CreatedBy, @@ -5389,7 +5341,6 @@ func (q *sqlQuerier) InsertChatModelConfig(ctx context.Context, arg InsertChatMo var i ChatModelConfig err := row.Scan( &i.ID, - &i.Provider, &i.Model, &i.DisplayName, &i.CreatedBy, @@ -5428,26 +5379,24 @@ const updateChatModelConfig = `-- name: UpdateChatModelConfig :one UPDATE chat_model_configs SET - provider = $1::text, - model = $2::text, - display_name = $3::text, - updated_by = $4::uuid, - enabled = $5::boolean, - is_default = $6::boolean, - context_limit = $7::bigint, - compression_threshold = $8::integer, - options = $9::jsonb, - ai_provider_id = $10::uuid, + model = $1::text, + display_name = $2::text, + updated_by = $3::uuid, + enabled = $4::boolean, + is_default = $5::boolean, + context_limit = $6::bigint, + compression_threshold = $7::integer, + options = $8::jsonb, + ai_provider_id = $9::uuid, updated_at = NOW() WHERE - id = $11::uuid + id = $10::uuid AND deleted = FALSE RETURNING - id, provider, model, display_name, created_by, updated_by, enabled, is_default, deleted, deleted_at, created_at, updated_at, context_limit, compression_threshold, options, ai_provider_id + id, model, display_name, created_by, updated_by, enabled, is_default, deleted, deleted_at, created_at, updated_at, context_limit, compression_threshold, options, ai_provider_id ` type UpdateChatModelConfigParams struct { - Provider string `db:"provider" json:"provider"` Model string `db:"model" json:"model"` DisplayName string `db:"display_name" json:"display_name"` UpdatedBy uuid.NullUUID `db:"updated_by" json:"updated_by"` @@ -5462,7 +5411,6 @@ type UpdateChatModelConfigParams struct { func (q *sqlQuerier) UpdateChatModelConfig(ctx context.Context, arg UpdateChatModelConfigParams) (ChatModelConfig, error) { row := q.db.QueryRowContext(ctx, updateChatModelConfig, - arg.Provider, arg.Model, arg.DisplayName, arg.UpdatedBy, @@ -5477,7 +5425,6 @@ func (q *sqlQuerier) UpdateChatModelConfig(ctx context.Context, arg UpdateChatMo var i ChatModelConfig err := row.Scan( &i.ID, - &i.Provider, &i.Model, &i.DisplayName, &i.CreatedBy, @@ -6967,7 +6914,7 @@ const getChatCostPerModel = `-- name: GetChatCostPerModel :many SELECT cmc.id AS model_config_id, cmc.display_name, - cmc.provider, + COALESCE(ap.type::text, '')::text AS provider, cmc.model, COALESCE(SUM(cm.total_cost_micros), 0)::bigint AS total_cost_micros, COUNT(*) FILTER ( @@ -6988,13 +6935,15 @@ JOIN chats c ON c.id = cm.chat_id JOIN chat_model_configs cmc ON cmc.id = cm.model_config_id +LEFT JOIN + ai_providers ap ON ap.id = cmc.ai_provider_id WHERE c.owner_id = $1::uuid AND cm.role = 'assistant' AND cm.created_at >= $2::timestamptz AND cm.created_at < $3::timestamptz GROUP BY - cmc.id, cmc.display_name, cmc.provider, cmc.model + cmc.id, cmc.display_name, ap.type, cmc.model ORDER BY total_cost_micros DESC ` @@ -7960,9 +7909,10 @@ func (q *sqlQuerier) GetChatMessagesForPromptByChatID(ctx context.Context, chatI } const getChatModelConfigsForTelemetry = `-- name: GetChatModelConfigsForTelemetry :many -SELECT id, provider, model, context_limit, enabled, is_default -FROM chat_model_configs -WHERE deleted = false +SELECT cmc.id, ap.type::text AS provider, cmc.model, cmc.context_limit, cmc.enabled, cmc.is_default +FROM chat_model_configs cmc +JOIN ai_providers ap ON ap.id = cmc.ai_provider_id +WHERE cmc.deleted = false ` type GetChatModelConfigsForTelemetryRow struct { @@ -7975,6 +7925,7 @@ type GetChatModelConfigsForTelemetryRow struct { } // Returns all model configurations for telemetry snapshot collection. +// deleted = false guarantees ai_provider_id is non-null, so INNER JOIN is safe. func (q *sqlQuerier) GetChatModelConfigsForTelemetry(ctx context.Context) ([]GetChatModelConfigsForTelemetryRow, error) { rows, err := q.db.QueryContext(ctx, getChatModelConfigsForTelemetry) if err != nil { diff --git a/coderd/database/queries/chatmodelconfigs.sql b/coderd/database/queries/chatmodelconfigs.sql index 4284521e1b..ae95083991 100644 --- a/coderd/database/queries/chatmodelconfigs.sql +++ b/coderd/database/queries/chatmodelconfigs.sql @@ -18,20 +18,23 @@ WHERE -- name: GetChatModelConfigs :many SELECT - * + cmc.* FROM - chat_model_configs + chat_model_configs cmc +LEFT JOIN + ai_providers ap ON ap.id = cmc.ai_provider_id WHERE - deleted = FALSE + cmc.deleted = FALSE ORDER BY - provider ASC, - model ASC, - updated_at DESC, - id DESC; + ap.type::text ASC, + cmc.model ASC, + cmc.updated_at DESC, + cmc.id DESC; -- name: GetEnabledChatModelConfigs :many SELECT - cmc.* + sqlc.embed(cmc), + ap.type::text AS provider FROM chat_model_configs cmc JOIN @@ -42,7 +45,7 @@ WHERE AND ap.enabled = TRUE AND ap.deleted = FALSE ORDER BY - cmc.provider ASC, + ap.type::text ASC, cmc.model ASC, cmc.updated_at DESC, cmc.id DESC; @@ -65,7 +68,6 @@ WHERE -- name: InsertChatModelConfig :one INSERT INTO chat_model_configs ( - provider, model, display_name, created_by, @@ -77,7 +79,6 @@ INSERT INTO chat_model_configs ( options, ai_provider_id ) VALUES ( - @provider::text, @model::text, @display_name::text, sqlc.narg('created_by')::uuid, @@ -96,7 +97,6 @@ RETURNING UPDATE chat_model_configs SET - provider = @provider::text, model = @model::text, display_name = @display_name::text, updated_by = sqlc.narg('updated_by')::uuid, @@ -133,38 +133,6 @@ SET WHERE id = @id::uuid; --- name: DeleteChatModelConfigsByProvider :exec -UPDATE - chat_model_configs -SET - deleted = TRUE, - deleted_at = NOW(), - updated_at = NOW() -WHERE - provider = @provider::text - AND deleted = FALSE; - --- name: BackfillChatModelConfigProvider :execresult --- old_provider is matched as text; new_provider is also cast to ai_provider_type --- for the EXISTS check against ai_providers.type. --- ai_provider_id IS NOT NULL is defensive; the check constraint already --- enforces that non-deleted rows always have a provider ID. -UPDATE - chat_model_configs -SET - provider = @new_provider::text, - updated_at = NOW() -WHERE - provider = @old_provider::text - AND deleted = FALSE - AND ai_provider_id IS NOT NULL - AND EXISTS ( - SELECT 1 FROM ai_providers - WHERE id = chat_model_configs.ai_provider_id - AND type = @new_provider::ai_provider_type - AND deleted = FALSE - ); - -- name: DeleteChatModelConfigsByAIProviderID :exec UPDATE chat_model_configs diff --git a/coderd/database/queries/chats.sql b/coderd/database/queries/chats.sql index a62ce06438..b7bbe8ecb3 100644 --- a/coderd/database/queries/chats.sql +++ b/coderd/database/queries/chats.sql @@ -2220,7 +2220,7 @@ WHERE SELECT cmc.id AS model_config_id, cmc.display_name, - cmc.provider, + COALESCE(ap.type::text, '')::text AS provider, cmc.model, COALESCE(SUM(cm.total_cost_micros), 0)::bigint AS total_cost_micros, COUNT(*) FILTER ( @@ -2241,13 +2241,15 @@ JOIN chats c ON c.id = cm.chat_id JOIN chat_model_configs cmc ON cmc.id = cm.model_config_id +LEFT JOIN + ai_providers ap ON ap.id = cmc.ai_provider_id WHERE c.owner_id = @owner_id::uuid AND cm.role = 'assistant' AND cm.created_at >= @start_date::timestamptz AND cm.created_at < @end_date::timestamptz GROUP BY - cmc.id, cmc.display_name, cmc.provider, cmc.model + cmc.id, cmc.display_name, ap.type, cmc.model ORDER BY total_cost_micros DESC; @@ -2580,9 +2582,11 @@ GROUP BY cm.chat_id; -- name: GetChatModelConfigsForTelemetry :many -- Returns all model configurations for telemetry snapshot collection. -SELECT id, provider, model, context_limit, enabled, is_default -FROM chat_model_configs -WHERE deleted = false; +-- deleted = false guarantees ai_provider_id is non-null, so INNER JOIN is safe. +SELECT cmc.id, ap.type::text AS provider, cmc.model, cmc.context_limit, cmc.enabled, cmc.is_default +FROM chat_model_configs cmc +JOIN ai_providers ap ON ap.id = cmc.ai_provider_id +WHERE cmc.deleted = false; -- name: GetActiveChatsByAgentID :many SELECT * FROM chats_expanded diff --git a/coderd/exp_chats.go b/coderd/exp_chats.go index 442ca10729..40d1d19223 100644 --- a/coderd/exp_chats.go +++ b/coderd/exp_chats.go @@ -774,7 +774,7 @@ func (api *API) chatPersonalModelOverrideDeploymentDefaults( type userChatModelAvailability struct { configuredProviders []chatprovider.ConfiguredProvider configuredModels []chatprovider.ConfiguredModel - enabledModels []database.ChatModelConfig + enabledModels []database.GetEnabledChatModelConfigsRow providerStatus map[string]chatprovider.ProviderAvailability providerStatusByID map[uuid.UUID]chatprovider.ProviderAvailability enabledProviderNames map[string]struct{} @@ -886,8 +886,8 @@ func (api *API) getUserChatProviderAvailability( if normalizedProvider == "" { continue } - if model.AIProviderID.Valid { - status, ok := availability.providerStatusByID[model.AIProviderID.UUID] + if model.ChatModelConfig.AIProviderID.Valid { + status, ok := availability.providerStatusByID[model.ChatModelConfig.AIProviderID.UUID] if ok { mergeProviderStatus(modelStatusByType, normalizedProvider, status) } @@ -904,8 +904,8 @@ func (api *API) getUserChatProviderAvailability( for _, model := range enabledModels { normalizedProvider := chatprovider.NormalizeProvider(model.Provider) - if model.AIProviderID.Valid { - status, ok := availability.providerStatusByID[model.AIProviderID.UUID] + if model.ChatModelConfig.AIProviderID.Valid { + status, ok := availability.providerStatusByID[model.ChatModelConfig.AIProviderID.UUID] if !ok { continue } @@ -915,8 +915,8 @@ func (api *API) getUserChatProviderAvailability( } availability.configuredModels = append(availability.configuredModels, chatprovider.ConfiguredModel{ Provider: model.Provider, - Model: model.Model, - DisplayName: model.DisplayName, + Model: model.ChatModelConfig.Model, + DisplayName: model.ChatModelConfig.DisplayName, }) } return availability, nil @@ -966,21 +966,10 @@ func (api *API) userCanUseChatModelConfig( } return chatModelConfigAvailable, nil } - provider, _, err := chatprovider.ResolveModelWithProviderHint(model.Model, model.Provider) - if err != nil { - return chatModelConfigUnavailableProviderDisabled, nil - } - if _, ok := availability.enabledProviderNames[provider]; !ok { - return chatModelConfigUnavailableProviderDisabled, nil - } - providerStatus, ok := availability.providerStatus[provider] - if !ok { - return chatModelConfigUnavailableProviderDisabled, nil - } - if !providerStatus.Available { - return chatModelConfigUnavailableCredentialsMissing, nil - } - return chatModelConfigAvailable, nil + // Active configs always carry a provider FK (CHECK + // chat_model_configs_ai_provider_required_when_active), so an unset FK + // means the config is not usable. + return chatModelConfigUnavailableModelNotFoundOrDisabled, nil } func (api *API) validateUserChatModelConfigAvailable( @@ -6828,7 +6817,12 @@ func (api *API) listChatModelConfigs(rw http.ResponseWriter, r *http.Request) { configs, err = api.Database.GetChatModelConfigs(ctx) } else { //nolint:gocritic // All authenticated users need to read enabled model configs to use the chat feature. - configs, err = api.Database.GetEnabledChatModelConfigs(dbauthz.AsChatd(ctx)) + rows, rowsErr := api.Database.GetEnabledChatModelConfigs(dbauthz.AsChatd(ctx)) + err = rowsErr + configs = make([]database.ChatModelConfig, 0, len(rows)) + for _, row := range rows { + configs = append(configs, row.ChatModelConfig) + } } if err != nil { httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ @@ -6900,7 +6894,6 @@ func (api *API) createChatModelConfig(rw http.ResponseWriter, r *http.Request) { httpapi.Write(ctx, rw, http.StatusPreconditionFailed, codersdk.Response{Message: "AI provider is disabled."}) return } - provider := string(aiProvider.Type) aiProviderID := uuid.NullUUID{UUID: aiProvider.ID, Valid: true} model := strings.TrimSpace(req.Model) @@ -6956,7 +6949,6 @@ func (api *API) createChatModelConfig(rw http.ResponseWriter, r *http.Request) { } insertParams := database.InsertChatModelConfigParams{ - Provider: provider, Model: model, DisplayName: strings.TrimSpace(req.DisplayName), Enabled: enabled, @@ -6982,7 +6974,6 @@ func (api *API) createChatModelConfig(rw http.ResponseWriter, r *http.Request) { if !lockedAIProvider.Enabled { return errChatProviderNotConfigured } - insertParams.Provider = string(lockedAIProvider.Type) if err := validateChatModelConfigProviderModel(lockedAIProvider, insertParams.Model); err != nil { return err } @@ -7087,19 +7078,6 @@ func (api *API) updateChatModelConfig(rw http.ResponseWriter, r *http.Request) { return } - if strings.TrimSpace(req.Provider) != "" && req.AIProviderID == nil { - requestedProvider := chatprovider.NormalizeProvider(req.Provider) - if requestedProvider == "" { - httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{Message: "Invalid provider."}) - return - } - if requestedProvider != existing.Provider { - httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{Message: "AI provider ID is required when updating provider."}) - return - } - } - - provider := existing.Provider aiProviderID := existing.AIProviderID if req.AIProviderID != nil { //nolint:gocritic // The route already authorized chat model config updates. @@ -7119,7 +7097,6 @@ func (api *API) updateChatModelConfig(rw http.ResponseWriter, r *http.Request) { httpapi.Write(ctx, rw, http.StatusPreconditionFailed, codersdk.Response{Message: "AI provider is disabled."}) return } - provider = string(aiProvider.Type) aiProviderID = uuid.NullUUID{UUID: aiProvider.ID, Valid: true} } @@ -7179,7 +7156,6 @@ func (api *API) updateChatModelConfig(rw http.ResponseWriter, r *http.Request) { } updateParams := database.UpdateChatModelConfigParams{ - Provider: provider, Model: model, DisplayName: displayName, Enabled: enabled, @@ -7208,7 +7184,6 @@ func (api *API) updateChatModelConfig(rw http.ResponseWriter, r *http.Request) { if !aiProvider.Enabled { return errChatProviderNotConfigured } - updateParams.Provider = string(aiProvider.Type) if err := validateChatModelConfigProviderModel(aiProvider, updateParams.Model); err != nil { return err } @@ -7388,7 +7363,6 @@ func chatModelConfigToUpdateParams( config database.ChatModelConfig, ) database.UpdateChatModelConfigParams { return database.UpdateChatModelConfigParams{ - Provider: config.Provider, Model: config.Model, DisplayName: config.DisplayName, Enabled: config.Enabled, @@ -7458,14 +7432,11 @@ func parseChatModelConfigID(rw http.ResponseWriter, r *http.Request) (uuid.UUID, } func convertChatModelConfig(config database.ChatModelConfig) codersdk.ChatModelConfig { - var aiProviderID *uuid.UUID - if config.AIProviderID.Valid { - aiProviderID = &config.AIProviderID.UUID - } + // Active configs always carry a non-null ai_provider_id (CHECK + // chat_model_configs_ai_provider_required_when_active). return codersdk.ChatModelConfig{ ID: config.ID, - Provider: config.Provider, - AIProviderID: aiProviderID, + AIProviderID: config.AIProviderID.UUID, Model: config.Model, DisplayName: config.DisplayName, Enabled: config.Enabled, diff --git a/coderd/exp_chats_test.go b/coderd/exp_chats_test.go index 4e98eab3ae..75b59fa603 100644 --- a/coderd/exp_chats_test.go +++ b/coderd/exp_chats_test.go @@ -1696,7 +1696,7 @@ func TestListChatModels(t *testing.T) { var openAIProvider *codersdk.ChatModelProvider for i := range models.Providers { - if models.Providers[i].Provider == modelConfig.Provider { + if models.Providers[i].Provider == coderdtest.TestChatProviderOpenAICompat { openAIProvider = &models.Providers[i] break } @@ -1706,7 +1706,7 @@ func TestListChatModels(t *testing.T) { foundModel := false for _, model := range openAIProvider.Models { - if model.Provider == modelConfig.Provider && model.Model == modelConfig.Model { + if model.Provider == coderdtest.TestChatProviderOpenAICompat && model.Model == modelConfig.Model { foundModel = true break } @@ -1772,14 +1772,14 @@ func TestListChatModels(t *testing.T) { ctx := testutil.Context(t, testutil.WaitLong) client := newChatClient(t) _ = coderdtest.CreateFirstUser(t, client.Client) - modelConfig := createChatModelConfig(t, client) + _ = createChatModelConfig(t, client) models, err := client.ListChatModels(ctx) require.NoError(t, err) var openAIProvider *codersdk.ChatModelProvider for i := range models.Providers { - if models.Providers[i].Provider == modelConfig.Provider { + if models.Providers[i].Provider == coderdtest.TestChatProviderOpenAICompat { openAIProvider = &models.Providers[i] break } @@ -1800,7 +1800,6 @@ func TestListChatModels(t *testing.T) { contextLimit := int64(4096) _, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: string(providerType), AIProviderID: &provider.ID, Model: "claude-sonnet", ContextLimit: &contextLimit, @@ -1851,7 +1850,6 @@ func TestListChatModels(t *testing.T) { contextLimit := int64(4096) _, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: "google", AIProviderID: &provider.ID, Model: "gemini-1.5-pro", ContextLimit: &contextLimit, @@ -1903,7 +1901,6 @@ func TestListChatModels(t *testing.T) { contextLimit := int64(4096) _, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: "openai", AIProviderID: &provider.ID, Model: "gpt-4o-mini", ContextLimit: &contextLimit, @@ -2397,14 +2394,14 @@ func TestListChatProviders(t *testing.T) { ctx := testutil.Context(t, testutil.WaitLong) client := newChatClient(t) _ = coderdtest.CreateFirstUser(t, client.Client) - modelConfig := createChatModelConfig(t, client) + _ = createChatModelConfig(t, client) providers, err := client.ListChatProviders(ctx) require.NoError(t, err) var openAIProvider *codersdk.ChatProviderConfig for i := range providers { - if providers[i].Provider == modelConfig.Provider { + if providers[i].Provider == coderdtest.TestChatProviderOpenAICompat { openAIProvider = &providers[i] break } @@ -3542,7 +3539,7 @@ func TestListChatModelConfigs(t *testing.T) { for _, config := range configs { if config.ID == modelConfig.ID { found = true - require.Equal(t, modelConfig.Provider, config.Provider) + require.Equal(t, modelConfig.AIProviderID, config.AIProviderID) require.Equal(t, modelConfig.Model, config.Model) require.True(t, config.IsDefault) } @@ -3562,7 +3559,6 @@ func TestListChatModelConfigs(t *testing.T) { contextLimit := int64(4096) enabled := false disabledConfig, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: "openai", AIProviderID: &aiProvider.ID, Model: "gpt-4o-disabled", DisplayName: "GPT-4o Disabled", @@ -3599,8 +3595,7 @@ func TestListChatModelConfigs(t *testing.T) { contextLimit := int64(4096) enabled := false _, err := adminClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: enabledConfig.Provider, - AIProviderID: enabledConfig.AIProviderID, + AIProviderID: &enabledConfig.AIProviderID, Model: "gpt-4o-disabled", DisplayName: "GPT-4o Disabled", Enabled: &enabled, @@ -3626,7 +3621,6 @@ func TestListChatModelConfigs(t *testing.T) { legacyOptions := json.RawMessage(`{"input_price_per_million_tokens":0.15,"output_price_per_million_tokens":0.6,"cache_read_price_per_million_tokens":0.03,"cache_write_price_per_million_tokens":0.3}`) storedConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", AIProviderID: uuid.NullUUID{UUID: aiProvider.ID, Valid: true}, Model: "gpt-4o-mini-legacy", DisplayName: "GPT-4o Mini Legacy", @@ -3670,7 +3664,7 @@ func TestListChatModelConfigs(t *testing.T) { for _, config := range configs { if config.ID == modelConfig.ID { found = true - require.Equal(t, modelConfig.Provider, config.Provider) + require.Equal(t, modelConfig.AIProviderID, config.AIProviderID) require.Equal(t, modelConfig.Model, config.Model) } } @@ -3701,7 +3695,6 @@ func TestCreateChatModelConfig(t *testing.T) { }, } modelConfig, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: "openai", AIProviderID: &aiProvider.ID, Model: "gpt-4o-mini", ContextLimit: &contextLimit, @@ -3710,7 +3703,7 @@ func TestCreateChatModelConfig(t *testing.T) { }) require.NoError(t, err) require.NotEqual(t, uuid.Nil, modelConfig.ID) - require.Equal(t, "openai", modelConfig.Provider) + require.Equal(t, aiProvider.ID, modelConfig.AIProviderID) require.Equal(t, "gpt-4o-mini", modelConfig.Model) require.EqualValues(t, 4096, modelConfig.ContextLimit) require.True(t, modelConfig.IsDefault) @@ -3733,7 +3726,6 @@ func TestCreateChatModelConfig(t *testing.T) { contextLimit := int64(4096) _, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: "openai", AIProviderID: &aiProvider.ID, Model: "gpt-4o-mini", ContextLimit: &contextLimit, @@ -3761,7 +3753,6 @@ func TestCreateChatModelConfig(t *testing.T) { aiProvider := createAIProviderForTest(t, client, "openai", "test-api-key") _, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: "openai", AIProviderID: &aiProvider.ID, Model: "gpt-4o-mini", }) @@ -3778,7 +3769,6 @@ func TestCreateChatModelConfig(t *testing.T) { contextLimit := int64(4096) _, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: "openai", Model: "gpt-4o-mini", ContextLimit: &contextLimit, }) @@ -3796,7 +3786,6 @@ func TestCreateChatModelConfig(t *testing.T) { contextLimit := int64(4096) missingProviderID := uuid.New() _, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: "openai", AIProviderID: &missingProviderID, Model: "gpt-4o-mini", ContextLimit: &contextLimit, @@ -3826,9 +3815,8 @@ func TestCreateChatModelConfig(t *testing.T) { ContextLimit: &contextLimit, }) require.NoError(t, err) - require.Equal(t, "openai", modelConfig.Provider) - require.NotNil(t, modelConfig.AIProviderID) - require.Equal(t, provider.ID, *modelConfig.AIProviderID) + require.NotEqual(t, uuid.Nil, modelConfig.AIProviderID) + require.Equal(t, provider.ID, modelConfig.AIProviderID) }) t.Run("AIProviderIDNotConfigured", func(t *testing.T) { @@ -3913,7 +3901,6 @@ func TestCreateChatModelConfig(t *testing.T) { contextLimit := int64(4096) _, err := memberClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: "openai", AIProviderID: &aiProvider.ID, Model: "gpt-4o-mini", ContextLimit: &contextLimit, @@ -3968,14 +3955,12 @@ func TestUpdateChatModelConfig(t *testing.T) { modelConfig := createChatModelConfig(t, client) updated, err := client.UpdateChatModelConfig(ctx, modelConfig.ID, codersdk.UpdateChatModelConfigRequest{ - Provider: modelConfig.Provider, - Model: "gpt-4o-mini-updated", + Model: "gpt-4o-mini-updated", }) require.NoError(t, err) require.Equal(t, modelConfig.ID, updated.ID) - require.Equal(t, modelConfig.Provider, updated.Provider) - require.NotNil(t, updated.AIProviderID) - require.Equal(t, *modelConfig.AIProviderID, *updated.AIProviderID) + require.NotEqual(t, uuid.Nil, updated.AIProviderID) + require.Equal(t, modelConfig.AIProviderID, updated.AIProviderID) require.Equal(t, "gpt-4o-mini-updated", updated.Model) }) @@ -4028,7 +4013,6 @@ func TestUpdateChatModelConfig(t *testing.T) { require.NoError(t, err) modelConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: string(database.AIProviderTypeOpenai), Model: "anthropic/claude-opus-4.6", AIProviderID: uuid.NullUUID{UUID: aiProvider.ID, Valid: true}, }) @@ -4132,7 +4116,6 @@ func TestUpdateChatModelConfig(t *testing.T) { contextLimit := int64(4096) enabled := false modelConfig, err := adminClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: "openai", AIProviderID: &aiProvider.ID, Model: "gpt-4o-reenable", DisplayName: "GPT-4o Re-enable", @@ -4218,9 +4201,8 @@ func TestUpdateChatModelConfig(t *testing.T) { Model: "claude-3-5-sonnet-latest", }) require.NoError(t, err) - require.Equal(t, "anthropic", updated.Provider) - require.NotNil(t, updated.AIProviderID) - require.Equal(t, provider.ID, *updated.AIProviderID) + require.NotEqual(t, uuid.Nil, updated.AIProviderID) + require.Equal(t, provider.ID, updated.AIProviderID) }) t.Run("UpdateProviderPreservesAIProviderIDWhenTypeUnchanged", func(t *testing.T) { @@ -4243,15 +4225,14 @@ func TestUpdateChatModelConfig(t *testing.T) { Model: "claude-3-5-sonnet-latest", }) require.NoError(t, err) - require.NotNil(t, updated.AIProviderID) + require.NotEqual(t, uuid.Nil, updated.AIProviderID) updated, err = client.UpdateChatModelConfig(ctx, modelConfig.ID, codersdk.UpdateChatModelConfigRequest{ - Provider: "anthropic", - Model: "claude-3-5-haiku-latest", + Model: "claude-3-5-haiku-latest", }) require.NoError(t, err) - require.NotNil(t, updated.AIProviderID) - require.Equal(t, provider.ID, *updated.AIProviderID) + require.NotEqual(t, uuid.Nil, updated.AIProviderID) + require.Equal(t, provider.ID, updated.AIProviderID) }) t.Run("UpdateAIProviderIDNotConfigured", func(t *testing.T) { @@ -4354,7 +4335,6 @@ func TestUpdateChatModelConfig(t *testing.T) { contextLimit := int64(4096) isDefault := false candidateConfig, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: "anthropic", AIProviderID: &aiProvider.ID, Model: "claude-3-5-sonnet", ContextLimit: &contextLimit, @@ -6635,7 +6615,7 @@ func TestSendMessageWithModelOverrideUpdatesLastModelConfigID(t *testing.T) { client, db := newChatClientWithDatabase(t) user := coderdtest.CreateFirstUser(t, client.Client) modelConfigA := createChatModelConfig(t, client) - modelConfigB := createAdditionalChatModelConfig(t, client, modelConfigA.Provider, "gpt-4o-mini-override-"+uuid.NewString()) + modelConfigB := createAdditionalChatModelConfig(t, client, coderdtest.TestChatProviderOpenAICompat, "gpt-4o-mini-override-"+uuid.NewString()) chat := dbgen.Chat(t, db, database.Chat{ OrganizationID: user.OrganizationID, @@ -6679,7 +6659,7 @@ func TestSendMessageQueuesEffectiveModelConfigID(t *testing.T) { client, db := newChatClientWithDatabase(t) user := coderdtest.CreateFirstUser(t, client.Client) modelConfigA := createChatModelConfig(t, client) - modelConfigB := createAdditionalChatModelConfig(t, client, modelConfigA.Provider, "gpt-4o-mini-queued-"+uuid.NewString()) + modelConfigB := createAdditionalChatModelConfig(t, client, coderdtest.TestChatProviderOpenAICompat, "gpt-4o-mini-queued-"+uuid.NewString()) chat := dbgen.Chat(t, db, database.Chat{ OrganizationID: user.OrganizationID, @@ -6730,7 +6710,7 @@ func TestQueuedMessageWithoutOverrideCapturesEnqueueTimeModel(t *testing.T) { client, db := newChatClientWithDatabase(t) user := coderdtest.CreateFirstUser(t, client.Client) modelConfigA := createChatModelConfig(t, client) - modelConfigB := createAdditionalChatModelConfig(t, client, modelConfigA.Provider, "gpt-4o-mini-later-"+uuid.NewString()) + modelConfigB := createAdditionalChatModelConfig(t, client, coderdtest.TestChatProviderOpenAICompat, "gpt-4o-mini-later-"+uuid.NewString()) chat := dbgen.Chat(t, db, database.Chat{ OrganizationID: user.OrganizationID, @@ -6824,7 +6804,7 @@ func TestWatchChatsStatusChangeCarriesUpdatedLastModelConfigID(t *testing.T) { client, db := newChatClientWithDatabase(t) user := coderdtest.CreateFirstUser(t, client.Client) modelConfigA := createChatModelConfig(t, client) - modelConfigB := createAdditionalChatModelConfig(t, client, modelConfigA.Provider, "gpt-4o-mini-watch-direct-"+uuid.NewString()) + modelConfigB := createAdditionalChatModelConfig(t, client, coderdtest.TestChatProviderOpenAICompat, "gpt-4o-mini-watch-direct-"+uuid.NewString()) chat := dbgen.Chat(t, db, database.Chat{ OrganizationID: user.OrganizationID, @@ -6857,7 +6837,7 @@ func TestWatchChatsStatusChangeCarriesUpdatedLastModelConfigID(t *testing.T) { client, db := newChatClientWithDatabase(t) user := coderdtest.CreateFirstUser(t, client.Client) modelConfigA := createChatModelConfig(t, client) - modelConfigB := createAdditionalChatModelConfig(t, client, modelConfigA.Provider, "gpt-4o-mini-watch-promote-"+uuid.NewString()) + modelConfigB := createAdditionalChatModelConfig(t, client, coderdtest.TestChatProviderOpenAICompat, "gpt-4o-mini-watch-promote-"+uuid.NewString()) chat := dbgen.Chat(t, db, database.Chat{ OrganizationID: user.OrganizationID, @@ -8177,7 +8157,7 @@ func TestPatchChatMessage(t *testing.T) { overrideModel := createAdditionalChatModelConfig( t, client, - defaultModel.Provider, + coderdtest.TestChatProviderOpenAICompat, "gpt-4o-mini-edit-override", ) @@ -10910,7 +10890,6 @@ func createAdditionalChatModelConfig( contextLimit := int64(4096) isDefault := false modelConfig, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: provider, AIProviderID: &aiProvider.ID, Model: model, ContextLimit: &contextLimit, @@ -11566,13 +11545,13 @@ func TestChatModelOverrides(t *testing.T) { openAIModel := createAdditionalChatModelConfig( t, adminClient, - defaultModel.Provider, + coderdtest.TestChatProviderOpenAICompat, "gpt-4.1-mini-"+string(setting.context), ) disabledModel := createDisabledChatModelConfig( t, adminClient, - defaultModel.Provider, + coderdtest.TestChatProviderOpenAICompat, "gpt-4.1-disabled-"+string(setting.context), ) memberClientRaw, _ := coderdtest.CreateAnotherUser(t, adminClient.Client, firstUser.OrganizationID) @@ -11784,7 +11763,7 @@ func TestUserChatPersonalModelOverrides(t *testing.T) { noKeyClient := codersdk.NewExperimentalClient(noKeyClientRaw) defaultModelConfig := createChatModelConfig(t, adminClient) - provider := enableUserChatProviderKey(t, adminClient, memberClient, defaultModelConfig.Provider) + provider := enableUserChatProviderKey(t, adminClient, memberClient, coderdtest.TestChatProviderOpenAICompat) modelProvider := createAIProviderForTest(t, adminClient, "anthropic", "") _, err := memberClient.UpsertUserAIProviderKey(ctx, "me", modelProvider.ID, codersdk.CreateUserAIProviderKeyRequest{ APIKey: "test-user-api-key-" + uuid.NewString(), @@ -11792,7 +11771,6 @@ func TestUserChatPersonalModelOverrides(t *testing.T) { require.NoError(t, err) contextLimit := int64(4096) modelConfig, err := adminClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: "anthropic", AIProviderID: &modelProvider.ID, Model: "claude-personal-" + uuid.NewString(), ContextLimit: &contextLimit, @@ -11810,13 +11788,12 @@ func TestUserChatPersonalModelOverrides(t *testing.T) { disabledModelConfig := createDisabledChatModelConfig( t, adminClient, - defaultModelConfig.Provider, + coderdtest.TestChatProviderOpenAICompat, "gpt-4o-personal-disabled-"+uuid.NewString(), ) disabledProvider := createAIProviderForTest(t, adminClient, "google", "test-api-key") contextLimit = int64(4096) disabledProviderModelConfig, err := adminClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: "google", AIProviderID: &disabledProvider.ID, Model: "gemini-personal-disabled-provider-" + uuid.NewString(), ContextLimit: &contextLimit, @@ -12153,7 +12130,7 @@ func TestCreateChatPersonalModelOverrideRoot(t *testing.T) { adminClient, db := newChatClientWithDatabase(t) firstUser := coderdtest.CreateFirstUser(t, adminClient.Client) defaultModel := createChatModelConfig(t, adminClient) - _ = enableUserChatProviderKey(t, adminClient, adminClient, defaultModel.Provider) + _ = enableUserChatProviderKey(t, adminClient, adminClient, coderdtest.TestChatProviderOpenAICompat) overrideProvider := createAIProviderForTest(t, adminClient, "anthropic", "") _, err := adminClient.UpsertUserAIProviderKey(ctx, "me", overrideProvider.ID, codersdk.CreateUserAIProviderKeyRequest{ APIKey: "test-user-api-key-" + uuid.NewString(), @@ -12161,7 +12138,6 @@ func TestCreateChatPersonalModelOverrideRoot(t *testing.T) { require.NoError(t, err) contextLimit := int64(4096) overrideModel, err := adminClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: "anthropic", AIProviderID: &overrideProvider.ID, Model: "claude-root-personal-" + uuid.NewString(), ContextLimit: &contextLimit, @@ -12170,7 +12146,7 @@ func TestCreateChatPersonalModelOverrideRoot(t *testing.T) { disabledModel := createDisabledChatModelConfig( t, adminClient, - defaultModel.Provider, + coderdtest.TestChatProviderOpenAICompat, "gpt-4o-root-personal-disabled-"+uuid.NewString(), ) memberClientRaw, member := coderdtest.CreateAnotherUser( diff --git a/coderd/telemetry/telemetry_test.go b/coderd/telemetry/telemetry_test.go index 85773f326a..ef82029b01 100644 --- a/coderd/telemetry/telemetry_test.go +++ b/coderd/telemetry/telemetry_test.go @@ -1600,18 +1600,18 @@ func TestChatsTelemetry(t *testing.T) { user := dbgen.User(t, db, database.User{}) // Create chat providers (required FK for model configs). - _ = dbgen.ChatProvider(t, db, database.ChatProvider{ + anthropicProvider := dbgen.ChatProvider(t, db, database.ChatProvider{ Provider: "anthropic", DisplayName: "Anthropic", }) - _ = dbgen.ChatProvider(t, db, database.ChatProvider{ + openaiProvider := dbgen.ChatProvider(t, db, database.ChatProvider{ Provider: "openai", DisplayName: "OpenAI", }) // Create a model config. modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "anthropic", + AIProviderID: uuid.NullUUID{UUID: anthropicProvider.ID, Valid: true}, Model: "claude-sonnet-4-20250514", DisplayName: "Claude Sonnet", IsDefault: true, @@ -1620,14 +1620,14 @@ func TestChatsTelemetry(t *testing.T) { // Create a second model config to test full dump. modelCfg2 := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", - Model: "gpt-4o", - DisplayName: "GPT-4o", + AIProviderID: uuid.NullUUID{UUID: openaiProvider.ID, Valid: true}, + Model: "gpt-4o", + DisplayName: "GPT-4o", }) // Create a soft-deleted model config — should NOT appear in telemetry. deletedCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "anthropic", + AIProviderID: uuid.NullUUID{UUID: anthropicProvider.ID, Valid: true}, Model: "claude-deleted", DisplayName: "Deleted Model", ContextLimit: 100000, @@ -1948,13 +1948,13 @@ func TestChatDiffStatusSummaryTelemetry(t *testing.T) { org, err := db.GetDefaultOrganization(ctx) require.NoError(t, err) - _ = dbgen.ChatProvider(t, db, database.ChatProvider{ + anthropicProvider := dbgen.ChatProvider(t, db, database.ChatProvider{ Provider: "anthropic", DisplayName: "Anthropic", }) modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "anthropic", + AIProviderID: uuid.NullUUID{UUID: anthropicProvider.ID, Valid: true}, Model: "claude-sonnet-4-20250514", DisplayName: "Claude Sonnet", IsDefault: true, diff --git a/coderd/workspaceagents_active_chat_internal_test.go b/coderd/workspaceagents_active_chat_internal_test.go index 1f4d317085..a366cb16ee 100644 --- a/coderd/workspaceagents_active_chat_internal_test.go +++ b/coderd/workspaceagents_active_chat_internal_test.go @@ -118,7 +118,6 @@ func insertAgentChatTestModelConfig( }) return dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, CreatedBy: createdBy, UpdatedBy: createdBy, diff --git a/coderd/x/chatd/advisor_internal_test.go b/coderd/x/chatd/advisor_internal_test.go index e76c287787..5744fbfd1d 100644 --- a/coderd/x/chatd/advisor_internal_test.go +++ b/coderd/x/chatd/advisor_internal_test.go @@ -248,7 +248,6 @@ func TestResolveAdvisorModelOverride(t *testing.T) { getEnabledChatModelConfigByID: func(context.Context, uuid.UUID) (database.ChatModelConfig, error) { return database.ChatModelConfig{ ID: configID, - Provider: "openai", Model: "gpt-5.2", Enabled: true, CreatedAt: time.Unix(0, 0).UTC(), @@ -283,7 +282,6 @@ func TestResolveAdvisorModelOverride(t *testing.T) { getEnabledChatModelConfigByID: func(context.Context, uuid.UUID) (database.ChatModelConfig, error) { return database.ChatModelConfig{ ID: configID, - Provider: "openai", Model: "gpt-5.2", Enabled: true, CreatedAt: time.Unix(0, 0).UTC(), @@ -322,6 +320,7 @@ func TestResolveAdvisorModelOverride(t *testing.T) { t.Parallel() ctx := testutil.Context(t, testutil.WaitShort) configID := uuid.New() + providerID := uuid.New() rawOptions, err := json.Marshal(codersdk.ChatModelCallConfig{ Temperature: func() *float64 { v := 0.42; return &v }(), }) @@ -329,16 +328,29 @@ func TestResolveAdvisorModelOverride(t *testing.T) { store := &advisorOverrideStubStore{ getEnabledChatModelConfigByID: func(context.Context, uuid.UUID) (database.ChatModelConfig, error) { return database.ChatModelConfig{ - ID: configID, - Provider: "openai", - Model: "gpt-5.2", - Enabled: true, - CreatedAt: time.Unix(0, 0).UTC(), - UpdatedAt: time.Unix(0, 0).UTC(), - Options: rawOptions, - DisplayName: "gpt-5.2", + ID: configID, + Model: "gpt-5.2", + Enabled: true, + CreatedAt: time.Unix(0, 0).UTC(), + UpdatedAt: time.Unix(0, 0).UTC(), + Options: rawOptions, + DisplayName: "gpt-5.2", + AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true}, }, nil }, + getAIProviderByID: func(context.Context, uuid.UUID) (database.AIProvider, error) { + return database.AIProvider{ + ID: providerID, + Type: database.AIProviderTypeOpenai, + Enabled: true, + }, nil + }, + getAIProviderKeysByProviderID: func(context.Context, uuid.UUID) ([]database.AIProviderKey, error) { + return []database.AIProviderKey{{ + ProviderID: providerID, + APIKey: "sk-test", + }}, nil + }, } p := newAdvisorTestServer(ctx, t, store) @@ -373,7 +385,6 @@ func TestResolveAdvisorModelOverride(t *testing.T) { getEnabledChatModelConfigByID: func(context.Context, uuid.UUID) (database.ChatModelConfig, error) { return database.ChatModelConfig{ ID: configID, - Provider: "openai", Model: "gpt-5.2", Enabled: true, CreatedAt: time.Unix(0, 0).UTC(), @@ -426,7 +437,6 @@ func TestResolveAdvisorModelOverridePromotesAIBridgeErrors(t *testing.T) { getEnabledChatModelConfigByID: func(context.Context, uuid.UUID) (database.ChatModelConfig, error) { return database.ChatModelConfig{ ID: configID, - Provider: "openai", Model: "gpt-5.2", Enabled: true, DisplayName: "gpt-5.2", diff --git a/coderd/x/chatd/chatd.go b/coderd/x/chatd/chatd.go index 4c8da40fd7..46482209ca 100644 --- a/coderd/x/chatd/chatd.go +++ b/coderd/x/chatd/chatd.go @@ -2588,6 +2588,16 @@ func (p *Server) prepareManualTitleDebugRun( finishDebugRun := func(error) {} route, routeErr := p.resolveModelRouteForConfig(ctx, chat.OwnerID, modelConfig, keys) + var routeProvider string + if routeErr == nil { + routeProvider, _ = route.providerHint() + } else if modelConfig.AIProviderID.Valid { + // Route resolution failed, but the linked provider still identifies the + // type for the debug run record. Best-effort: leave empty if disabled. + if provider, err := p.enabledAIProviderByID(ctx, modelConfig.AIProviderID.UUID); err == nil { + routeProvider = string(provider.Type) + } + } debugOpts := modelOpts debugOpts.RecordHTTP = true var debugModelErr error @@ -2606,21 +2616,19 @@ func (p *Server) prepareManualTitleDebugRun( case debugModelErr != nil: p.logger.Warn(ctx, "failed to create debug-aware manual title model", slog.F("chat_id", chat.ID), - slog.F("provider", modelConfig.Provider), slog.F("model", modelConfig.Model), slog.Error(debugModelErr), ) case debugModel == nil: p.logger.Warn(ctx, "manual title debug model creation returned nil", slog.F("chat_id", chat.ID), - slog.F("provider", modelConfig.Provider), slog.F("model", modelConfig.Model), ) default: titleModel = chatdebug.WrapModel(debugModel, debugSvc, chatdebug.RecorderOptions{ ChatID: chat.ID, OwnerID: chat.OwnerID, - Provider: modelConfig.Provider, + Provider: routeProvider, Model: modelConfig.Model, }) } @@ -2651,7 +2659,7 @@ func (p *Server) prepareManualTitleDebugRun( debugRun, createRunErr := debugSvc.CreateRun(createRunCtx, chatdebug.CreateRunParams{ ChatID: chat.ID, ModelConfigID: modelConfig.ID, - Provider: modelConfig.Provider, + Provider: routeProvider, Model: modelConfig.Model, Kind: chatdebug.KindTitleGeneration, Status: chatdebug.StatusInProgress, @@ -2663,7 +2671,6 @@ func (p *Server) prepareManualTitleDebugRun( if createRunErr != nil { p.logger.Warn(ctx, "failed to create manual title debug run", slog.F("chat_id", chat.ID), - slog.F("provider", modelConfig.Provider), slog.F("model", modelConfig.Model), slog.Error(createRunErr), ) @@ -2789,7 +2796,6 @@ func (p *Server) resolveManualTitleModel( if err != nil { p.logger.Debug(ctx, "manual title preferred model unavailable", slog.F("chat_id", chat.ID), - slog.F("provider", config.Provider), slog.F("model", config.Model), slog.Error(err), ) @@ -2804,7 +2810,6 @@ func (p *Server) resolveManualTitleModel( if err != nil { p.logger.Debug(ctx, "manual title preferred model unavailable", slog.F("chat_id", chat.ID), - slog.F("provider", config.Provider), slog.F("model", config.Model), slog.Error(err), ) diff --git a/coderd/x/chatd/chatd_chainmode_test.go b/coderd/x/chatd/chatd_chainmode_test.go index 00af33e1fe..60f4b2e90a 100644 --- a/coderd/x/chatd/chatd_chainmode_test.go +++ b/coderd/x/chatd/chatd_chainmode_test.go @@ -326,7 +326,6 @@ func seedAnthropicChatDependencies(t *testing.T, db database.Store, baseURL stri }) dbgen.AIProviderKey(t, db, database.AIProviderKey{ProviderID: provider.ID}) model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "anthropic", Model: "claude-sonnet-4-20250514", IsDefault: true, AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, @@ -453,7 +452,6 @@ func updateModelForChainMode(t *testing.T, db database.Store, model database.Cha ID: model.ID, DisplayName: model.DisplayName, Model: model.Model, - Provider: model.Provider, Enabled: model.Enabled, ContextLimit: model.ContextLimit, CompressionThreshold: model.CompressionThreshold, diff --git a/coderd/x/chatd/chatd_internal_test.go b/coderd/x/chatd/chatd_internal_test.go index caaaebf82f..89d9a2ca75 100644 --- a/coderd/x/chatd/chatd_internal_test.go +++ b/coderd/x/chatd/chatd_internal_test.go @@ -26,6 +26,7 @@ import ( coderdpubsub "github.com/coder/coder/v2/coderd/pubsub" "github.com/coder/coder/v2/coderd/rbac" "github.com/coder/coder/v2/coderd/workspacestats" + "github.com/coder/coder/v2/coderd/x/chatd/chatdebug" "github.com/coder/coder/v2/coderd/x/chatd/chatloop" openaicomputeruse "github.com/coder/coder/v2/coderd/x/chatd/chatopenai/computeruse" "github.com/coder/coder/v2/coderd/x/chatd/chatprompt" @@ -794,11 +795,12 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts(t *testing.T) { WorkerID: uuid.NullUUID{UUID: workerID, Valid: true}, Title: fallbackChatTitle(userPrompt), } + providerID := uuid.New() modelConfig := database.ChatModelConfig{ ID: modelConfigID, - Provider: "openai", Model: "gpt-4o-mini", ContextLimit: 8192, + AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true}, } updatedChat := chat updatedChat.Title = wantTitle @@ -832,14 +834,20 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts(t *testing.T) { } db.EXPECT().GetChatModelConfigByID(gomock.Any(), modelConfigID).Return(modelConfig, nil) - providerID := uuid.New() + db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(database.AIProvider{ + ID: providerID, + Type: database.AIProviderTypeOpenai, + Enabled: true, + BaseUrl: serverURL, + }, nil).AnyTimes() db.EXPECT().GetAIProviders(gomock.Any(), gomock.Any()).Return([]database.AIProvider{{ ID: providerID, Type: database.AIProviderTypeOpenai, Enabled: true, BaseUrl: serverURL, - }}, nil) - db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), []uuid.UUID{providerID}).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil) + }}, nil).AnyTimes() + db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes() + db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), gomock.Any()).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes() db.EXPECT().GetChatUsageLimitConfig(gomock.Any()).Return(database.ChatUsageLimitConfig{}, sql.ErrNoRows) db.EXPECT().GetChatMessagesByChatIDAscPaginated( gomock.Any(), @@ -957,11 +965,12 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts_IdleChatReleasesManualLock(t lockedChat := chat lockedChat.WorkerID = uuid.NullUUID{UUID: manualTitleLockWorkerID, Valid: true} lockedChat.StartedAt = sql.NullTime{Time: time.Now(), Valid: true} + providerID := uuid.New() modelConfig := database.ChatModelConfig{ ID: modelConfigID, - Provider: "openai", Model: "gpt-4o-mini", ContextLimit: 8192, + AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true}, } updatedChat := lockedChat updatedChat.Title = wantTitle @@ -998,14 +1007,20 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts_IdleChatReleasesManualLock(t } db.EXPECT().GetChatModelConfigByID(gomock.Any(), modelConfigID).Return(modelConfig, nil) - providerID := uuid.New() + db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(database.AIProvider{ + ID: providerID, + Type: database.AIProviderTypeOpenai, + Enabled: true, + BaseUrl: serverURL, + }, nil).AnyTimes() db.EXPECT().GetAIProviders(gomock.Any(), gomock.Any()).Return([]database.AIProvider{{ ID: providerID, Type: database.AIProviderTypeOpenai, Enabled: true, BaseUrl: serverURL, - }}, nil) - db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), []uuid.UUID{providerID}).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil) + }}, nil).AnyTimes() + db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes() + db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), gomock.Any()).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes() db.EXPECT().GetChatUsageLimitConfig(gomock.Any()).Return(database.ChatUsageLimitConfig{}, sql.ErrNoRows) db.EXPECT().GetChatMessagesByChatIDAscPaginated( gomock.Any(), @@ -3495,3 +3510,73 @@ func TestServer_inflightContext(t *testing.T) { t.Fatal("inflight context not canceled on server shutdown") } } + +// TestPrepareManualTitleDebugRun_RouteFailureDerivesProviderFromConfig drives +// the fallback branch in prepareManualTitleDebugRun: AI-gateway route +// resolution fails (the BYOK key lookup returns a non-ErrNoRows error) while +// the linked provider stays enabled, so the debug run records the provider +// type derived from modelConfig.AIProviderID instead of an empty string. +func TestPrepareManualTitleDebugRun_RouteFailureDerivesProviderFromConfig(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitShort) + ctrl := gomock.NewController(t) + db := dbmock.NewMockStore(ctrl) + logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) + + ownerID := uuid.New() + providerID := uuid.New() + chat := database.Chat{ID: uuid.New(), OwnerID: ownerID} + modelConfig := database.ChatModelConfig{ + ID: uuid.New(), + Model: "claude-sonnet-4", + AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true}, + } + provider := database.AIProvider{ + ID: providerID, + Type: database.AIProviderTypeAnthropic, + Name: "anthropic", + Enabled: true, + } + + // Resolved twice: once by gatewayProviderForConfig during route resolution, + // once by the fallback's own enabledAIProviderByID lookup. + db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes() + // A non-ErrNoRows BYOK error fails route resolution while the provider stays + // enabled, which is exactly the gap the fallback covers. + db.EXPECT().GetUserAIProviderKeyByProviderID(gomock.Any(), database.GetUserAIProviderKeyByProviderIDParams{ + UserID: ownerID, + AIProviderID: providerID, + }).Return(database.UserAIProviderKey{}, sql.ErrConnDone) + + var gotProvider sql.NullString + db.EXPECT().InsertChatDebugRun(gomock.Any(), gomock.Any()).DoAndReturn( + func(_ context.Context, params database.InsertChatDebugRunParams) (database.ChatDebugRun, error) { + gotProvider = params.Provider + return database.ChatDebugRun{ChatID: params.ChatID, Provider: params.Provider}, nil + }, + ) + + server := &Server{ + db: db, + logger: logger, + aiGatewayRoutingEnabled: true, + allowBYOK: true, + } + debugSvc := chatdebug.NewService(db, logger, nil) + fallbackModel := &chattest.FakeModel{ProviderName: "stub", ModelName: "stub"} + + server.prepareManualTitleDebugRun( + ctx, + debugSvc, + chat, + modelConfig, + chatprovider.ProviderAPIKeys{}, + modelBuildOptions{}, + nil, + fallbackModel, + ) + + require.True(t, gotProvider.Valid, "debug run provider should be populated from the linked config") + require.Equal(t, "anthropic", gotProvider.String) +} diff --git a/coderd/x/chatd/chatd_retry_test.go b/coderd/x/chatd/chatd_retry_test.go index 24457dabd4..969929c13d 100644 --- a/coderd/x/chatd/chatd_retry_test.go +++ b/coderd/x/chatd/chatd_retry_test.go @@ -103,7 +103,6 @@ func TestActiveServer_RetryStreamSilenceTimeoutAndClassification(t *testing.T) { }) user, org, _ := seedChatDependenciesWithProvider(t, db, "openai", openAIURL) model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", Model: "gpt-4o", Enabled: true, CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, diff --git a/coderd/x/chatd/chatd_test.go b/coderd/x/chatd/chatd_test.go index fa29db36c8..862fca3fa5 100644 --- a/coderd/x/chatd/chatd_test.go +++ b/coderd/x/chatd/chatd_test.go @@ -7713,7 +7713,6 @@ func updateChatModelCompressionThreshold(t *testing.T, db database.Store, model ID: model.ID, DisplayName: model.DisplayName, Model: model.Model, - Provider: model.Provider, Enabled: model.Enabled, ContextLimit: model.ContextLimit, CompressionThreshold: model.CompressionThreshold, @@ -7730,7 +7729,6 @@ func updateChatModelContextLimit(t *testing.T, db database.Store, model database ID: model.ID, DisplayName: model.DisplayName, Model: model.Model, - Provider: model.Provider, Enabled: model.Enabled, ContextLimit: model.ContextLimit, CompressionThreshold: model.CompressionThreshold, @@ -7749,7 +7747,6 @@ func updateChatModelCallConfig(t *testing.T, db database.Store, model database.C ID: model.ID, DisplayName: model.DisplayName, Model: model.Model, - Provider: model.Provider, Enabled: model.Enabled, ContextLimit: model.ContextLimit, CompressionThreshold: model.CompressionThreshold, @@ -8333,6 +8330,8 @@ func TestProposeChatTitle_DebugRun(t *testing.T) { if tt.wantTitleGenerationRuns > 0 { require.Equal(t, string(codersdk.ChatDebugRunKindTitleGeneration), runs[0].Kind) require.Equal(t, string(tt.wantDebugStatus), runs[0].Status) + require.True(t, runs[0].Provider.Valid) + require.Equal(t, "openai", runs[0].Provider.String) require.True(t, runs[0].FinishedAt.Valid) require.True(t, runs[0].HistoryTipMessageID.Valid) require.Equal(t, message.ID, runs[0].HistoryTipMessageID.Int64) @@ -8378,14 +8377,14 @@ func seedChatDependenciesWithProvider( UserID: user.ID, OrganizationID: org.ID, }) - dbgen.ChatProvider(t, db, database.ChatProvider{ + providerConfig := dbgen.ChatProvider(t, db, database.ChatProvider{ Provider: provider, DisplayName: provider, BaseUrl: baseURL, }) model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: provider, - IsDefault: true, + AIProviderID: uuid.NullUUID{UUID: providerConfig.ID, Valid: true}, + IsDefault: true, }) return user, org, model } @@ -8423,8 +8422,8 @@ func seedChatDependenciesWithProviderPolicy( }) model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: provider, - IsDefault: true, + AIProviderID: uuid.NullUUID{UUID: providerConfig.ID, Valid: true}, + IsDefault: true, }) return user, org, providerConfig, model @@ -8482,13 +8481,32 @@ func insertChatModelConfigWithCallConfig( options, err := json.Marshal(callConfig) require.NoError(t, err) + // Reuse the newest AI provider of this type (creating a bare one only when + // none exists) so the config links the seeded provider carrying the mock + // base URL and API key rather than a fresh credential-less one. + providers, err := db.GetAIProviders(context.Background(), database.GetAIProvidersParams{IncludeDisabled: true}) + require.NoError(t, err) + var aiProvider database.AIProvider + for _, candidate := range providers { + if candidate.Type != database.AIProviderType(provider) { + continue + } + if aiProvider.ID == uuid.Nil || candidate.CreatedAt.After(aiProvider.CreatedAt) { + aiProvider = candidate + } + } + if aiProvider.ID == uuid.Nil { + aiProvider = dbgen.AIProvider(t, db, database.AIProvider{ + Type: database.AIProviderType(provider), + }) + } return dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: provider, - Model: model, - DisplayName: model, - CreatedBy: uuid.NullUUID{UUID: userID, Valid: true}, - UpdatedBy: uuid.NullUUID{UUID: userID, Valid: true}, - Options: options, + AIProviderID: uuid.NullUUID{UUID: aiProvider.ID, Valid: true}, + Model: model, + DisplayName: model, + CreatedBy: uuid.NullUUID{UUID: userID, Valid: true}, + UpdatedBy: uuid.NullUUID{UUID: userID, Valid: true}, + Options: options, }) } @@ -9671,7 +9689,6 @@ func seedAIGatewayOpenAITestDependencies( BaseUrl: openAIURL, }) model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: string(database.AIProviderTypeOpenai), Model: "gpt-4o-mini", IsDefault: true, AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, @@ -12548,14 +12565,12 @@ func TestProviderSwitchSanitizesAndRestoresPEToolHistory(t *testing.T) { }) mA := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai-compat", Model: "gpt-4o-mini", DisplayName: "Model A", Enabled: true, AIProviderID: uuid.NullUUID{UUID: cpA.ID, Valid: true}, }) mB := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai-compat", Model: "gpt-4o-mini", DisplayName: "Model B", Enabled: true, diff --git a/coderd/x/chatd/chatdebug/service_test.go b/coderd/x/chatd/chatdebug/service_test.go index 358ff0e36b..df39abb300 100644 --- a/coderd/x/chatd/chatdebug/service_test.go +++ b/coderd/x/chatd/chatdebug/service_test.go @@ -25,13 +25,14 @@ import ( ) type testFixture struct { - ctx context.Context - db database.Store - svc *chatdebug.Service - org database.Organization - owner database.User - chat database.Chat - model database.ChatModelConfig + ctx context.Context + db database.Store + svc *chatdebug.Service + org database.Organization + owner database.User + chat database.Chat + model database.ChatModelConfig + provider string } func TestService_IsEnabled(t *testing.T) { @@ -115,7 +116,7 @@ func TestService_CreateRun(t *testing.T) { HistoryTipMessageID: historyTipMsg.ID, Kind: chatdebug.KindChatTurn, Status: chatdebug.StatusInProgress, - Provider: fixture.model.Provider, + Provider: fixture.provider, Model: fixture.model.Model, Summary: map[string]any{ "phase": "create", @@ -126,7 +127,7 @@ func TestService_CreateRun(t *testing.T) { assertRunMatches(t, run, fixture.chat.ID, rootChat.ID, parentChat.ID, fixture.model.ID, triggerMsg.ID, historyTipMsg.ID, chatdebug.KindChatTurn, chatdebug.StatusInProgress, - fixture.model.Provider, fixture.model.Model, + fixture.provider, fixture.model.Model, `{"count":1,"phase":"create"}`) stored, err := fixture.db.GetChatDebugRunByID(fixture.ctx, run.ID) @@ -470,7 +471,7 @@ func TestService_UpdateStep(t *testing.T) { ResponseStatus: 200, DurationMs: 25, }}, - Metadata: map[string]any{"provider": fixture.model.Provider}, + Metadata: map[string]any{"provider": fixture.provider}, FinishedAt: finishedAt, }) require.NoError(t, err) @@ -487,7 +488,7 @@ func TestService_UpdateStep(t *testing.T) { `[{"number":1,"response_status":200,"duration_ms":25}]`, string(updated.Attempts), ) - require.JSONEq(t, `{"provider":"`+fixture.model.Provider+`"}`, + require.JSONEq(t, `{"provider":"`+fixture.provider+`"}`, string(updated.Metadata)) require.True(t, updated.FinishedAt.Valid) storedSteps, err := fixture.db.GetChatDebugStepsByRunID(fixture.ctx, run.ID) @@ -1070,14 +1071,17 @@ func newFixture(t *testing.T) testFixture { ctx := testutil.Context(t, testutil.WaitLong) db, _ := dbtestutil.NewDB(t) org, owner, chat, model := seedChat(t, db) + prov, err := db.GetAIProviderByID(ctx, model.AIProviderID.UUID) + require.NoError(t, err) return testFixture{ - ctx: ctx, - db: db, - svc: chatdebug.NewService(db, testutil.Logger(t), nil), - org: org, - owner: owner, - chat: chat, - model: model, + ctx: ctx, + db: db, + svc: chatdebug.NewService(db, testutil.Logger(t), nil), + org: org, + owner: owner, + chat: chat, + model: model, + provider: string(prov.Type), } } @@ -1158,7 +1162,7 @@ func createRun(t *testing.T, fixture testFixture) database.ChatDebugRun { ModelConfigID: fixture.model.ID, Kind: chatdebug.KindChatTurn, Status: chatdebug.StatusInProgress, - Provider: fixture.model.Provider, + Provider: fixture.provider, Model: fixture.model.Model, }) require.NoError(t, err) diff --git a/coderd/x/chatd/chatprompt/chatprompt_test.go b/coderd/x/chatd/chatprompt/chatprompt_test.go index bba38785ad..fa9a58869b 100644 --- a/coderd/x/chatd/chatprompt/chatprompt_test.go +++ b/coderd/x/chatd/chatprompt/chatprompt_test.go @@ -2516,12 +2516,12 @@ func TestMediaToolResultRoundTrip(t *testing.T) { OrganizationID: org.ID, }) - dbgen.ChatProvider(t, db, database.ChatProvider{ + anthropicProvider := dbgen.ChatProvider(t, db, database.ChatProvider{ Provider: "anthropic", }) model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "anthropic", + AIProviderID: uuid.NullUUID{UUID: anthropicProvider.ID, Valid: true}, Model: "test-model", IsDefault: true, ContextLimit: 200000, diff --git a/coderd/x/chatd/chatstate/family_test.go b/coderd/x/chatd/chatstate/family_test.go index b7781d83b6..7fcbd18310 100644 --- a/coderd/x/chatd/chatstate/family_test.go +++ b/coderd/x/chatd/chatstate/family_test.go @@ -211,7 +211,6 @@ func seedFamilyDeps(t *testing.T, db database.Store) (database.User, database.Or BaseUrl: "http://example.invalid", }) model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", IsDefault: true, }) return user, org, model diff --git a/coderd/x/chatd/chatstate/machine_test.go b/coderd/x/chatd/chatstate/machine_test.go index 65e96b0f8a..c4def78174 100644 --- a/coderd/x/chatd/chatstate/machine_test.go +++ b/coderd/x/chatd/chatstate/machine_test.go @@ -58,7 +58,6 @@ func newTestFixture(t *testing.T) *testFixture { BaseUrl: "http://example.invalid", }) model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", IsDefault: true, }) apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID}) diff --git a/coderd/x/chatd/chatstate/trigger_test.go b/coderd/x/chatd/chatstate/trigger_test.go index dc31651ac2..5d0bfcb04c 100644 --- a/coderd/x/chatd/chatstate/trigger_test.go +++ b/coderd/x/chatd/chatstate/trigger_test.go @@ -38,7 +38,6 @@ func newTriggerFixture(t *testing.T) *triggerFixture { BaseUrl: "http://example.invalid", }) model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", IsDefault: true, }) f := &testFixture{ diff --git a/coderd/x/chatd/configcache_internal_test.go b/coderd/x/chatd/configcache_internal_test.go index f868c321a9..375a307f94 100644 --- a/coderd/x/chatd/configcache_internal_test.go +++ b/coderd/x/chatd/configcache_internal_test.go @@ -683,7 +683,6 @@ func testAIProvider(name string) database.AIProvider { func testChatModelConfig(id uuid.UUID, model string) database.ChatModelConfig { return database.ChatModelConfig{ ID: id, - Provider: "openai", Model: model, DisplayName: model, Enabled: true, diff --git a/coderd/x/chatd/generation_preparer.go b/coderd/x/chatd/generation_preparer.go index 73498b0047..107607a639 100644 --- a/coderd/x/chatd/generation_preparer.go +++ b/coderd/x/chatd/generation_preparer.go @@ -262,7 +262,11 @@ func (server *Server) prepareGeneration( acceptsFilePart := func(mediaType string) bool { return chatprovider.AcceptsFilePartMediaType(model.Provider(), model.Model(), mediaType) } - prompt, err = chatprompt.ConvertMessagesWithFiles(ctx, promptRows, server.chatFileResolver(modelConfig.Provider), logger, acceptsFilePart) + providerType, err := modelRoute.providerHint() + if err != nil { + return xerrors.Errorf("resolve provider type: %w", err) + } + prompt, err = chatprompt.ConvertMessagesWithFiles(ctx, promptRows, server.chatFileResolver(providerType), logger, acceptsFilePart) if err != nil { return xerrors.Errorf("build chat prompt: %w", err) } diff --git a/coderd/x/chatd/generation_preparer_internal_test.go b/coderd/x/chatd/generation_preparer_internal_test.go index c5fe09bd4c..493dcfce35 100644 --- a/coderd/x/chatd/generation_preparer_internal_test.go +++ b/coderd/x/chatd/generation_preparer_internal_test.go @@ -113,7 +113,6 @@ func TestDeriveFinalTurnRunResult(t *testing.T) { CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, }) modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", Model: "gpt-4o-mini", DisplayName: "gpt-4o-mini", Options: json.RawMessage(`{}`), @@ -233,7 +232,6 @@ func TestDeriveFinalTurnRunResult(t *testing.T) { // degraded path that still returns the re-derived text and IDs. provider := insertInternalAIProvider(t, db, database.AIProviderTypeOpenai, "provider-api-key", false) modelCfg := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", Model: "gpt-4o-mini", DisplayName: "gpt-4o-mini", AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, diff --git a/coderd/x/chatd/helpers_test.go b/coderd/x/chatd/helpers_test.go index 352392f26c..bb295728b0 100644 --- a/coderd/x/chatd/helpers_test.go +++ b/coderd/x/chatd/helpers_test.go @@ -133,7 +133,6 @@ func newWorkerTestFixture(t *testing.T) *workerTestFixture { BaseUrl: "http://example.invalid", }) model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", IsDefault: true, }) apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID}) diff --git a/coderd/x/chatd/integration_test.go b/coderd/x/chatd/integration_test.go index 249d72ec3d..9bf30b4e3e 100644 --- a/coderd/x/chatd/integration_test.go +++ b/coderd/x/chatd/integration_test.go @@ -98,7 +98,6 @@ func TestAnthropicWebSearchRoundTrip(t *testing.T) { contextLimit := int64(200000) isDefault := true _, err := expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: string(provider.Type), AIProviderID: &provider.ID, Model: "claude-sonnet-4-20250514", ContextLimit: &contextLimit, @@ -358,7 +357,6 @@ func TestOpenAIReasoningRoundTrip(t *testing.T) { isDefault := true reasoningSummary := "auto" _, err := expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: string(provider.Type), AIProviderID: &provider.ID, Model: "o4-mini", ContextLimit: &contextLimit, @@ -508,7 +506,6 @@ func TestOpenAIReasoningRoundTripStoreFalse(t *testing.T) { isDefault := true reasoningSummary := "auto" _, err := expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: string(provider.Type), AIProviderID: &provider.ID, Model: "o4-mini", ContextLimit: &contextLimit, diff --git a/coderd/x/chatd/model_routing_direct.go b/coderd/x/chatd/model_routing_direct.go index 8173aa75c9..0ba6dacc7a 100644 --- a/coderd/x/chatd/model_routing_direct.go +++ b/coderd/x/chatd/model_routing_direct.go @@ -2,6 +2,7 @@ package chatd import ( "context" + "database/sql" "net/http" "charm.land/fantasy" @@ -83,7 +84,7 @@ func (p *Server) directProviderHintAndProviderForConfig( modelConfig database.ChatModelConfig, ) (string, *database.AIProvider, error) { if !modelConfig.AIProviderID.Valid { - return modelConfig.Provider, nil, nil + return "", nil, sql.ErrNoRows } provider, err := p.enabledAIProviderByID(ctx, modelConfig.AIProviderID.UUID) if err != nil { diff --git a/coderd/x/chatd/model_routing_internal_test.go b/coderd/x/chatd/model_routing_internal_test.go index 52929003fb..2843bfebe0 100644 --- a/coderd/x/chatd/model_routing_internal_test.go +++ b/coderd/x/chatd/model_routing_internal_test.go @@ -125,7 +125,6 @@ func TestResolveModelRouteForConfigPreservesBaseURL(t *testing.T) { server := &Server{db: db} route, err := server.resolveModelRouteForConfig(ctx, ownerID, database.ChatModelConfig{ - Provider: "openai", AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true}, }, chatprovider.ProviderAPIKeys{}) require.NoError(t, err) @@ -238,7 +237,6 @@ func TestResolveModelRouteForConfigAIGatewayProviderAuth(t *testing.T) { modelConfig := database.ChatModelConfig{ ID: uuid.New(), Model: "gpt-4", - Provider: "openai", AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true}, } diff --git a/coderd/x/chatd/quickgen.go b/coderd/x/chatd/quickgen.go index 96fad11f3f..f31a35f70f 100644 --- a/coderd/x/chatd/quickgen.go +++ b/coderd/x/chatd/quickgen.go @@ -102,17 +102,17 @@ func (p *Server) preferredShortTextCandidates( } func selectPreferredConfiguredShortTextModelConfig( - configs []database.ChatModelConfig, + configs []database.GetEnabledChatModelConfigsRow, ) (database.ChatModelConfig, bool) { for _, preferred := range preferredTitleModels { for _, config := range configs { if chatprovider.NormalizeProvider(config.Provider) != preferred.provider { continue } - if !strings.EqualFold(strings.TrimSpace(config.Model), preferred.model) { + if !strings.EqualFold(strings.TrimSpace(config.ChatModelConfig.Model), preferred.model) { continue } - return config, true + return config.ChatModelConfig, true } } return database.ChatModelConfig{}, false @@ -181,11 +181,18 @@ func (p *Server) GenerateChatTitleAsync(ctx context.Context, chat database.Chat) ) return } + providerType, err := route.providerHint() + if err != nil { + logger.Debug(titleCtx, "failed to resolve provider type for automatic title generation", + slog.Error(err), + ) + return + } p.maybeGenerateChatTitle( turnCtx, chat, messages, - modelConfig.Provider, + providerType, modelConfig.Model, model, route, @@ -259,8 +266,16 @@ func (p *Server) maybeGenerateChatTitle( var candidates []shortTextCandidate if overrideSet { + overrideProvider, err := overrideRoute.providerHint() + if err != nil { + logger.Debug(ctx, "failed to resolve provider type for title generation override", + slog.F("chat_id", chat.ID), + slog.Error(err), + ) + return + } candidates = []shortTextCandidate{{ - provider: overrideConfig.Provider, + provider: overrideProvider, model: overrideConfig.Model, route: overrideRoute, lm: overrideModel, diff --git a/coderd/x/chatd/quickgen_internal_test.go b/coderd/x/chatd/quickgen_internal_test.go index 0e46ccc0f7..d725a382ea 100644 --- a/coderd/x/chatd/quickgen_internal_test.go +++ b/coderd/x/chatd/quickgen_internal_test.go @@ -391,8 +391,7 @@ func TestMaybeGenerateChatTitlePreservesUpdatedAt(t *testing.T) { CentralApiKeyEnabled: true, }) modelConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", - Model: "test-model", + Model: "test-model", }) userPrompt := "summarize failed workspace build logs" @@ -586,24 +585,23 @@ func Test_selectPreferredConfiguredShortTextModelConfig(t *testing.T) { t.Run("chooses the highest-priority configured lightweight model", func(t *testing.T) { t.Parallel() - configs := []database.ChatModelConfig{ - {Provider: preferredTitleModels[2].provider, Model: preferredTitleModels[2].model}, - {Provider: preferredTitleModels[1].provider, Model: preferredTitleModels[1].model}, - {Provider: "openai", Model: "gpt-4.1"}, + configs := []database.GetEnabledChatModelConfigsRow{ + {ChatModelConfig: database.ChatModelConfig{Model: preferredTitleModels[2].model}, Provider: preferredTitleModels[2].provider}, + {ChatModelConfig: database.ChatModelConfig{Model: preferredTitleModels[1].model}, Provider: preferredTitleModels[1].provider}, + {ChatModelConfig: database.ChatModelConfig{Model: "gpt-4.1"}, Provider: "openai"}, } got, ok := selectPreferredConfiguredShortTextModelConfig(configs) require.True(t, ok) - require.Equal(t, preferredTitleModels[1].provider, got.Provider) require.Equal(t, preferredTitleModels[1].model, got.Model) }) t.Run("returns false when no preferred lightweight model is configured", func(t *testing.T) { t.Parallel() - got, ok := selectPreferredConfiguredShortTextModelConfig([]database.ChatModelConfig{{ - Provider: "openai", - Model: "gpt-4.1", + got, ok := selectPreferredConfiguredShortTextModelConfig([]database.GetEnabledChatModelConfigsRow{{ + ChatModelConfig: database.ChatModelConfig{Model: "gpt-4.1"}, + Provider: "openai", }}) require.False(t, ok) require.Equal(t, database.ChatModelConfig{}, got) diff --git a/coderd/x/chatd/subagent.go b/coderd/x/chatd/subagent.go index 5056c8e9a8..b5e9888793 100644 --- a/coderd/x/chatd/subagent.go +++ b/coderd/x/chatd/subagent.go @@ -161,39 +161,6 @@ func personalModelOverrideContextForSubagent( } } -func validateModelConfigAndResolveProvider( - modelConfig database.ChatModelConfig, -) (database.ChatModelConfig, string, error) { - if !modelConfig.Enabled { - return database.ChatModelConfig{}, "", sql.ErrNoRows - } - providerName, _, err := chatprovider.ResolveModelWithProviderHint( - modelConfig.Model, - modelConfig.Provider, - ) - if err != nil { - return database.ChatModelConfig{}, "", xerrors.Errorf( - "%w: %v", - errInvalidModelOverrideMetadata, - err, - ) - } - return modelConfig, providerName, nil -} - -func enabledProviderContainsName( - providers []database.AIProvider, - providerName string, -) bool { - normalizedProviderName := chatprovider.NormalizeProvider(providerName) - for _, provider := range providers { - if chatprovider.NormalizeProvider(string(provider.Type)) == normalizedProviderName { - return true - } - } - return false -} - func userCanUseProviderKeys( providerKeys chatprovider.ProviderAPIKeys, providerName string, @@ -542,18 +509,8 @@ func (p *Server) resolveModelConfigAndNormalizedProvider( } return modelConfig, providerName, nil } - modelConfig, providerName, err := validateModelConfigAndResolveProvider(modelConfig) - if err != nil { - return database.ChatModelConfig{}, "", err - } - enabledProviders, err := p.configCache.EnabledProviders(ctx) - if err != nil { - return database.ChatModelConfig{}, "", err - } - if !enabledProviderContainsName(enabledProviders, providerName) { - return database.ChatModelConfig{}, "", sql.ErrNoRows - } - return modelConfig, providerName, nil + // Active configs carry a provider FK; resolved above. Missing FK means no usable config. + return database.ChatModelConfig{}, "", sql.ErrNoRows } func (p *Server) subagentTools( diff --git a/coderd/x/chatd/subagent_internal_test.go b/coderd/x/chatd/subagent_internal_test.go index 6db1692aa2..f5ad4f0174 100644 --- a/coderd/x/chatd/subagent_internal_test.go +++ b/coderd/x/chatd/subagent_internal_test.go @@ -213,7 +213,6 @@ func seedInternalChatDeps( }) model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, IsDefault: true, }) @@ -580,8 +579,7 @@ func TestResolveChatModel_AIProviderDisabled(t *testing.T) { user, org, _ := seedInternalChatDeps(t, db) provider := insertInternalAIProvider(t, db, database.AIProviderTypeOpenai, "provider-api-key", false) modelConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", - Model: "gpt-4o-mini", + Model: "gpt-4o-mini", AIProviderID: uuid.NullUUID{ UUID: provider.ID, Valid: true, @@ -714,11 +712,30 @@ func insertInternalChatModelConfigWithOptions( ) database.ChatModelConfig { t.Helper() + // Reuse the newest AI provider of this type (creating a bare credential-less + // one only when none exists) so the config links the provider already + // carrying the test's credentials, or lack thereof, rather than a fresh one. + providers, err := db.GetAIProviders(context.Background(), database.GetAIProvidersParams{IncludeDisabled: true}) + require.NoError(t, err) + var aiProvider database.AIProvider + for _, candidate := range providers { + if candidate.Type != database.AIProviderType(provider) { + continue + } + if aiProvider.ID == uuid.Nil || candidate.CreatedAt.After(aiProvider.CreatedAt) { + aiProvider = candidate + } + } + if aiProvider.ID == uuid.Nil { + aiProvider = dbgen.AIProvider(t, db, database.AIProvider{ + Type: database.AIProviderType(provider), + }) + } modelConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: provider, - Model: model, - DisplayName: model, - Options: options, + AIProviderID: uuid.NullUUID{UUID: aiProvider.ID, Valid: true}, + Model: model, + DisplayName: model, + Options: options, }, func(p *database.InsertChatModelConfigParams) { p.Enabled = enabled }) @@ -1448,7 +1465,6 @@ func TestResolveConfiguredModelOverride_AcceptsAmbientCredentialsProvider( ownerID := uuid.New() modelConfig := database.ChatModelConfig{ ID: uuid.New(), - Provider: "bedrock", Model: "anthropic.claude-haiku-4-5-20251001-v1:0", DisplayName: "Ambient Bedrock Override", Enabled: true, @@ -2115,7 +2131,7 @@ func TestSpawnAgent_ExploreFallsBackWhenOverrideCredentialsAreUnavailable(t *tes currentTurnModel := insertInternalChatModelConfig( t, db, "explore-missing-user-key-current-"+uuid.NewString(), true, ) - dbgen.ChatProvider(t, db, database.ChatProvider{ + overrideProvider := dbgen.ChatProvider(t, db, database.ChatProvider{ Provider: "openai-compat", DisplayName: "OpenAI Compat", }, func(p *database.InsertChatProviderParams) { @@ -2126,9 +2142,9 @@ func TestSpawnAgent_ExploreFallsBackWhenOverrideCredentialsAreUnavailable(t *tes }) overrideModel := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai-compat", - Model: "gpt-4o-mini", - DisplayName: "Explore Override Missing User Key", + AIProviderID: uuid.NullUUID{UUID: overrideProvider.ID, Valid: true}, + Model: "gpt-4o-mini", + DisplayName: "Explore Override Missing User Key", }) require.NoError(t, db.UpsertChatExploreModelOverride(ctx, overrideModel.ID.String())) parentChat := createInternalParentChat( @@ -2764,7 +2780,9 @@ func TestSpawnAgent_ComputerUseUsesComputerUseModelNotParent(t *testing.T) { insertEnabledAnthropicProvider(t, db, user.ID) workspace, build, agent := seedWorkspaceBinding(t, db, user.ID) - require.Equal(t, "openai", model.Provider, "seed helper must create an OpenAI model") + seedProvider, err := db.GetAIProviderByID(ctx, model.AIProviderID.UUID) + require.NoError(t, err) + require.Equal(t, "openai", string(seedProvider.Type), "seed helper must create an OpenAI model") parent, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, @@ -2806,7 +2824,7 @@ func TestSpawnAgent_ComputerUseUsesComputerUseModelNotParent(t *testing.T) { assert.Equal(t, database.ChatModeComputerUse, childChat.Mode.ChatMode) computerUseModelProvider, computerUseModelName, ok := chattool.DefaultComputerUseModel(chattool.ComputerUseProviderAnthropic) require.True(t, ok) - assert.NotEqual(t, model.Provider, computerUseModelProvider, + assert.NotEqual(t, string(seedProvider.Type), computerUseModelProvider, "computer use model provider must differ from parent model provider") assert.Equal(t, "anthropic", computerUseModelProvider) assert.NotEmpty(t, computerUseModelName) @@ -3668,7 +3686,6 @@ func TestAwaitSubagentCompletion(t *testing.T) { BaseUrl: providerServer.URL, }) model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ - Provider: "openai", Model: "gpt-4o-mini", AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, }) diff --git a/coderd/x/chatd/tasks_test.go b/coderd/x/chatd/tasks_test.go index b2115dfe9a..4f549554e2 100644 --- a/coderd/x/chatd/tasks_test.go +++ b/coderd/x/chatd/tasks_test.go @@ -873,7 +873,7 @@ func newTaskTestFixture(t *testing.T) *taskTestFixture { DisplayName: "openai", BaseUrl: "http://example.invalid", }) - model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{Provider: "openai", IsDefault: true}) + model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{IsDefault: true}) apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID}) return &taskTestFixture{db: db, pubsub: newTaskRecordingPubsub(ps), sqlDB: sqlDB, user: user, org: org, model: model, apiKey: apiKey} } diff --git a/coderd/x/chatd/title_override_internal_test.go b/coderd/x/chatd/title_override_internal_test.go index 4c2d49e9c4..b93baba649 100644 --- a/coderd/x/chatd/title_override_internal_test.go +++ b/coderd/x/chatd/title_override_internal_test.go @@ -348,6 +348,8 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideCallFailureSkipsFallback( logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) chat, messages := titleOverrideTestChatAndMessages(t) overrideConfig := titleOverrideModelConfig("gpt-4.1", true) + providerID := uuid.New() + overrideConfig.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true} var requestCount atomic.Int32 serverURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse { @@ -355,6 +357,12 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideCallFailureSkipsFallback( require.Equal(t, overrideConfig.Model, req.Model) return chattest.OpenAINonStreamingResponse(`{"title":""}`) }) + provider := database.AIProvider{ + ID: providerID, + Type: database.AIProviderTypeOpenai, + Enabled: true, + BaseUrl: serverURL, + } keys := titleOverrideOpenAIKeys(serverURL) fallbackModel := &chattest.FakeModel{ GenerateObjectFn: func(context.Context, fantasy.ObjectCall) (*fantasy.ObjectResponse, error) { @@ -365,8 +373,11 @@ func TestMaybeGenerateChatTitle_TitleGenerationOverrideCallFailureSkipsFallback( db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil) db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil) - db.EXPECT().GetAIProviders(gomock.Any(), gomock.Any()).Return([]database.AIProvider{{Type: database.AIProviderTypeOpenai, Enabled: true}}, nil) - db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), []uuid.UUID{uuid.Nil}).Return(nil, nil) + db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes() + db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ + ProviderID: providerID, + APIKey: "test-key", + }}, nil).AnyTimes() generated := &generatedChatTitle{} server := titleOverrideTestServer(db, logger) @@ -398,18 +409,34 @@ func TestResolveManualTitleModel_TitleGenerationOverrideUnset(t *testing.T) { db := dbmock.NewMockStore(ctrl) logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) chat, _ := titleOverrideTestChatAndMessages(t) + providerID := uuid.New() preferredConfig := database.ChatModelConfig{ - ID: uuid.New(), - Provider: preferredTitleModels[1].provider, - Model: preferredTitleModels[1].model, - Enabled: true, + ID: uuid.New(), + AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true}, + Model: preferredTitleModels[1].model, + Enabled: true, + } + serverURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse { + t.Fatal("model construction should not call the provider") + return chattest.OpenAIResponse{} + }) + provider := database.AIProvider{ + ID: providerID, + Type: database.AIProviderTypeOpenai, + Enabled: true, + BaseUrl: serverURL, } db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return("", nil) - db.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return([]database.ChatModelConfig{ - {Provider: "openai", Model: "gpt-4.1", Enabled: true}, - preferredConfig, + db.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return([]database.GetEnabledChatModelConfigsRow{ + {ChatModelConfig: database.ChatModelConfig{Model: "gpt-4.1", Enabled: true}, Provider: "openai"}, + {ChatModelConfig: preferredConfig, Provider: preferredTitleModels[1].provider}, }, nil) + db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes() + db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ + ProviderID: providerID, + APIKey: "test-key", + }}, nil).AnyTimes() server := titleOverrideTestServer(db, logger) model, gotConfig, _, err := server.resolveManualTitleModel( @@ -435,7 +462,6 @@ func TestResolveManualTitleModel_TitleGenerationOverrideUnsetAIProvider(t *testi providerID := uuid.New() preferredConfig := database.ChatModelConfig{ ID: uuid.New(), - Provider: preferredTitleModels[1].provider, AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true}, Model: preferredTitleModels[1].model, Enabled: true, @@ -452,8 +478,8 @@ func TestResolveManualTitleModel_TitleGenerationOverrideUnsetAIProvider(t *testi } db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return("", nil) - db.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return([]database.ChatModelConfig{ - preferredConfig, + db.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return([]database.GetEnabledChatModelConfigsRow{ + {ChatModelConfig: preferredConfig, Provider: preferredTitleModels[1].provider}, }, nil) db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil) db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ @@ -483,18 +509,34 @@ func TestResolveManualTitleModel_TitleGenerationOverrideReadDBError(t *testing.T db := dbmock.NewMockStore(ctrl) logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) chat, _ := titleOverrideTestChatAndMessages(t) + providerID := uuid.New() preferredConfig := database.ChatModelConfig{ - ID: uuid.New(), - Provider: preferredTitleModels[1].provider, - Model: preferredTitleModels[1].model, - Enabled: true, + ID: uuid.New(), + AIProviderID: uuid.NullUUID{UUID: providerID, Valid: true}, + Model: preferredTitleModels[1].model, + Enabled: true, + } + serverURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse { + t.Fatal("model construction should not call the provider") + return chattest.OpenAIResponse{} + }) + provider := database.AIProvider{ + ID: providerID, + Type: database.AIProviderTypeOpenai, + Enabled: true, + BaseUrl: serverURL, } db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return("", sql.ErrConnDone) - db.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return([]database.ChatModelConfig{ - {Provider: "openai", Model: "gpt-4.1", Enabled: true}, - preferredConfig, + db.EXPECT().GetEnabledChatModelConfigs(gomock.Any()).Return([]database.GetEnabledChatModelConfigsRow{ + {ChatModelConfig: database.ChatModelConfig{Model: "gpt-4.1", Enabled: true}, Provider: "openai"}, + {ChatModelConfig: preferredConfig, Provider: preferredTitleModels[1].provider}, }, nil) + db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes() + db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ + ProviderID: providerID, + APIKey: "test-key", + }}, nil).AnyTimes() server := titleOverrideTestServer(db, logger) model, gotConfig, _, err := server.resolveManualTitleModel( @@ -518,11 +560,26 @@ func TestResolveManualTitleModel_TitleGenerationOverrideSetUsable(t *testing.T) logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) chat, _ := titleOverrideTestChatAndMessages(t) overrideConfig := titleOverrideModelConfig("gpt-4.1", true) + providerID := uuid.New() + overrideConfig.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true} + serverURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse { + t.Fatal("model construction should not call the provider") + return chattest.OpenAIResponse{} + }) + provider := database.AIProvider{ + ID: providerID, + Type: database.AIProviderTypeOpenai, + Enabled: true, + BaseUrl: serverURL, + } db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil) db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil) - db.EXPECT().GetAIProviders(gomock.Any(), gomock.Any()).Return([]database.AIProvider{{Type: database.AIProviderTypeOpenai, Enabled: true}}, nil) - db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), gomock.Any()).Return(nil, nil).AnyTimes() + db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes() + db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ + ProviderID: providerID, + APIKey: "test-key", + }}, nil).AnyTimes() server := titleOverrideTestServer(db, logger) model, gotConfig, _, err := server.resolveManualTitleModel( @@ -546,11 +603,18 @@ func TestResolveManualTitleModel_TitleGenerationOverrideMissingCredentials(t *te logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) chat, _ := titleOverrideTestChatAndMessages(t) overrideConfig := titleOverrideModelConfig("gpt-4.1", true) + providerID := uuid.New() + overrideConfig.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true} + provider := database.AIProvider{ + ID: providerID, + Type: database.AIProviderTypeOpenai, + Enabled: true, + } db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil) db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil) - db.EXPECT().GetAIProviders(gomock.Any(), gomock.Any()).Return([]database.AIProvider{{Type: database.AIProviderTypeOpenai, Enabled: true}}, nil) - db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), gomock.Any()).Return(nil, nil).AnyTimes() + db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes() + db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return(nil, nil).AnyTimes() server := titleOverrideTestServer(db, logger) model, gotConfig, _, err := server.resolveManualTitleModel( @@ -730,10 +794,9 @@ func titleOverrideTestServer(db database.Store, logger slog.Logger) *Server { func titleOverrideModelConfig(model string, enabled bool) database.ChatModelConfig { return database.ChatModelConfig{ - ID: uuid.New(), - Provider: "openai", - Model: model, - Enabled: enabled, + ID: uuid.New(), + Model: model, + Enabled: enabled, } } diff --git a/coderd/x/chatd/turn_summary_internal_test.go b/coderd/x/chatd/turn_summary_internal_test.go index 9547f9a759..d5d545b2ed 100644 --- a/coderd/x/chatd/turn_summary_internal_test.go +++ b/coderd/x/chatd/turn_summary_internal_test.go @@ -43,7 +43,6 @@ func TestUpdateLastTurnSummaryRejectsStaleWrites(t *testing.T) { modelCfg, err := db.InsertChatModelConfig(ctx, database.InsertChatModelConfigParams{ AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, - Provider: "openai", Model: "test-model", DisplayName: "Test Model", CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true}, @@ -138,7 +137,6 @@ func TestPendingChatPersistsSummaryButSkipsWebPush(t *testing.T) { modelCfg, err := db.InsertChatModelConfig(ctx, database.InsertChatModelConfigParams{ AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, - Provider: "openai", Model: "test-model", DisplayName: "Test Model", CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true}, @@ -222,7 +220,6 @@ func TestSuccessfulChildChatOutcomeSkipsSummaryAndWebPush(t *testing.T) { modelCfg, err := db.InsertChatModelConfig(ctx, database.InsertChatModelConfigParams{ AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, - Provider: "openai", Model: "test-model", DisplayName: "Test Model", CreatedBy: uuid.NullUUID{UUID: owner.ID, Valid: true}, diff --git a/codersdk/chats.go b/codersdk/chats.go index 649c681987..b374c5f244 100644 --- a/codersdk/chats.go +++ b/codersdk/chats.go @@ -1249,8 +1249,7 @@ type CreateUserChatProviderKeyRequest struct { // ChatModelConfig is an admin-managed model configuration. type ChatModelConfig struct { ID uuid.UUID `json:"id" format:"uuid"` - Provider string `json:"provider"` - AIProviderID *uuid.UUID `json:"ai_provider_id,omitempty" format:"uuid"` + AIProviderID uuid.UUID `json:"ai_provider_id" format:"uuid"` Model string `json:"model"` DisplayName string `json:"display_name"` Enabled bool `json:"enabled"` @@ -1461,7 +1460,6 @@ func (c *ChatModelCallConfig) UnmarshalJSON(data []byte) error { // CreateChatModelConfigRequest creates a chat model config. type CreateChatModelConfigRequest struct { - Provider string `json:"provider,omitempty"` AIProviderID *uuid.UUID `json:"ai_provider_id,omitempty" format:"uuid"` Model string `json:"model"` DisplayName string `json:"display_name,omitempty"` @@ -1474,7 +1472,6 @@ type CreateChatModelConfigRequest struct { // UpdateChatModelConfigRequest updates a chat model config. type UpdateChatModelConfigRequest struct { - Provider string `json:"provider,omitempty"` AIProviderID *uuid.UUID `json:"ai_provider_id,omitempty" format:"uuid"` Model string `json:"model,omitempty"` DisplayName string `json:"display_name,omitempty"` diff --git a/enterprise/coderd/exp_chats_test.go b/enterprise/coderd/exp_chats_test.go index cf6f958f79..364b82fc68 100644 --- a/enterprise/coderd/exp_chats_test.go +++ b/enterprise/coderd/exp_chats_test.go @@ -55,7 +55,6 @@ func createOpenAIModelConfigForTest( t.Helper() provider := createOpenAIProviderForTest(ctx, t, client, apiKey, baseURL) model, err := client.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: string(provider.Type), AIProviderID: &provider.ID, Model: "gpt-4", DisplayName: "GPT-4", @@ -954,7 +953,6 @@ func TestChatModelConfigDefault(t *testing.T) { firstModel, err := expClient.CreateChatModelConfig( ctx, codersdk.CreateChatModelConfigRequest{ - Provider: string(provider.Type), AIProviderID: &provider.ID, Model: "gpt-5-a", DisplayName: "GPT 5 A", @@ -969,7 +967,6 @@ func TestChatModelConfigDefault(t *testing.T) { secondModel, err := expClient.CreateChatModelConfig( ctx, codersdk.CreateChatModelConfigRequest{ - Provider: string(provider.Type), AIProviderID: &provider.ID, Model: "gpt-5-b", DisplayName: "GPT 5 B", @@ -1099,7 +1096,6 @@ func TestCreateChatNonDefaultOrg(t *testing.T) { provider := createOpenAIProviderForTest(ctx, t, expClient, "test-key", "https://example.com") _, err := expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: string(provider.Type), AIProviderID: &provider.ID, Model: "gpt-4o-mini", DisplayName: "Test Model", @@ -1169,7 +1165,6 @@ func TestListChats_OrgAdminOnlySeesOwnChats(t *testing.T) { provider := createOpenAIProviderForTest(ctx, t, expClient, "test-key", "https://example.com") _, err := expClient.CreateChatModelConfig(ctx, codersdk.CreateChatModelConfigRequest{ - Provider: string(provider.Type), AIProviderID: &provider.ID, Model: "gpt-4o-mini", DisplayName: "Test Model", diff --git a/scaletest/chat/provider.go b/scaletest/chat/provider.go index 156339d8bb..79baf1d4a2 100644 --- a/scaletest/chat/provider.go +++ b/scaletest/chat/provider.go @@ -99,7 +99,7 @@ func ensureScaletestChatModelConfig(ctx context.Context, client chatModelConfigC } for i := range modelConfigs { - matchesProvider := modelConfigs[i].AIProviderID != nil && *modelConfigs[i].AIProviderID == provider.ID + matchesProvider := modelConfigs[i].AIProviderID == provider.ID matchesModel := modelConfigs[i].Model == scaletestModelName if !matchesProvider || !matchesModel { continue diff --git a/site/src/api/typesGenerated.ts b/site/src/api/typesGenerated.ts index d1e9aa8f49..e0de4e1c9f 100644 --- a/site/src/api/typesGenerated.ts +++ b/site/src/api/typesGenerated.ts @@ -2485,8 +2485,7 @@ export interface ChatModelCallConfig { */ export interface ChatModelConfig { readonly id: string; - readonly provider: string; - readonly ai_provider_id?: string; + readonly ai_provider_id: string; readonly model: string; readonly display_name: string; readonly enabled: boolean; @@ -3517,7 +3516,6 @@ export interface CreateChatMessageResponse { * CreateChatModelConfigRequest creates a chat model config. */ export interface CreateChatModelConfigRequest { - readonly provider?: string; readonly ai_provider_id?: string; readonly model: string; readonly display_name?: string; @@ -8900,7 +8898,6 @@ export interface UpdateChatDebugRetentionDaysRequest { * UpdateChatModelConfigRequest updates a chat model config. */ export interface UpdateChatModelConfigRequest { - readonly provider?: string; readonly ai_provider_id?: string; readonly model?: string; readonly display_name?: string; diff --git a/site/src/modules/aiModels/providerStates.test.ts b/site/src/modules/aiModels/providerStates.test.ts index 8b24c5480a..60d165cbc9 100644 --- a/site/src/modules/aiModels/providerStates.test.ts +++ b/site/src/modules/aiModels/providerStates.test.ts @@ -9,7 +9,6 @@ import { canManageProviderModels, deriveProviderStates, type ProviderState, - resolveModelProviderKey, } from "./providerStates"; const baseProviderState: ProviderState = { @@ -28,7 +27,7 @@ const baseProviderState: ProviderState = { }; describe("deriveProviderStates", () => { - it("orders provider configs first, then catalog-only, then model-only providers", () => { + it("orders provider configs first, then catalog-only providers", () => { const providerConfigs = [ { ...MockChatProviderConfig, @@ -44,23 +43,14 @@ describe("deriveProviderStates", () => { ], unsupported_providers: [], }; - const modelConfigs = [ - { ...MockChatModelConfig, id: "m-vercel", provider: "vercel" }, - ]; - const states = deriveProviderStates(modelConfigs, providerConfigs, catalog); + const states = deriveProviderStates([], providerConfigs, catalog); - expect(states.map((s) => s.provider)).toEqual([ - "anthropic", - "google", - "vercel", - ]); + expect(states.map((s) => s.provider)).toEqual(["anthropic", "google"]); expect(states[0].key).toBe("prov-anthropic"); expect(states[1].key).toBe("google"); - expect(states[2].key).toBe("vercel"); expect(states[0].hasEffectiveAPIKey).toBe(true); expect(states[1].hasEffectiveAPIKey).toBe(true); - expect(states[2].hasEffectiveAPIKey).toBe(false); }); it("matches model configs to provider configs by ai_provider_id", () => { @@ -71,10 +61,9 @@ describe("deriveProviderStates", () => { { ...MockChatModelConfig, id: "m1", - provider: "openai", ai_provider_id: "prov-openai", }, - { ...MockChatModelConfig, id: "m2", provider: "openai" }, + { ...MockChatModelConfig, id: "m2", ai_provider_id: "prov-openai" }, ]; const states = deriveProviderStates(modelConfigs, providerConfigs, null); @@ -100,13 +89,13 @@ describe("deriveProviderStates", () => { expect(states[0].hasEffectiveAPIKey).toBe(true); }); - it("drops models without ai_provider_id when multiple configs exist for the same provider", () => { + it("drops models without ai_provider_id", () => { const providerConfigs = [ { ...MockChatProviderConfig, id: "prov-a", provider: "openai" }, { ...MockChatProviderConfig, id: "prov-b", provider: "openai" }, ]; const modelConfigs = [ - { ...MockChatModelConfig, id: "m1", provider: "openai" }, + { ...MockChatModelConfig, id: "m1", ai_provider_id: "" }, ]; const states = deriveProviderStates(modelConfigs, providerConfigs, null); @@ -240,47 +229,3 @@ describe("canManageProviderModels", () => { expect(canManageProviderModels(undefined)).toBe(false); }); }); - -describe("resolveModelProviderKey", () => { - const states: ProviderState[] = [ - { ...baseProviderState, key: "prov-a", provider: "openai" }, - { ...baseProviderState, key: "prov-b", provider: "openai" }, - { ...baseProviderState, key: "prov-anthropic", provider: "anthropic" }, - ]; - - it("prefers ai_provider_id when present", () => { - expect( - resolveModelProviderKey( - { ...MockChatModelConfig, ai_provider_id: "prov-explicit" }, - states, - ), - ).toBe("prov-explicit"); - }); - - it("falls back to the single matching provider state key", () => { - expect( - resolveModelProviderKey( - { ...MockChatModelConfig, provider: "anthropic" }, - states, - ), - ).toBe("prov-anthropic"); - }); - - it("returns an empty key when multiple states match the provider", () => { - expect( - resolveModelProviderKey( - { ...MockChatModelConfig, provider: "openai" }, - states, - ), - ).toBe(""); - }); - - it("falls back to the provider name when no states match", () => { - expect( - resolveModelProviderKey( - { ...MockChatModelConfig, provider: "google" }, - states, - ), - ).toBe("google"); - }); -}); diff --git a/site/src/modules/aiModels/providerStates.ts b/site/src/modules/aiModels/providerStates.ts index 79d8732142..a94fc6b9f3 100644 --- a/site/src/modules/aiModels/providerStates.ts +++ b/site/src/modules/aiModels/providerStates.ts @@ -107,31 +107,20 @@ export const deriveProviderStates = ( catalogProvidersByProvider.set(provider, cp); } - const providerConfigKeysByProvider = new Map(); const providerTypesWithConfigs = new Set(); + const providerConfigsByKey = new Map(); for (const pc of providerConfigs ?? []) { const provider = normalizeProvider(pc.provider); if (!provider) continue; const key = providerConfigStateKey(pc); providerTypesWithConfigs.add(provider); - providerConfigKeysByProvider.set(provider, [ - ...(providerConfigKeysByProvider.get(provider) ?? []), - key, - ]); + if (key) { + providerConfigsByKey.set(key, pc); + } includeEntry(key, provider); } - const modelStateKey = (modelConfig: TypesGen.ChatModelConfig): string => { - const aiProviderID = readOptionalString(modelConfig.ai_provider_id); - if (aiProviderID) { - return aiProviderID; - } - const provider = normalizeProvider(modelConfig.provider); - const providerConfigKeys = providerConfigKeysByProvider.get(provider) ?? []; - if (providerConfigKeys.length === 1) { - return providerConfigKeys[0]; - } - return providerConfigKeys.length === 0 ? provider : ""; - }; + const modelStateKey = (modelConfig: TypesGen.ChatModelConfig): string => + readOptionalString(modelConfig.ai_provider_id) ?? ""; for (const cp of catalogProviders) { const provider = normalizeProvider(cp.provider); @@ -139,14 +128,8 @@ export const deriveProviderStates = ( includeEntry(provider, provider); } for (const mc of modelConfigs) { - includeEntry(modelStateKey(mc), mc.provider); - } - - const providerConfigsByKey = new Map(); - for (const pc of providerConfigs ?? []) { - const key = providerConfigStateKey(pc); - if (!key) continue; - providerConfigsByKey.set(key, pc); + const key = modelStateKey(mc); + includeEntry(key, providerConfigsByKey.get(key)?.provider ?? ""); } const modelConfigsByKey = new Map(); @@ -216,22 +199,3 @@ export const canManageProviderModels = ( providerState.providerConfig.allow_user_api_key), ); }; - -export const resolveModelProviderKey = ( - modelConfig: TypesGen.ChatModelConfig, - providerStates: readonly ProviderState[], -): string => { - const providerID = readOptionalString(modelConfig.ai_provider_id); - if (providerID) { - return providerID; - } - const provider = normalizeProvider(modelConfig.provider); - const matches = providerStates.filter((s) => s.provider === provider); - if (matches.length === 1) { - return matches[0].key; - } - if (matches.length > 1) { - return ""; - } - return provider; -}; diff --git a/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPage.tsx b/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPage.tsx index 03d4af953b..16b10d3e4d 100644 --- a/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPage.tsx +++ b/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPage.tsx @@ -11,6 +11,7 @@ import { chatComputerUseProvider, chatModelConfigs, chatPersonalModelOverridesAdminSettings, + chatProviderConfigs, updateChatAdvisorConfig, updateChatComputerUseProvider, updateChatPersonalModelOverridesAdminSettings, @@ -19,6 +20,7 @@ import type * as TypesGen from "#/api/typesGenerated"; import { useAuthenticated } from "#/hooks/useAuthenticated"; import { useDashboard } from "#/modules/dashboard/useDashboard"; import { RequirePermission } from "#/modules/permissions/RequirePermission"; +import { providerTypeByIDFromConfigs } from "#/pages/AgentsPage/utils/modelOptions"; import { pageTitle } from "#/utils/page"; import { CoderAgentsPageView } from "./CoderAgentsPageView"; @@ -86,6 +88,10 @@ const CoderAgentsPage: FC = () => { ...chatComputerUseProvider(), enabled: canEditDeploymentConfig && showVirtualDesktopSettings, }); + const providerConfigsQuery = useQuery({ + ...chatProviderConfigs(), + enabled: canEditDeploymentConfig, + }); const savePersonalModelOverridesAdminSettingsMutation = useMutation( updateChatPersonalModelOverridesAdminSettings(queryClient), ); @@ -108,6 +114,10 @@ const CoderAgentsPage: FC = () => { updateChatComputerUseProvider(queryClient), ); + const providerTypeByID = providerTypeByIDFromConfigs( + providerConfigsQuery.data, + ); + return ( {pageTitle("Coder Agents", "AI Settings")} @@ -133,6 +143,7 @@ const CoderAgentsPage: FC = () => { titleGenerationModelOverrideData={titleGenerationModelQuery.data} exploreModelOverrideData={exploreModelOverrideQuery.data} modelConfigsData={modelConfigsQuery.data} + providerTypeByID={providerTypeByID} modelConfigsError={modelConfigsQuery.error} isLoadingModelConfigs={modelConfigsQuery.isLoading} isFetchingModelConfigs={modelConfigsQuery.isFetching} diff --git a/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.stories.tsx b/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.stories.tsx index 1523550e11..973c194d78 100644 --- a/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.stories.tsx +++ b/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.stories.tsx @@ -49,7 +49,7 @@ const generalModelConfig = buildModelConfig({ const claudeSonnetModelConfig = buildModelConfig({ id: "model-claude-sonnet-4", - provider: "anthropic", + ai_provider_id: "provider-anthropic", model: "claude-sonnet-4", display_name: "Claude Sonnet 4", context_limit: 200_000, @@ -64,7 +64,7 @@ const titleModelConfig = buildModelConfig({ const exploreFallbackModelConfig = buildModelConfig({ id: "model-explore-blank-display", - provider: "anthropic", + ai_provider_id: "provider-anthropic", model: "claude-sonnet-4-20250514", display_name: "", context_limit: 200_000, @@ -87,7 +87,7 @@ const titleDisabledModelConfig = buildModelConfig({ const exploreDisabledModelConfig = buildModelConfig({ id: "model-explore-disabled", - provider: "anthropic", + ai_provider_id: "provider-anthropic", model: "claude-haiku-legacy", display_name: "Claude Haiku Legacy", enabled: false, @@ -104,6 +104,11 @@ const allModelConfigs: TypesGen.ChatModelConfig[] = [ exploreDisabledModelConfig, ]; +const providerTypeByID = new Map([ + ["provider-1", "openai"], + ["provider-anthropic", "anthropic"], +]); + const buildArgs = ( overrides: Partial = {}, ): CoderAgentsPageViewProps => ({ @@ -118,6 +123,7 @@ const buildArgs = ( titleGenerationModelOverrideData: buildTitleGenerationModelOverrideData(), exploreModelOverrideData: buildOverrideData("explore"), modelConfigsData: allModelConfigs, + providerTypeByID, modelConfigsError: undefined, isLoadingModelConfigs: false, isFetchingModelConfigs: false, diff --git a/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.tsx b/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.tsx index 153dec7aa4..91ed5d0016 100644 --- a/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.tsx +++ b/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.tsx @@ -34,6 +34,7 @@ export interface CoderAgentsPageViewProps { titleGenerationModelOverrideData?: TypesGen.ChatModelOverrideResponse; exploreModelOverrideData?: TypesGen.ChatModelOverrideResponse; modelConfigsData: TypesGen.ChatModelConfig[] | undefined; + providerTypeByID: ReadonlyMap; modelConfigsError: unknown; isLoadingModelConfigs: boolean; isFetchingModelConfigs: boolean; @@ -83,6 +84,7 @@ export const CoderAgentsPageView: FC = ({ titleGenerationModelOverrideData, exploreModelOverrideData, modelConfigsData, + providerTypeByID, modelConfigsError, isLoadingModelConfigs, isFetchingModelConfigs, @@ -145,6 +147,7 @@ export const CoderAgentsPageView: FC = ({ description="Used by delegated agents that can edit files or run commands." modelOverrideData={generalModelOverrideData} enabledModelConfigs={enabledModelConfigs} + providerTypeByID={providerTypeByID} modelConfigsError={modelConfigsError} isLoading={isLoadingModelConfigs} onSaveModelOverride={onSaveGeneralModelOverride} @@ -158,6 +161,7 @@ export const CoderAgentsPageView: FC = ({ description="Leave unset to use Coder's title default, which prefers fast models from configured providers." modelOverrideData={titleGenerationModelOverrideData} enabledModelConfigs={enabledModelConfigs} + providerTypeByID={providerTypeByID} modelConfigsError={modelConfigsError} isLoading={isLoadingModelConfigs} onSaveModelOverride={onSaveTitleGenerationModel} @@ -172,6 +176,7 @@ export const CoderAgentsPageView: FC = ({ description="Used for read-only codebase exploration before work returns to the main agent." modelOverrideData={exploreModelOverrideData} enabledModelConfigs={enabledModelConfigs} + providerTypeByID={providerTypeByID} modelConfigsError={modelConfigsError} isLoading={isLoadingModelConfigs} onSaveModelOverride={onSaveExploreModelOverride} diff --git a/site/src/pages/AISettingsPage/CoderAgentsPage/components/SubagentModelOverrideSettings.tsx b/site/src/pages/AISettingsPage/CoderAgentsPage/components/SubagentModelOverrideSettings.tsx index 06826f200a..8a1294c0c5 100644 --- a/site/src/pages/AISettingsPage/CoderAgentsPage/components/SubagentModelOverrideSettings.tsx +++ b/site/src/pages/AISettingsPage/CoderAgentsPage/components/SubagentModelOverrideSettings.tsx @@ -27,6 +27,7 @@ interface SubagentModelOverrideSettingsProps { description?: ReactNode; modelOverrideData: ModelOverrideData | undefined; enabledModelConfigs: readonly TypesGen.ChatModelConfig[]; + providerTypeByID: ReadonlyMap; modelConfigsError: unknown; isLoading: boolean; onSaveModelOverride: ( @@ -43,9 +44,10 @@ interface SubagentModelOverrideSettingsProps { const toModelSelectorOption = ( modelConfig: TypesGen.ChatModelConfig, + providerTypeByID: ReadonlyMap, ): ModelSelectorOption => ({ id: modelConfig.id, - provider: modelConfig.provider, + provider: providerTypeByID.get(modelConfig.ai_provider_id) ?? "", model: modelConfig.model, displayName: modelConfig.display_name.trim() || modelConfig.model, contextLimit: modelConfig.context_limit, @@ -58,6 +60,7 @@ export const SubagentModelOverrideSettings: FC< description, modelOverrideData, enabledModelConfigs, + providerTypeByID, modelConfigsError, isLoading, onSaveModelOverride, @@ -71,7 +74,9 @@ export const SubagentModelOverrideSettings: FC< const { isSavedVisible, showSavedState } = useTemporarySavedState(); const hasLoadedModelOverride = modelOverrideData !== undefined; const isMalformedOverride = modelOverrideData?.is_malformed ?? false; - const enabledModelOptions = enabledModelConfigs.map(toModelSelectorOption); + const enabledModelOptions = enabledModelConfigs.map((modelConfig) => + toModelSelectorOption(modelConfig, providerTypeByID), + ); const form = useFormik({ enableReinitialize: true, diff --git a/site/src/pages/AISettingsPage/ModelsPage/ModelsPage.tsx b/site/src/pages/AISettingsPage/ModelsPage/ModelsPage.tsx index 49a28aa98d..67b54f7cea 100644 --- a/site/src/pages/AISettingsPage/ModelsPage/ModelsPage.tsx +++ b/site/src/pages/AISettingsPage/ModelsPage/ModelsPage.tsx @@ -8,6 +8,7 @@ import { import { useAuthenticated } from "#/hooks/useAuthenticated"; import { deriveProviderStates } from "#/modules/aiModels/providerStates"; import { RequirePermission } from "#/modules/permissions/RequirePermission"; +import { providerTypeByIDFromConfigs } from "#/pages/AgentsPage/utils/modelOptions"; import { pageTitle } from "#/utils/page"; import ModelsPageView from "./ModelsPageView"; @@ -21,8 +22,14 @@ const ModelsPage: FC = () => { const modelConfigsQuery = useQuery(chatModelConfigs()); const modelCatalogQuery = useQuery(chatModels()); + const providerTypeByID = providerTypeByIDFromConfigs( + providerConfigsQuery.data, + ); + const models = (modelConfigsQuery.data ?? []).slice().sort((a, b) => { - const cmp = a.provider.localeCompare(b.provider); + const aProvider = providerTypeByID.get(a.ai_provider_id) ?? ""; + const bProvider = providerTypeByID.get(b.ai_provider_id) ?? ""; + const cmp = aProvider.localeCompare(bProvider); return cmp !== 0 ? cmp : a.model.localeCompare(b.model); }); const providerStates = deriveProviderStates( @@ -48,6 +55,7 @@ const ModelsPage: FC = () => { } models={models} providerStates={providerStates} + providerTypeByID={providerTypeByID} /> ); diff --git a/site/src/pages/AISettingsPage/ModelsPage/ModelsPageView.stories.tsx b/site/src/pages/AISettingsPage/ModelsPage/ModelsPageView.stories.tsx index 6e5a32a16b..c4de1c4ffd 100644 --- a/site/src/pages/AISettingsPage/ModelsPage/ModelsPageView.stories.tsx +++ b/site/src/pages/AISettingsPage/ModelsPage/ModelsPageView.stories.tsx @@ -25,6 +25,11 @@ const meta: Meta = { MockAnthropicProviderState, MockBedrockProviderState, ], + providerTypeByID: new Map([ + ["prov-openai", "openai"], + ["prov-anthropic", "anthropic"], + ["prov-bedrock", "bedrock"], + ]), }, parameters: { reactRouter: reactRouterParameters({ diff --git a/site/src/pages/AISettingsPage/ModelsPage/ModelsPageView.tsx b/site/src/pages/AISettingsPage/ModelsPage/ModelsPageView.tsx index 321ae70dd5..80426c630a 100644 --- a/site/src/pages/AISettingsPage/ModelsPage/ModelsPageView.tsx +++ b/site/src/pages/AISettingsPage/ModelsPage/ModelsPageView.tsx @@ -99,6 +99,7 @@ interface ModelsPageViewProps { error: unknown; models: readonly ChatModelConfig[]; providerStates: readonly ProviderState[]; + providerTypeByID: ReadonlyMap; } const ModelsPageView: FC = ({ @@ -106,6 +107,7 @@ const ModelsPageView: FC = ({ error, models, providerStates, + providerTypeByID, }) => { const navigate = useNavigate(); const [page, setPage] = useState(1); @@ -264,6 +266,7 @@ const ModelsPageView: FC = ({ key={model.id} model={model} providerLabel={providerLabelByModelId.get(model.id) ?? ""} + providerTypeByID={providerTypeByID} onClick={() => void navigate(`/ai/settings/models/${model.id}`)} /> )) diff --git a/site/src/pages/AISettingsPage/ModelsPage/components/ModelForm.tsx b/site/src/pages/AISettingsPage/ModelsPage/components/ModelForm.tsx index 1f16c44d31..e6062e15d1 100644 --- a/site/src/pages/AISettingsPage/ModelsPage/components/ModelForm.tsx +++ b/site/src/pages/AISettingsPage/ModelsPage/components/ModelForm.tsx @@ -145,14 +145,12 @@ export const ModelForm: FC = ({ const selectedProviderConfigID = selectedProviderState?.providerConfig?.id; - const editingProviderConfigID = - editingModel?.ai_provider_id?.trim() ?? ""; + const editingProviderConfigID = editingModel?.ai_provider_id.trim() ?? ""; if (isEditing && editingModel) { const req: TypesGen.UpdateChatModelConfigRequest = { ...(selectedProviderConfigID && selectedProviderConfigID !== editingProviderConfigID && { - provider: selectedProviderState.provider, ai_provider_id: selectedProviderConfigID, }), ...(trimmedModel !== editingModel.model && { @@ -181,7 +179,6 @@ export const ModelForm: FC = ({ if (!selectedProviderState?.providerConfig) return; const req: TypesGen.CreateChatModelConfigRequest = { - provider: selectedProviderState.provider, ai_provider_id: selectedProviderState.providerConfig.id, model: trimmedModel, enabled: values.enabled, diff --git a/site/src/pages/AISettingsPage/ModelsPage/components/ModelRow.tsx b/site/src/pages/AISettingsPage/ModelsPage/components/ModelRow.tsx index d0212f0057..7aa7cc8d4c 100644 --- a/site/src/pages/AISettingsPage/ModelsPage/components/ModelRow.tsx +++ b/site/src/pages/AISettingsPage/ModelsPage/components/ModelRow.tsx @@ -10,6 +10,7 @@ import { ProviderIcon } from "#/pages/AISettingsPage/ProvidersPage/components/Pr type ModelRowProps = { model: ChatModelConfig; providerLabel: string; + providerTypeByID: ReadonlyMap; onClick: () => void; }; @@ -23,6 +24,7 @@ const formatContextLimit = (contextLimit: number): string => { export const ModelRow: FC = ({ model, providerLabel, + providerTypeByID, onClick, }) => { const clickableProps = useClickableTableRow({ onClick }); @@ -36,7 +38,9 @@ export const ModelRow: FC = ({ size="lg" className="flex shrink-0 items-center justify-center" > - +
= { spyOn(API, "getApiKey").mockRejectedValue(new Error("missing API key")); spyOn(API.experimental, "updateChat").mockResolvedValue(); spyOn(API.experimental, "getMCPServerConfigs").mockResolvedValue([]); + spyOn(API.experimental, "getUserAIProviderKeyConfigs").mockResolvedValue([ + { + provider: { + id: "provider-1", + type: "openai", + name: "openai", + display_name: "OpenAI", + enabled: true, + deleted: false, + }, + has_user_api_key: false, + has_provider_api_key: true, + byok_enabled: true, + }, + ]); return () => localStorage.removeItem(RIGHT_PANEL_OPEN_KEY); }, }; diff --git a/site/src/pages/AgentsPage/AgentChatPage.tsx b/site/src/pages/AgentsPage/AgentChatPage.tsx index dc6ee1ffc0..bafe6cacff 100644 --- a/site/src/pages/AgentsPage/AgentChatPage.tsx +++ b/site/src/pages/AgentsPage/AgentChatPage.tsx @@ -34,6 +34,7 @@ import { updateChatWorkspace, updateInfiniteChatsCache, userChatDebugLogging, + userChatProviderConfigs, userCompactionThresholds, } from "#/api/queries/chats"; import { deploymentSSHConfig } from "#/api/queries/deployment"; @@ -87,12 +88,11 @@ import { getAgentChatSendShortcut } from "./utils/agentChatSendShortcut"; import { type ParsedDraft, parseStoredDraft } from "./utils/draftStorage"; import { countConfiguredProviderConfigs, - getModelOptionsFromConfigs, getModelSelectorPlaceholder, getUnsupportedProviderNames, - hasConfiguredModelsInCatalog, hasUserFixableProviders, resolveModelOptionId, + resolveModelSelector, } from "./utils/modelOptions"; import { parsePullRequestUrl } from "./utils/pullRequest"; import { @@ -776,6 +776,7 @@ const AgentChatPage: FC = () => { ...chatProviderConfigs(), enabled: permissions.editDeploymentConfig, }); + const userProviderConfigsQuery = useQuery(userChatProviderConfigs()); const userThresholdsQuery = useQuery(userCompactionThresholds()); const preferencesQuery = useQuery(preferenceSettings()); const userDebugLoggingQuery = useQuery(userChatDebugLogging()); @@ -805,9 +806,15 @@ const AgentChatPage: FC = () => { void mcpServersQuery.refetch(); }; - const modelOptions = getModelOptionsFromConfigs( - chatModelConfigsQuery.data, - chatModelsQuery.data, + const { + options: modelOptions, + isModelCatalogLoading, + modelCatalog, + hasConfiguredModels, + } = resolveModelSelector( + chatModelConfigsQuery, + chatModelsQuery, + userProviderConfigsQuery, ); const modelConfigs = chatModelConfigsQuery.data ?? []; const providerCount = @@ -826,8 +833,6 @@ const AgentChatPage: FC = () => { const unsupportedProviderNames = getUnsupportedProviderNames( chatModelsQuery.data, ); - const modelCatalog = chatModelsQuery.data; - const isModelCatalogLoading = chatModelsQuery.isLoading; // Subscribe to live workspace updates so that agent status changes // (e.g. connected/disconnected) are reflected without a page refresh. @@ -1136,7 +1141,6 @@ const AgentChatPage: FC = () => { modelConfigs, ); const hasModelOptions = modelOptions.length > 0; - const hasConfiguredModels = hasConfiguredModelsInCatalog(modelCatalog); const hasUserFixableModelProviders = hasUserFixableProviders(modelCatalog); const modelSelectorPlaceholder = getModelSelectorPlaceholder( modelOptions, diff --git a/site/src/pages/AgentsPage/AgentCreatePage.tsx b/site/src/pages/AgentsPage/AgentCreatePage.tsx index b587b14121..6508f47356 100644 --- a/site/src/pages/AgentsPage/AgentCreatePage.tsx +++ b/site/src/pages/AgentsPage/AgentCreatePage.tsx @@ -10,6 +10,7 @@ import { createChat, mcpServerConfigs, userChatPersonalModelOverrides, + userChatProviderConfigs, } from "#/api/queries/chats"; import { preferenceSettings } from "#/api/queries/users"; import { workspaces } from "#/api/queries/workspaces"; @@ -27,8 +28,8 @@ import { getAgentChatSendShortcut } from "./utils/agentChatSendShortcut"; import { getChimeEnabled, setChimeEnabled } from "./utils/chime"; import { countConfiguredProviderConfigs, - getModelOptionsFromConfigs, getUnsupportedProviderNames, + resolveModelSelector, } from "./utils/modelOptions"; import { buildAgentChatPath } from "./utils/navigation"; @@ -46,6 +47,7 @@ const AgentCreatePage: FC = () => { ...chatProviderConfigs(), enabled: permissions.editDeploymentConfig, }); + const userProviderConfigsQuery = useQuery(userChatProviderConfigs()); const personalModelOverridesQuery = useQuery( userChatPersonalModelOverrides(), ); @@ -56,10 +58,12 @@ const AgentCreatePage: FC = () => { const webPush = useWebpushNotifications(); const [chimeEnabled, setChimeEnabledState] = useState(getChimeEnabled); - const catalogModelOptions = getModelOptionsFromConfigs( - chatModelConfigsQuery.data, - chatModelsQuery.data, - ); + const { options: catalogModelOptions, isModelCatalogLoading } = + resolveModelSelector( + chatModelConfigsQuery, + chatModelsQuery, + userProviderConfigsQuery, + ); const providerCount = permissions.editDeploymentConfig && chatProviderConfigsQuery.isSuccess && @@ -166,7 +170,7 @@ const AgentCreatePage: FC = () => { modelCount={modelCount} unsupportedProviderNames={unsupportedProviderNames} modelConfigs={chatModelConfigsQuery.data ?? []} - isModelCatalogLoading={chatModelsQuery.isLoading} + isModelCatalogLoading={isModelCatalogLoading} isModelConfigsLoading={chatModelConfigsQuery.isLoading} rootPersonalModelOverride={rootPersonalModelOverride} isPersonalModelOverridesLoading={personalModelOverridesQuery.isLoading} diff --git a/site/src/pages/AgentsPage/AgentSettingsAPIKeysPage.stories.tsx b/site/src/pages/AgentsPage/AgentSettingsAPIKeysPage.stories.tsx index 33f8175a37..456e9549c9 100644 --- a/site/src/pages/AgentsPage/AgentSettingsAPIKeysPage.stories.tsx +++ b/site/src/pages/AgentsPage/AgentSettingsAPIKeysPage.stories.tsx @@ -24,7 +24,7 @@ const createProvider = ( const createModel = ( overrides: Partial & - Pick, + Pick, ): ChatModelConfig => ({ ...MockChatModelConfig, created_at: "2026-03-01T00:00:00.000Z", @@ -40,7 +40,7 @@ const baseProvider = createProvider({ const baseModel = createModel({ id: "model-1", - provider: "openai", + ai_provider_id: "prov-1", display_name: "GPT-4o", model: "gpt-4o", }); @@ -127,7 +127,7 @@ export const WithFallback: Story = { models: [ createModel({ id: "model-1", - provider: "anthropic", + ai_provider_id: "prov-1", display_name: "Claude Sonnet 4", model: "claude-sonnet-4-20250514", }), @@ -160,19 +160,19 @@ export const MultipleProviders: Story = { models: [ createModel({ id: "model-openai-1", - provider: "openai", + ai_provider_id: "prov-openai", display_name: "GPT-4o", model: "gpt-4o", }), createModel({ id: "model-anthropic-1", - provider: "anthropic", + ai_provider_id: "prov-anthropic", display_name: "Claude Sonnet 4", model: "claude-sonnet-4-20250514", }), createModel({ id: "model-anthropic-2", - provider: "anthropic", + ai_provider_id: "prov-anthropic", display_name: "Claude Opus 4", model: "claude-opus-4-20250514", }), @@ -226,7 +226,7 @@ export const SavingSingleProvider: Story = { ...baseModels, createModel({ id: "model-2", - provider: "anthropic", + ai_provider_id: "prov-2", display_name: "Claude Sonnet 4", model: "claude-sonnet-4-20250514", }), @@ -330,19 +330,19 @@ export const ShowsProviderStatuses: Story = { models: [ createModel({ id: "model-openai-1", - provider: "openai", + ai_provider_id: "prov-openai", display_name: "GPT-4o", model: "gpt-4o", }), createModel({ id: "model-anthropic-1", - provider: "anthropic", + ai_provider_id: "prov-anthropic", display_name: "Claude Sonnet 4", model: "claude-sonnet-4-20250514", }), createModel({ id: "model-google-1", - provider: "google", + ai_provider_id: "prov-google", display_name: "Gemini 2.5 Pro", model: "gemini-2.5-pro", }), diff --git a/site/src/pages/AgentsPage/AgentSettingsAPIKeysPageView.tsx b/site/src/pages/AgentsPage/AgentSettingsAPIKeysPageView.tsx index a9638a7cf3..fe7fd73b34 100644 --- a/site/src/pages/AgentsPage/AgentSettingsAPIKeysPageView.tsx +++ b/site/src/pages/AgentsPage/AgentSettingsAPIKeysPageView.tsx @@ -63,13 +63,11 @@ interface ProviderKeyPanelProps { isRemoving: boolean; onSave: (providerConfigId: string, apiKey: string) => void; onRemove: (providerConfigId: string) => void; - hasAmbiguousProviderType: boolean; } const ProviderKeyPanel: FC = ({ provider, models, - hasAmbiguousProviderType, isModelsLoading, areModelsUnavailable, isSaving, @@ -85,15 +83,9 @@ const ProviderKeyPanel: FC = ({ const [isDeleteDialogOpen, setIsDeleteDialogOpen] = useState(false); const status = getProviderStatus(provider); - const enabledModels = models.filter((model) => { - return ( - model.enabled && - (model.ai_provider_id === provider.provider_id || - (!model.ai_provider_id && - !hasAmbiguousProviderType && - model.provider === provider.provider)) - ); - }); + const enabledModels = models.filter( + (model) => model.enabled && model.ai_provider_id === provider.provider_id, + ); const hasApiKeyValue = apiKey.trim().length > 0; const hasAPIKeyWhitespace = apiKey !== API_KEY_PLACEHOLDER && apiKey.trim() !== apiKey; @@ -273,14 +265,6 @@ export const AgentSettingsAPIKeysPageView: FC< onSave, onRemove, }) => { - const providerTypeCounts = new Map(); - for (const item of providerItems) { - providerTypeCounts.set( - item.provider.provider, - (providerTypeCounts.get(item.provider.provider) ?? 0) + 1, - ); - } - return (
@@ -311,9 +295,6 @@ export const AgentSettingsAPIKeysPageView: FC< isRemoving={item.isRemoving} onSave={onSave} onRemove={onRemove} - hasAmbiguousProviderType={ - (providerTypeCounts.get(item.provider.provider) ?? 0) > 1 - } /> ))}
diff --git a/site/src/pages/AgentsPage/AgentSettingsCompactionPage.tsx b/site/src/pages/AgentsPage/AgentSettingsCompactionPage.tsx index 1109ee8f9d..2701339724 100644 --- a/site/src/pages/AgentsPage/AgentSettingsCompactionPage.tsx +++ b/site/src/pages/AgentsPage/AgentSettingsCompactionPage.tsx @@ -4,13 +4,16 @@ import { chatModelConfigs, deleteUserCompactionThreshold, updateUserCompactionThreshold, + userChatProviderConfigs, userCompactionThresholds, } from "#/api/queries/chats"; import { AgentSettingsCompactionPageView } from "./AgentSettingsCompactionPageView"; +import { providerTypeByIDFromUserConfigs } from "./utils/modelOptions"; const AgentSettingsCompactionPage: FC = () => { const queryClient = useQueryClient(); const modelConfigsQuery = useQuery(chatModelConfigs()); + const providerConfigsQuery = useQuery(userChatProviderConfigs()); const thresholdsQuery = useQuery(userCompactionThresholds()); const saveThresholdMutation = useMutation( updateUserCompactionThreshold(queryClient), @@ -31,9 +34,14 @@ const AgentSettingsCompactionPage: FC = () => { const handleResetThreshold = (modelConfigId: string) => resetThresholdMutation.mutateAsync(modelConfigId); + const providerTypeByID = providerTypeByIDFromUserConfigs( + providerConfigsQuery.data, + ); + return ( ([["prov-openai", "openai"]]), modelConfigsError: undefined, isLoadingModelConfigs: false, thresholds: [ diff --git a/site/src/pages/AgentsPage/AgentSettingsCompactionPageView.tsx b/site/src/pages/AgentsPage/AgentSettingsCompactionPageView.tsx index 4dd47ee35a..8df3a5f4d5 100644 --- a/site/src/pages/AgentsPage/AgentSettingsCompactionPageView.tsx +++ b/site/src/pages/AgentsPage/AgentSettingsCompactionPageView.tsx @@ -5,6 +5,7 @@ import { UserCompactionThresholdSettings } from "./components/UserCompactionThre export interface AgentSettingsCompactionPageViewProps { modelConfigsData: TypesGen.ChatModelConfig[] | undefined; + providerTypeByID: ReadonlyMap; modelConfigsError: unknown; isLoadingModelConfigs: boolean; thresholds: readonly TypesGen.UserChatCompactionThreshold[] | undefined; @@ -21,6 +22,7 @@ export const AgentSettingsCompactionPageView: FC< AgentSettingsCompactionPageViewProps > = ({ modelConfigsData, + providerTypeByID, modelConfigsError, isLoadingModelConfigs, thresholds, @@ -37,6 +39,7 @@ export const AgentSettingsCompactionPageView: FC< /> { const queryClient = useQueryClient(); const overridesQuery = useQuery(userChatPersonalModelOverrides()); const chatModelsQuery = useQuery(chatModels()); const modelConfigsQuery = useQuery(chatModelConfigs()); + const providerConfigsQuery = useQuery(userChatProviderConfigs()); const saveRootModelOverrideMutation = useMutation( updateUserChatPersonalModelOverride(queryClient), ); @@ -24,13 +26,12 @@ const AgentSettingsUserAgentsPage: FC = () => { const saveExploreModelOverrideMutation = useMutation( updateUserChatPersonalModelOverride(queryClient), ); - const modelOptions = getModelOptionsFromConfigs( - modelConfigsQuery.data, - chatModelsQuery.data, + const { options: modelOptions, isModelCatalogLoading } = resolveModelSelector( + modelConfigsQuery, + chatModelsQuery, + providerConfigsQuery, ); const modelConfigsError = modelConfigsQuery.error ?? chatModelsQuery.error; - const isLoadingModels = - chatModelsQuery.isLoading || modelConfigsQuery.isLoading; const saveModelOverride = ( context: TypesGen.ChatPersonalModelOverrideContext, @@ -56,7 +57,7 @@ const AgentSettingsUserAgentsPage: FC = () => { modelOptions={modelOptions} modelConfigs={modelConfigsQuery.data ?? []} modelConfigsError={modelConfigsError} - isLoadingModels={isLoadingModels} + isLoadingModels={isModelCatalogLoading} onSaveRootModelOverride={saveModelOverride( "root", saveRootModelOverrideMutation, diff --git a/site/src/pages/AgentsPage/AgentSettingsUserAgentsPageView.stories.tsx b/site/src/pages/AgentsPage/AgentSettingsUserAgentsPageView.stories.tsx index a7e56cbd5a..14b455ecc5 100644 --- a/site/src/pages/AgentsPage/AgentSettingsUserAgentsPageView.stories.tsx +++ b/site/src/pages/AgentsPage/AgentSettingsUserAgentsPageView.stories.tsx @@ -64,7 +64,7 @@ const defaultModelConfig = buildModelConfig({ const claudeModelConfig = buildModelConfig({ id: "model-claude-sonnet-4", - provider: "anthropic", + ai_provider_id: "provider-anthropic", model: "claude-sonnet-4", display_name: "Claude Sonnet 4", context_limit: 200_000, @@ -79,7 +79,7 @@ const disabledModelConfig = buildModelConfig({ const inaccessibleModelConfig = buildModelConfig({ id: "model-inaccessible", - provider: "bedrock", + ai_provider_id: "provider-bedrock", model: "claude-3-5-sonnet", display_name: "Bedrock Claude", }); @@ -94,14 +94,14 @@ const modelConfigs = [ const modelOptions: ModelSelectorOption[] = [ { id: defaultModelConfig.id, - provider: defaultModelConfig.provider, + provider: "openai", model: defaultModelConfig.model, displayName: defaultModelConfig.display_name, contextLimit: defaultModelConfig.context_limit, }, { id: claudeModelConfig.id, - provider: claudeModelConfig.provider, + provider: "anthropic", model: claudeModelConfig.model, displayName: claudeModelConfig.display_name, contextLimit: claudeModelConfig.context_limit, diff --git a/site/src/pages/AgentsPage/AgentsPage.tsx b/site/src/pages/AgentsPage/AgentsPage.tsx index e2b241e678..03ec1b8403 100644 --- a/site/src/pages/AgentsPage/AgentsPage.tsx +++ b/site/src/pages/AgentsPage/AgentsPage.tsx @@ -39,6 +39,7 @@ import { updateChatTitle, updateInfiniteChatsCache, userChatPersonalModelOverrides, + userChatProviderConfigs, } from "#/api/queries/chats"; import { invalidateWorkspaceMutationQueries, @@ -64,7 +65,10 @@ import { shouldNavigateAfterArchive, } from "./utils/agentWorkspaceUtils"; import { maybePlayChime } from "./utils/chime"; -import { getModelOptionsFromConfigs } from "./utils/modelOptions"; +import { + getModelOptionsFromConfigs, + providerTypeByIDFromUserConfigs, +} from "./utils/modelOptions"; import { clearPersistedRightPanelState } from "./utils/rightPanelTabStorage"; import { clearPersistedSidebarTabId } from "./utils/sidebarTabStorage"; import { @@ -165,6 +169,7 @@ const AgentsPage: FC = () => { // deduplicates the requests. const chatModelsQuery = useQuery(chatModels()); const chatModelConfigsQuery = useQuery(chatModelConfigs()); + const chatProviderConfigsQuery = useQuery(userChatProviderConfigs()); const personalModelOverridesQuery = useQuery( userChatPersonalModelOverrides(), ); @@ -318,6 +323,7 @@ const AgentsPage: FC = () => { const catalogModelOptions = getModelOptionsFromConfigs( chatModelConfigsQuery.data, chatModelsQuery.data, + providerTypeByIDFromUserConfigs(chatProviderConfigsQuery.data), ); const chatList = chatsQuery.data?.pages.flat() ?? []; const isArchiving = diff --git a/site/src/pages/AgentsPage/AgentsPageView.stories.tsx b/site/src/pages/AgentsPage/AgentsPageView.stories.tsx index aa34f6ead4..de3972ca45 100644 --- a/site/src/pages/AgentsPage/AgentsPageView.stories.tsx +++ b/site/src/pages/AgentsPage/AgentsPageView.stories.tsx @@ -69,7 +69,7 @@ const defaultModelOptions: ModelSelectorOption[] = [ const defaultModelConfigs: TypesGen.ChatModelConfig[] = [ { id: defaultModelConfigID, - provider: "openai", + ai_provider_id: "provider-openai", model: "gpt-4o", display_name: "GPT-4o", enabled: true, @@ -180,6 +180,7 @@ const AgentsRouteElement = () => ( is_malformed: false, }} modelConfigsData={[]} + providerTypeByID={new Map()} modelConfigsError={undefined} isLoadingModelConfigs={false} isFetchingModelConfigs={false} @@ -436,7 +437,7 @@ const meta: Meta = { spyOn(API.experimental, "getChatModelConfigs").mockResolvedValue([ { id: defaultModelConfigID, - provider: "openai", + ai_provider_id: "provider-openai", model: "gpt-4o", display_name: "GPT-4o", enabled: true, @@ -447,6 +448,21 @@ const meta: Meta = { updated_at: "2026-02-18T00:00:00.000Z", }, ]); + spyOn(API.experimental, "getUserAIProviderKeyConfigs").mockResolvedValue([ + { + provider: { + id: "provider-openai", + type: "openai", + name: "openai", + display_name: "OpenAI", + enabled: true, + deleted: false, + }, + has_user_api_key: false, + has_provider_api_key: true, + byok_enabled: true, + }, + ]); spyOn(API.experimental, "getMCPServerConfigs").mockResolvedValue([]); spyOn(API.experimental, "getChatDebugLogging").mockResolvedValue({ allow_users: false, diff --git a/site/src/pages/AgentsPage/components/AgentChatInput.tsx b/site/src/pages/AgentsPage/components/AgentChatInput.tsx index 0ef3c7bb87..cec54bf727 100644 --- a/site/src/pages/AgentsPage/components/AgentChatInput.tsx +++ b/site/src/pages/AgentsPage/components/AgentChatInput.tsx @@ -66,7 +66,6 @@ import { chatAttachmentAcceptAttribute, isChatAttachmentFile, } from "../utils/chatAttachments"; -import { formatProviderLabel } from "../utils/modelOptions"; import { AgentSetupNotice } from "./AgentSetupNotice"; import { AttachmentPreview, @@ -1408,7 +1407,6 @@ export const AgentChatInput: FC = ({ options={modelOptions} disabled={isDisabled} placeholder={modelSelectorPlaceholder} - formatProviderLabel={formatProviderLabel} className="md:shrink" dropdownSide="top" dropdownAlign="center" diff --git a/site/src/pages/AgentsPage/components/AgentCreateForm.stories.tsx b/site/src/pages/AgentsPage/components/AgentCreateForm.stories.tsx index 52af7098aa..34df9879d3 100644 --- a/site/src/pages/AgentsPage/components/AgentCreateForm.stories.tsx +++ b/site/src/pages/AgentsPage/components/AgentCreateForm.stories.tsx @@ -61,7 +61,7 @@ const defaultModelConfigs: TypesGen.ChatModelConfig[] = [ buildModelConfig({ is_default: true }), buildModelConfig({ id: claudeModelConfigID, - provider: "anthropic", + ai_provider_id: "provider-anthropic", model: "claude-sonnet-4", display_name: "Claude Sonnet 4", context_limit: 200_000, diff --git a/site/src/pages/AgentsPage/components/ChatConversation/chatHelpers.test.ts b/site/src/pages/AgentsPage/components/ChatConversation/chatHelpers.test.ts index 6212f1f47d..3d3f2c544c 100644 --- a/site/src/pages/AgentsPage/components/ChatConversation/chatHelpers.test.ts +++ b/site/src/pages/AgentsPage/components/ChatConversation/chatHelpers.test.ts @@ -191,40 +191,6 @@ describe("resolveModelFromChatConfig", () => { ); }); - it("matches by provider:model combined candidate", () => { - // The model field alone doesn't match an option id, but - // provider + model concatenated does. - const config = { model: "gpt-4", provider: "openai" }; - expect(resolveModelFromChatConfig(config, options)).toBe("openai:gpt-4"); - }); - - it("falls back to model field match on option.model property", () => { - // Neither `model` nor `provider:model` match an option id, - // so the function falls through to matching option.model. - const altOptions: ModelSelectorOption[] = [ - buildOption("custom-id-1", "openai", "gpt-4"), - ]; - const config = { model: "gpt-4", provider: "openai" }; - expect(resolveModelFromChatConfig(config, altOptions)).toBe("custom-id-1"); - }); - - it("falls back to model field match ignoring provider when provider is absent", () => { - const altOptions: ModelSelectorOption[] = [ - buildOption("custom-id-1", "openai", "gpt-4"), - ]; - const config = { model: "gpt-4" }; - expect(resolveModelFromChatConfig(config, altOptions)).toBe("custom-id-1"); - }); - - it("respects provider when matching on option.model", () => { - const altOptions: ModelSelectorOption[] = [ - buildOption("id-a", "azure", "gpt-4"), - buildOption("id-b", "openai", "gpt-4"), - ]; - const config = { model: "gpt-4", provider: "openai" }; - expect(resolveModelFromChatConfig(config, altOptions)).toBe("id-b"); - }); - it("returns first option when no match is found", () => { const config = { model: "unknown-model" }; expect(resolveModelFromChatConfig(config, options)).toBe("openai:gpt-4"); diff --git a/site/src/pages/AgentsPage/components/ChatConversation/chatHelpers.ts b/site/src/pages/AgentsPage/components/ChatConversation/chatHelpers.ts index e0877b0525..083c4d8210 100644 --- a/site/src/pages/AgentsPage/components/ChatConversation/chatHelpers.ts +++ b/site/src/pages/AgentsPage/components/ChatConversation/chatHelpers.ts @@ -81,27 +81,11 @@ export const resolveModelFromChatConfig = ( const typedModelConfig = modelConfig as Record; const model = asString(typedModelConfig.model); - const provider = asString(typedModelConfig.provider); - - const candidates = [model]; - if (provider && model) { - candidates.push(`${provider}:${model}`); - } - - for (const candidate of candidates) { - const match = modelOptions.find((option) => option.id === candidate); - if (match) { - return match.id; - } - } if (model) { - const modelMatch = modelOptions.find( - (option) => - option.model === model && (!provider || option.provider === provider), - ); - if (modelMatch) { - return modelMatch.id; + const match = modelOptions.find((option) => option.id === model); + if (match) { + return match.id; } } diff --git a/site/src/pages/AgentsPage/components/ChatElements/ModelSelector.tsx b/site/src/pages/AgentsPage/components/ChatElements/ModelSelector.tsx index 8b07d1b242..1e5da57ccc 100644 --- a/site/src/pages/AgentsPage/components/ChatElements/ModelSelector.tsx +++ b/site/src/pages/AgentsPage/components/ChatElements/ModelSelector.tsx @@ -13,6 +13,7 @@ import { TooltipProvider, TooltipTrigger, } from "#/components/Tooltip/Tooltip"; +import { formatProviderLabel as defaultFormatProviderLabel } from "#/utils/aiProviders"; import { cn } from "#/utils/cn"; export interface ModelSelectorOption { @@ -41,14 +42,6 @@ interface ModelSelectorProps { enableMobileFullWidthDropdown?: boolean; } -const defaultFormatProviderLabel = (provider: string): string => { - const normalized = provider.trim().toLowerCase(); - if (!normalized) { - return "Unknown"; - } - return `${normalized[0].toUpperCase()}${normalized.slice(1)}`; -}; - const formatContextLimit = (tokens: number): string => { if (tokens >= 1_000_000) { const m = tokens / 1_000_000; diff --git a/site/src/pages/AgentsPage/components/ChatsSidebar/ChatsSidebar.stories.tsx b/site/src/pages/AgentsPage/components/ChatsSidebar/ChatsSidebar.stories.tsx index 01c6e79b7c..7b08ccf03d 100644 --- a/site/src/pages/AgentsPage/components/ChatsSidebar/ChatsSidebar.stories.tsx +++ b/site/src/pages/AgentsPage/components/ChatsSidebar/ChatsSidebar.stories.tsx @@ -45,7 +45,7 @@ const defaultModelOptions: ModelSelectorOption[] = [ const defaultModelConfigs: TypesGen.ChatModelConfig[] = [ { id: "config-openai-gpt-4o", - provider: "openai", + ai_provider_id: "prov-1", model: "gpt-4o", display_name: "GPT-4o", enabled: true, diff --git a/site/src/pages/AgentsPage/components/ChatsSidebar/ChatsSidebar.test.tsx b/site/src/pages/AgentsPage/components/ChatsSidebar/ChatsSidebar.test.tsx index 70e5d53efa..c7351226a8 100644 --- a/site/src/pages/AgentsPage/components/ChatsSidebar/ChatsSidebar.test.tsx +++ b/site/src/pages/AgentsPage/components/ChatsSidebar/ChatsSidebar.test.tsx @@ -584,7 +584,7 @@ describe("ChatsSidebar model display names", () => { const modelConfigs: TypesGen.ChatModelConfig[] = [ { id: "config-fast", - provider: "openai", + ai_provider_id: "prov-openai", model: "gpt-4o", display_name: "GPT-4o (Fast)", enabled: true, @@ -596,7 +596,7 @@ describe("ChatsSidebar model display names", () => { }, { id: "config-quality", - provider: "openai", + ai_provider_id: "prov-openai", model: "gpt-4o", display_name: "GPT-4o (Quality)", enabled: true, @@ -629,33 +629,6 @@ describe("ChatsSidebar model display names", () => { expect(queryByText("GPT-4o (Fast)")).not.toBeInTheDocument(); }); - it("falls back to legacy provider/model matching when no config ID match exists", () => { - const { getByText } = render( - - - , - ); - - expect(getByText("GPT-4o (Quality)")).toBeInTheDocument(); - }); - it("shows Default model when last_model_config_id is a nil UUID", () => { const { getByText } = render( diff --git a/site/src/pages/AgentsPage/components/ChatsSidebar/tree/modelDisplayName.ts b/site/src/pages/AgentsPage/components/ChatsSidebar/tree/modelDisplayName.ts index 72d5ac6a0c..38c9276696 100644 --- a/site/src/pages/AgentsPage/components/ChatsSidebar/tree/modelDisplayName.ts +++ b/site/src/pages/AgentsPage/components/ChatsSidebar/tree/modelDisplayName.ts @@ -1,5 +1,4 @@ import type { Chat, ChatModelConfig } from "#/api/typesGenerated"; -import { getNormalizedModelRef } from "../../../utils/modelOptions"; import type { ModelSelectorOption } from "../../ChatElements"; import { asString } from "../../ChatElements/runtimeTypeUtils"; @@ -24,13 +23,6 @@ export const getModelDisplayName = ( (config) => config.id === normalizedModelConfigID, ); if (!modelConfig) { - const legacyModelOption = modelOptions.find( - (option) => - `${option.provider}:${option.model}` === normalizedModelConfigID, - ); - if (legacyModelOption?.displayName) { - return legacyModelOption.displayName; - } return "Default model"; } @@ -39,17 +31,6 @@ export const getModelDisplayName = ( return displayName; } - const { provider, model } = getNormalizedModelRef(modelConfig); - if (!provider || !model) { - return "Default model"; - } - - const fallbackModelOption = modelOptions.find( - (option) => option.provider === provider && option.model === model, - ); - if (fallbackModelOption?.displayName) { - return fallbackModelOption.displayName; - } - - return model; + const model = asString(modelConfig.model).trim(); + return model || "Default model"; }; diff --git a/site/src/pages/AgentsPage/components/UserCompactionThresholdSettings.stories.tsx b/site/src/pages/AgentsPage/components/UserCompactionThresholdSettings.stories.tsx index 75af74c08d..5b1d2771a1 100644 --- a/site/src/pages/AgentsPage/components/UserCompactionThresholdSettings.stories.tsx +++ b/site/src/pages/AgentsPage/components/UserCompactionThresholdSettings.stories.tsx @@ -24,7 +24,7 @@ const mockModelConfigs: TypesGen.ChatModelConfig[] = [ { ...MockChatModelConfig, id: "model-2", - provider: "anthropic", + ai_provider_id: "provider-anthropic", model: "claude-sonnet", display_name: "Claude Sonnet", created_at: "2025-01-01T00:00:00Z", @@ -49,6 +49,10 @@ const meta = { decorators: [withAuthProvider, withDashboardProvider], args: { modelConfigs: mockModelConfigs, + providerTypeByID: new Map([ + ["provider-1", "openai"], + ["provider-anthropic", "anthropic"], + ]), thresholds: [], isThresholdsLoading: false, thresholdsError: undefined, diff --git a/site/src/pages/AgentsPage/components/UserCompactionThresholdSettings.tsx b/site/src/pages/AgentsPage/components/UserCompactionThresholdSettings.tsx index 70deedf17b..99a6a11063 100644 --- a/site/src/pages/AgentsPage/components/UserCompactionThresholdSettings.tsx +++ b/site/src/pages/AgentsPage/components/UserCompactionThresholdSettings.tsx @@ -29,6 +29,7 @@ import { ProviderIcon } from "./ChatModelAdminPanel/ProviderIcon"; interface UserCompactionThresholdSettingsProps { modelConfigs: readonly TypesGen.ChatModelConfig[]; + providerTypeByID: ReadonlyMap; modelConfigsError?: unknown; isLoadingModelConfigs?: boolean; thresholds: readonly TypesGen.UserChatCompactionThreshold[] | undefined; @@ -71,6 +72,7 @@ export const UserCompactionThresholdSettings: FC< UserCompactionThresholdSettingsProps > = ({ modelConfigs, + providerTypeByID, modelConfigsError, isLoadingModelConfigs, thresholds, @@ -282,7 +284,9 @@ export const UserCompactionThresholdSettings: FC< {modelName} diff --git a/site/src/pages/AgentsPage/utils/modelOptions.test.ts b/site/src/pages/AgentsPage/utils/modelOptions.test.ts index e8245f55bb..0f6d7bf57f 100644 --- a/site/src/pages/AgentsPage/utils/modelOptions.test.ts +++ b/site/src/pages/AgentsPage/utils/modelOptions.test.ts @@ -13,16 +13,18 @@ import { formatProviderLabel, getModelOptionsFromConfigs, getModelSelectorPlaceholder, - getNormalizedModelRef, getUnsupportedProviderNames, hasConfiguredProviderConfigs, hasUserFixableProviders, + providerTypeByIDFromConfigs, + providerTypeByIDFromUserConfigs, resolveModelOptionId, + resolveModelSelector, } from "./modelOptions"; const createConfig = ( overrides: Partial & - Pick, + Pick, ): ChatModelConfig => ({ ...MockChatModelConfig, context_limit: 0, @@ -32,6 +34,12 @@ const createConfig = ( ...overrides, }); +const providerTypeByID = new Map([ + ["prov-openai", "openai"], + ["prov-anthropic", "anthropic"], + ["prov-openrouter", "openrouter"], +]); + const createCatalog = ( providers: ChatModelsResponse["providers"], unsupportedProviders: ChatModelsResponse["unsupported_providers"] = [], @@ -54,20 +62,6 @@ const createProviderConfig = ( ...overrides, }); -describe("getNormalizedModelRef", () => { - it("returns empty strings for malformed values", () => { - expect(getNormalizedModelRef({ provider: undefined, model: null })).toEqual( - { provider: "", model: "" }, - ); - }); - - it("trims and normalizes provider values", () => { - expect( - getNormalizedModelRef({ provider: " OpenAI ", model: " gpt-4o " }), - ).toEqual({ provider: "openai", model: "gpt-4o" }); - }); -}); - describe("hasUserFixableProviders", () => { it("returns true when a provider needs a user API key", () => { const catalog = createCatalog([ @@ -267,31 +261,9 @@ describe("resolveModelOptionId", () => { expect(resolveModelOptionId("config-2", modelOptions)).toBe("config-2"); }); - it("returns the config ID for a legacy provider:model match", () => { - expect(resolveModelOptionId("openai:gpt-4o", modelOptions)).toBe( - "config-1", - ); - }); - it("returns an empty string when no option matches", () => { expect(resolveModelOptionId("openai:gpt-5", modelOptions)).toBe(""); }); - - it("returns the first duplicate legacy match deterministically", () => { - const duplicateModelOptions = [ - ...modelOptions, - { - id: "config-3", - provider: "openai", - model: "gpt-4o", - displayName: "GPT-4o duplicate", - }, - ] as const; - - expect(resolveModelOptionId("openai:gpt-4o", duplicateModelOptions)).toBe( - "config-1", - ); - }); }); describe("getModelOptionsFromConfigs", () => { @@ -299,14 +271,14 @@ describe("getModelOptionsFromConfigs", () => { const configs = [ createConfig({ id: "config-1", - provider: "openai", + ai_provider_id: "prov-openai", model: "gpt-4o", display_name: "GPT-4o (Fast)", context_limit: 128_000, }), createConfig({ id: "config-2", - provider: "openai", + ai_provider_id: "prov-openai", model: "gpt-4o", display_name: "GPT-4o (Quality)", context_limit: 128_000, @@ -320,7 +292,9 @@ describe("getModelOptionsFromConfigs", () => { }, ]); - expect(getModelOptionsFromConfigs(configs, catalog)).toEqual([ + expect( + getModelOptionsFromConfigs(configs, catalog, providerTypeByID), + ).toEqual([ { id: "config-1", provider: "openai", @@ -342,7 +316,7 @@ describe("getModelOptionsFromConfigs", () => { const configs = [ createConfig({ id: "config-1", - provider: "anthropic", + ai_provider_id: "prov-anthropic", model: "claude-sonnet-4-20250514", display_name: "Claude Sonnet", context_limit: 200_000, @@ -356,14 +330,16 @@ describe("getModelOptionsFromConfigs", () => { }, ]); - expect(getModelOptionsFromConfigs(configs, catalog)).toEqual([]); + expect( + getModelOptionsFromConfigs(configs, catalog, providerTypeByID), + ).toEqual([]); }); it("excludes disabled configs", () => { const configs = [ createConfig({ id: "config-1", - provider: "openai", + ai_provider_id: "prov-openai", model: "gpt-4o", display_name: "GPT-4o", enabled: false, @@ -371,7 +347,7 @@ describe("getModelOptionsFromConfigs", () => { }), createConfig({ id: "config-2", - provider: "openai", + ai_provider_id: "prov-openai", model: "gpt-4.1", display_name: "GPT-4.1", context_limit: 128_000, @@ -385,7 +361,9 @@ describe("getModelOptionsFromConfigs", () => { }, ]); - expect(getModelOptionsFromConfigs(configs, catalog)).toEqual([ + expect( + getModelOptionsFromConfigs(configs, catalog, providerTypeByID), + ).toEqual([ { id: "config-2", provider: "openai", @@ -400,7 +378,7 @@ describe("getModelOptionsFromConfigs", () => { const configs = [ createConfig({ id: "config-1", - provider: " openai ", + ai_provider_id: "prov-openai", model: " gpt-4o ", display_name: " ", context_limit: 0, @@ -414,7 +392,9 @@ describe("getModelOptionsFromConfigs", () => { }, ]); - expect(getModelOptionsFromConfigs(configs, catalog)).toEqual([ + expect( + getModelOptionsFromConfigs(configs, catalog, providerTypeByID), + ).toEqual([ { id: "config-1", provider: "openai", @@ -426,29 +406,33 @@ describe("getModelOptionsFromConfigs", () => { }); it("returns an empty array for null and undefined inputs", () => { - expect(getModelOptionsFromConfigs(null, null)).toEqual([]); - expect(getModelOptionsFromConfigs(undefined, undefined)).toEqual([]); + expect(getModelOptionsFromConfigs(null, null, providerTypeByID)).toEqual( + [], + ); + expect( + getModelOptionsFromConfigs(undefined, undefined, providerTypeByID), + ).toEqual([]); }); it("sorts options by provider and display name", () => { const configs = [ createConfig({ id: "config-openai-zeta", - provider: "openai", + ai_provider_id: "prov-openai", model: "gpt-z", display_name: "Zeta", context_limit: 32_000, }), createConfig({ id: "config-anthropic", - provider: "anthropic", + ai_provider_id: "prov-anthropic", model: "claude-sonnet-4-20250514", display_name: "Claude Sonnet", context_limit: 200_000, }), createConfig({ id: "config-openai-alpha", - provider: "openai", + ai_provider_id: "prov-openai", model: "gpt-a", display_name: "Alpha", context_limit: 32_000, @@ -468,7 +452,9 @@ describe("getModelOptionsFromConfigs", () => { ]); expect( - getModelOptionsFromConfigs(configs, catalog).map((option) => option.id), + getModelOptionsFromConfigs(configs, catalog, providerTypeByID).map( + (option) => option.id, + ), ).toEqual([ "config-anthropic", "config-openai-alpha", @@ -480,14 +466,14 @@ describe("getModelOptionsFromConfigs", () => { const configs = [ createConfig({ id: "config-1", - provider: "openrouter", + ai_provider_id: "prov-openrouter", model: "openai/gpt-4o", display_name: "GPT-4o via OpenRouter", context_limit: 128_000, }), createConfig({ id: "config-2", - provider: "openrouter", + ai_provider_id: "prov-openrouter", model: "anthropic/claude-sonnet-4-20250514", display_name: "Claude via OpenRouter", context_limit: 200_000, @@ -501,7 +487,9 @@ describe("getModelOptionsFromConfigs", () => { }, ]); - expect(getModelOptionsFromConfigs(configs, catalog)).toEqual([ + expect( + getModelOptionsFromConfigs(configs, catalog, providerTypeByID), + ).toEqual([ { id: "config-2", provider: "openrouter", @@ -518,6 +506,98 @@ describe("getModelOptionsFromConfigs", () => { }, ]); }); + + it("drops configs whose ai_provider_id is absent from the provider map", () => { + const configs = [ + createConfig({ + id: "config-1", + ai_provider_id: "prov-openai", + model: "gpt-4o", + display_name: "GPT-4o", + context_limit: 128_000, + }), + ]; + const catalog = createCatalog([ + { provider: "openai", available: true, models: [] }, + ]); + + expect(getModelOptionsFromConfigs(configs, catalog, new Map())).toEqual([]); + }); + + it("keeps only configs whose ai_provider_id resolves in the provider map", () => { + const configs = [ + createConfig({ + id: "config-openai", + ai_provider_id: "prov-openai", + model: "gpt-4o", + display_name: "GPT-4o", + context_limit: 128_000, + }), + createConfig({ + id: "config-orphan", + ai_provider_id: "prov-missing", + model: "claude-sonnet-4-20250514", + display_name: "Claude Sonnet", + context_limit: 200_000, + }), + ]; + const catalog = createCatalog([ + { provider: "openai", available: true, models: [] }, + { provider: "anthropic", available: true, models: [] }, + ]); + const partialMap = new Map([["prov-openai", "openai"]]); + + expect( + getModelOptionsFromConfigs(configs, catalog, partialMap).map( + (option) => option.id, + ), + ).toEqual(["config-openai"]); + }); +}); + +describe("providerTypeByIDFromConfigs", () => { + it("maps ChatProviderConfig.id to its provider type", () => { + const map = providerTypeByIDFromConfigs([ + { ...MockChatProviderConfig, id: "prov-openai", provider: "openai" }, + { + ...MockChatProviderConfig, + id: "prov-anthropic", + provider: "anthropic", + }, + ]); + + expect(map.get("prov-openai")).toBe("openai"); + expect(map.get("prov-anthropic")).toBe("anthropic"); + expect(map.size).toBe(2); + }); + + it("returns an empty map for nullish input", () => { + expect(providerTypeByIDFromConfigs(undefined).size).toBe(0); + expect(providerTypeByIDFromConfigs(null).size).toBe(0); + }); +}); + +describe("providerTypeByIDFromUserConfigs", () => { + it("maps UserChatProviderConfig.provider_id to its provider type", () => { + const map = providerTypeByIDFromUserConfigs([ + { + provider_id: "prov-openai", + provider: "openai", + display_name: "OpenAI", + has_user_api_key: false, + has_central_api_key_fallback: true, + byok_enabled: true, + }, + ]); + + expect(map.get("prov-openai")).toBe("openai"); + expect(map.size).toBe(1); + }); + + it("returns an empty map for nullish input", () => { + expect(providerTypeByIDFromUserConfigs(undefined).size).toBe(0); + expect(providerTypeByIDFromUserConfigs(null).size).toBe(0); + }); }); describe("getUnsupportedProviderNames", () => { @@ -563,3 +643,60 @@ describe("getUnsupportedProviderNames", () => { expect(getUnsupportedProviderNames(null)).toEqual([]); }); }); + +describe("resolveModelSelector", () => { + const config = createConfig({ + id: "config-openai", + ai_provider_id: "prov-openai", + model: "gpt-4o", + display_name: "GPT-4o", + context_limit: 128_000, + }); + const catalog = createCatalog([ + { provider: "openai", available: true, models: [] }, + ]); + const userProviderConfigs = [ + { + provider_id: "prov-openai", + provider: "openai", + display_name: "OpenAI", + has_user_api_key: false, + has_central_api_key_fallback: true, + byok_enabled: true, + }, + ]; + + it("stays loading and drops options while the provider query is pending", () => { + // Catalog + configs have resolved, but provider identity has not, so + // the provider map is empty. Options must be dropped and the flag must + // stay loading rather than flashing "No Models". + const state = resolveModelSelector( + { data: [config], isLoading: false }, + { data: catalog, isLoading: false }, + { data: undefined, isLoading: true }, + ); + + expect(state.isModelCatalogLoading).toBe(true); + expect(state.options).toEqual([]); + }); + + it("resolves options once every query settles", () => { + const state = resolveModelSelector( + { data: [config], isLoading: false }, + { data: catalog, isLoading: false }, + { data: userProviderConfigs, isLoading: false }, + ); + + expect(state.isModelCatalogLoading).toBe(false); + expect(state.modelCatalog).toBe(catalog); + expect(state.options).toEqual([ + { + id: "config-openai", + provider: "openai", + model: "gpt-4o", + displayName: "GPT-4o", + contextLimit: 128_000, + }, + ]); + }); +}); diff --git a/site/src/pages/AgentsPage/utils/modelOptions.ts b/site/src/pages/AgentsPage/utils/modelOptions.ts index c33fc0abef..3f1d1c799a 100644 --- a/site/src/pages/AgentsPage/utils/modelOptions.ts +++ b/site/src/pages/AgentsPage/utils/modelOptions.ts @@ -1,3 +1,4 @@ +import type { UseQueryResult } from "react-query"; import type * as TypesGen from "#/api/typesGenerated"; import type { ModelSelectorOption } from "../components/ChatElements"; import { @@ -5,22 +6,12 @@ import { asString, } from "../components/ChatElements/runtimeTypeUtils"; -type RuntimeModelRef = { - readonly provider?: unknown; - readonly model?: unknown; -}; - -type ModelRefLike = - | Pick - | Pick - | RuntimeModelRef; - type CatalogModelLike = | TypesGen.ChatModel - | (RuntimeModelRef & { + | { readonly id?: unknown; readonly display_name?: unknown; - }); + }; type CatalogProviderLike = Omit & { readonly models?: readonly CatalogModelLike[]; @@ -30,15 +21,6 @@ type ModelCatalogLike = { readonly providers?: readonly CatalogProviderLike[]; }; -type ModelOptionConfigLike = - | TypesGen.ChatModelConfig - | (RuntimeModelRef & { - readonly id?: unknown; - readonly display_name?: unknown; - readonly enabled?: unknown; - readonly context_limit?: unknown; - }); - export const hasConfiguredProviderConfigs = ( providerConfigs: readonly TypesGen.ChatProviderConfig[] | null | undefined, catalog: TypesGen.ChatModelsResponse | null | undefined, @@ -65,16 +47,6 @@ export const countConfiguredProviderConfigs = ( ); }; -export const getNormalizedModelRef = ( - value: ModelRefLike, -): { readonly provider: string; readonly model: string } => { - const modelRef = value ?? {}; - return { - provider: asString(modelRef.provider).trim().toLowerCase(), - model: asString(modelRef.model).trim(), - }; -}; - const getCatalogProviders = ( catalog: ModelCatalogLike | null | undefined, ): readonly CatalogProviderLike[] => { @@ -165,10 +137,9 @@ const getAvailableProviders = ( }; /** - * Resolves a stored model reference (config ID or legacy - * "provider:model" string) to the ID of a matching model option. - * Returns the matched option ID, or an empty string if no match is - * found. + * Resolves a stored model config ID to the ID of a matching model + * option. Returns the matched option ID, or an empty string when the + * stored ID is blank or no longer matches an available option. */ export const resolveModelOptionId = ( storedRef: string | null | undefined, @@ -184,19 +155,41 @@ export const resolveModelOptionId = ( return directMatch.id; } - const legacyMatch = modelOptions.find( - (option) => `${option.provider}:${option.model}` === normalized, - ); - if (legacyMatch) { - return legacyMatch.id; - } - return ""; }; +// providerTypeByIDFromConfigs and providerTypeByIDFromUserConfigs build +// the ai_provider_id -> provider-type lookup that getModelOptionsFromConfigs +// needs. The admin and user provider endpoints expose the provider id under +// different field names (id vs provider_id), so each source has its own +// helper to bake in the correct field and keep callers from mixing them up. +export const providerTypeByIDFromConfigs = ( + providerConfigs: readonly TypesGen.ChatProviderConfig[] | null | undefined, +): ReadonlyMap => + new Map( + (providerConfigs ?? []).map((providerConfig) => [ + providerConfig.id, + providerConfig.provider, + ]), + ); + +export const providerTypeByIDFromUserConfigs = ( + providerConfigs: + | readonly TypesGen.UserChatProviderConfig[] + | null + | undefined, +): ReadonlyMap => + new Map( + (providerConfigs ?? []).map((providerConfig) => [ + providerConfig.provider_id, + providerConfig.provider, + ]), + ); + export const getModelOptionsFromConfigs = ( configs: readonly TypesGen.ChatModelConfig[] | null | undefined, catalog: TypesGen.ChatModelsResponse | null | undefined, + providerTypeByID: ReadonlyMap, ): readonly ModelSelectorOption[] => { if (!configs || !catalog) { return []; @@ -205,13 +198,16 @@ export const getModelOptionsFromConfigs = ( const availableProviders = getAvailableProviders(catalog); const options: ModelSelectorOption[] = []; - for (const config of configs as readonly ModelOptionConfigLike[]) { - if (config.enabled !== true) { + for (const config of configs) { + if (!config.enabled) { continue; } - const configID = asString(config.id).trim(); - const { provider, model } = getNormalizedModelRef(config); + const configID = config.id.trim(); + const provider = asString(providerTypeByID.get(config.ai_provider_id)) + .trim() + .toLowerCase(); + const model = config.model.trim(); if (!configID || !provider || !model) { continue; } @@ -219,7 +215,7 @@ export const getModelOptionsFromConfigs = ( continue; } - const displayName = asString(config.display_name).trim() || model; + const displayName = config.display_name.trim() || model; const contextLimit = asNumber(config.context_limit); options.push({ id: configID, @@ -239,6 +235,45 @@ export const getModelOptionsFromConfigs = ( }); }; +// Read slice of a react-query result. The field types come from UseQueryResult +// by indexed access, not Pick (which would distribute over v5's status union), +// so they track the library rather than being hand-maintained. +type SelectorQuery = { + readonly data: UseQueryResult["data"]; + readonly isLoading: UseQueryResult["isLoading"]; +}; + +interface ModelSelectorState { + readonly options: readonly ModelSelectorOption[]; + readonly isModelCatalogLoading: boolean; + readonly modelCatalog: TypesGen.ChatModelsResponse | undefined; + readonly hasConfiguredModels: boolean; +} + +// Provider identity comes from a separate query (userChatProviderConfigs). +// Folding all three loading states into one flag here spares every caller the +// "configs loaded but providers still pending" window that would otherwise +// build an empty provider map, drop every option, and flash "No Models". +export const resolveModelSelector = ( + modelConfigs: SelectorQuery, + catalog: SelectorQuery, + userProviderConfigs: SelectorQuery< + readonly TypesGen.UserChatProviderConfig[] + >, +): ModelSelectorState => ({ + options: getModelOptionsFromConfigs( + modelConfigs.data, + catalog.data, + providerTypeByIDFromUserConfigs(userProviderConfigs.data), + ), + isModelCatalogLoading: + modelConfigs.isLoading || + catalog.isLoading || + userProviderConfigs.isLoading, + modelCatalog: catalog.data, + hasConfiguredModels: hasConfiguredModelsInCatalog(catalog.data), +}); + // getProviderForModelOption returns the provider string for the // currently-selected model option, or undefined when the selection // is not (yet) in the options list. Extracted so resize/budget logic diff --git a/site/src/testHelpers/chatModels.ts b/site/src/testHelpers/chatModels.ts index bc4e9a62ec..ceb557272a 100644 --- a/site/src/testHelpers/chatModels.ts +++ b/site/src/testHelpers/chatModels.ts @@ -7,7 +7,7 @@ import { MOCK_TIMESTAMP } from "./chatEntities"; export const MockChatModelConfig: ChatModelConfig = { id: "model-1", - provider: "openai", + ai_provider_id: "provider-1", model: "gpt-5", display_name: "gpt-5", enabled: true,