diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index 6c01c9a40ea71..bb6590deba1ce 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -3450,6 +3450,13 @@ func (q *querier) GetChatStreamSyncRows(ctx context.Context, ids []uuid.UUID) ([ return q.db.GetChatStreamSyncRows(ctx, ids) } +func (q *querier) GetChatSummaryGenerationModelOverride(ctx context.Context) (string, error) { + if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceDeploymentConfig); err != nil { + return "", err + } + return q.db.GetChatSummaryGenerationModelOverride(ctx) +} + func (q *querier) GetChatSystemPrompt(ctx context.Context) (string, error) { // The system prompt is a deployment-wide setting read during chat // creation by every authenticated user, so no RBAC policy check @@ -8711,6 +8718,13 @@ func (q *querier) UpsertChatRetentionDays(ctx context.Context, retentionDays int return q.db.UpsertChatRetentionDays(ctx, retentionDays) } +func (q *querier) UpsertChatSummaryGenerationModelOverride(ctx context.Context, value string) error { + if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil { + return err + } + return q.db.UpsertChatSummaryGenerationModelOverride(ctx, value) +} + func (q *querier) UpsertChatSystemPrompt(ctx context.Context, value string) error { if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil { return err diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index f1bbd747d00a4..b76bbb79960a2 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -1201,6 +1201,10 @@ func (s *MethodTestSuite) TestChats() { dbm.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return("", nil).AnyTimes() check.Args().Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead) })) + s.Run("GetChatSummaryGenerationModelOverride", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + dbm.EXPECT().GetChatSummaryGenerationModelOverride(gomock.Any()).Return("", nil).AnyTimes() + check.Args().Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead) + })) s.Run("GetChatPlanModeInstructions", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { dbm.EXPECT().GetChatPlanModeInstructions(gomock.Any()).Return("", nil).AnyTimes() check.Args().Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate) @@ -1661,6 +1665,10 @@ func (s *MethodTestSuite) TestChats() { dbm.EXPECT().UpsertChatTitleGenerationModelOverride(gomock.Any(), "").Return(nil).AnyTimes() check.Args("").Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate) })) + s.Run("UpsertChatSummaryGenerationModelOverride", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + dbm.EXPECT().UpsertChatSummaryGenerationModelOverride(gomock.Any(), "").Return(nil).AnyTimes() + check.Args("").Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate) + })) s.Run("UpsertChatPlanModeInstructions", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { dbm.EXPECT().UpsertChatPlanModeInstructions(gomock.Any(), "").Return(nil).AnyTimes() check.Args("").Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate) diff --git a/coderd/database/dbmetrics/querymetrics.go b/coderd/database/dbmetrics/querymetrics.go index 11e61d95f1802..dd0e4a1ffc86e 100644 --- a/coderd/database/dbmetrics/querymetrics.go +++ b/coderd/database/dbmetrics/querymetrics.go @@ -1770,6 +1770,14 @@ func (m queryMetricsStore) GetChatStreamSyncRows(ctx context.Context, ids []uuid return r0, r1 } +func (m queryMetricsStore) GetChatSummaryGenerationModelOverride(ctx context.Context) (string, error) { + start := time.Now() + r0, r1 := m.s.GetChatSummaryGenerationModelOverride(ctx) + m.queryLatencies.WithLabelValues("GetChatSummaryGenerationModelOverride").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetChatSummaryGenerationModelOverride").Inc() + return r0, r1 +} + func (m queryMetricsStore) GetChatSystemPrompt(ctx context.Context) (string, error) { start := time.Now() r0, r1 := m.s.GetChatSystemPrompt(ctx) @@ -6226,6 +6234,14 @@ func (m queryMetricsStore) UpsertChatRetentionDays(ctx context.Context, retentio return r0 } +func (m queryMetricsStore) UpsertChatSummaryGenerationModelOverride(ctx context.Context, value string) error { + start := time.Now() + r0 := m.s.UpsertChatSummaryGenerationModelOverride(ctx, value) + m.queryLatencies.WithLabelValues("UpsertChatSummaryGenerationModelOverride").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "UpsertChatSummaryGenerationModelOverride").Inc() + return r0 +} + func (m queryMetricsStore) UpsertChatSystemPrompt(ctx context.Context, value string) error { start := time.Now() r0 := m.s.UpsertChatSystemPrompt(ctx, value) diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index 8c33e2a589f7b..acc86d366d638 100644 --- a/coderd/database/dbmock/dbmock.go +++ b/coderd/database/dbmock/dbmock.go @@ -3267,6 +3267,21 @@ func (mr *MockStoreMockRecorder) GetChatStreamSyncRows(ctx, ids any) *gomock.Cal return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatStreamSyncRows", reflect.TypeOf((*MockStore)(nil).GetChatStreamSyncRows), ctx, ids) } +// GetChatSummaryGenerationModelOverride mocks base method. +func (m *MockStore) GetChatSummaryGenerationModelOverride(ctx context.Context) (string, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetChatSummaryGenerationModelOverride", ctx) + ret0, _ := ret[0].(string) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetChatSummaryGenerationModelOverride indicates an expected call of GetChatSummaryGenerationModelOverride. +func (mr *MockStoreMockRecorder) GetChatSummaryGenerationModelOverride(ctx any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatSummaryGenerationModelOverride", reflect.TypeOf((*MockStore)(nil).GetChatSummaryGenerationModelOverride), ctx) +} + // GetChatSystemPrompt mocks base method. func (m *MockStore) GetChatSystemPrompt(ctx context.Context) (string, error) { m.ctrl.T.Helper() @@ -11652,6 +11667,20 @@ func (mr *MockStoreMockRecorder) UpsertChatRetentionDays(ctx, retentionDays any) return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpsertChatRetentionDays", reflect.TypeOf((*MockStore)(nil).UpsertChatRetentionDays), ctx, retentionDays) } +// UpsertChatSummaryGenerationModelOverride mocks base method. +func (m *MockStore) UpsertChatSummaryGenerationModelOverride(ctx context.Context, value string) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "UpsertChatSummaryGenerationModelOverride", ctx, value) + ret0, _ := ret[0].(error) + return ret0 +} + +// UpsertChatSummaryGenerationModelOverride indicates an expected call of UpsertChatSummaryGenerationModelOverride. +func (mr *MockStoreMockRecorder) UpsertChatSummaryGenerationModelOverride(ctx, value any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpsertChatSummaryGenerationModelOverride", reflect.TypeOf((*MockStore)(nil).UpsertChatSummaryGenerationModelOverride), ctx, value) +} + // UpsertChatSystemPrompt mocks base method. func (m *MockStore) UpsertChatSystemPrompt(ctx context.Context, value string) error { m.ctrl.T.Helper() diff --git a/coderd/database/querier.go b/coderd/database/querier.go index 4609115e40704..da532b2ab2de7 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -469,6 +469,7 @@ type sqlcQuerier interface { // A value of 0 disables chat purging entirely. GetChatRetentionDays(ctx context.Context) (int32, error) GetChatStreamSyncRows(ctx context.Context, ids []uuid.UUID) ([]GetChatStreamSyncRowsRow, error) + GetChatSummaryGenerationModelOverride(ctx context.Context) (string, error) GetChatSystemPrompt(ctx context.Context) (string, error) // GetChatSystemPromptConfig returns both chat system prompt settings in a // single read to avoid torn reads between separate site-config lookups. @@ -1546,6 +1547,7 @@ type sqlcQuerier interface { UpsertChatPersonalModelOverridesEnabled(ctx context.Context, enabled bool) error UpsertChatPlanModeInstructions(ctx context.Context, value string) error UpsertChatRetentionDays(ctx context.Context, retentionDays int32) error + UpsertChatSummaryGenerationModelOverride(ctx context.Context, value string) error UpsertChatSystemPrompt(ctx context.Context, value string) error UpsertChatTemplateAllowlist(ctx context.Context, templateAllowlist string) error UpsertChatTitleGenerationModelOverride(ctx context.Context, value string) error diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index 214ed5833367b..a2ebf1b7bb157 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -24613,6 +24613,18 @@ func (q *sqlQuerier) GetChatRetentionDays(ctx context.Context) (int32, error) { return retention_days, err } +const getChatSummaryGenerationModelOverride = `-- name: GetChatSummaryGenerationModelOverride :one +SELECT + COALESCE((SELECT value FROM site_configs WHERE key = 'agents_chat_summary_generation_model_override'), '') :: text AS model_config_id +` + +func (q *sqlQuerier) GetChatSummaryGenerationModelOverride(ctx context.Context) (string, error) { + row := q.db.QueryRowContext(ctx, getChatSummaryGenerationModelOverride) + var model_config_id string + err := row.Scan(&model_config_id) + return model_config_id, err +} + const getChatSystemPrompt = `-- name: GetChatSystemPrompt :one SELECT COALESCE((SELECT value FROM site_configs WHERE key = 'agents_chat_system_prompt'), '') :: text AS chat_system_prompt @@ -25062,6 +25074,16 @@ func (q *sqlQuerier) UpsertChatRetentionDays(ctx context.Context, retentionDays return err } +const upsertChatSummaryGenerationModelOverride = `-- name: UpsertChatSummaryGenerationModelOverride :exec +INSERT INTO site_configs (key, value) VALUES ('agents_chat_summary_generation_model_override', $1) +ON CONFLICT (key) DO UPDATE SET value = $1 WHERE site_configs.key = 'agents_chat_summary_generation_model_override' +` + +func (q *sqlQuerier) UpsertChatSummaryGenerationModelOverride(ctx context.Context, value string) error { + _, err := q.db.ExecContext(ctx, upsertChatSummaryGenerationModelOverride, value) + return err +} + const upsertChatSystemPrompt = `-- name: UpsertChatSystemPrompt :exec INSERT INTO site_configs (key, value) VALUES ('agents_chat_system_prompt', $1) ON CONFLICT (key) DO UPDATE SET value = $1 WHERE site_configs.key = 'agents_chat_system_prompt' diff --git a/coderd/database/queries/siteconfig.sql b/coderd/database/queries/siteconfig.sql index 709cd287ca610..93877b9b51408 100644 --- a/coderd/database/queries/siteconfig.sql +++ b/coderd/database/queries/siteconfig.sql @@ -191,6 +191,14 @@ SELECT INSERT INTO site_configs (key, value) VALUES ('agents_chat_title_generation_model_override', $1) ON CONFLICT (key) DO UPDATE SET value = $1 WHERE site_configs.key = 'agents_chat_title_generation_model_override'; +-- name: GetChatSummaryGenerationModelOverride :one +SELECT + COALESCE((SELECT value FROM site_configs WHERE key = 'agents_chat_summary_generation_model_override'), '') :: text AS model_config_id; + +-- name: UpsertChatSummaryGenerationModelOverride :exec +INSERT INTO site_configs (key, value) VALUES ('agents_chat_summary_generation_model_override', $1) +ON CONFLICT (key) DO UPDATE SET value = $1 WHERE site_configs.key = 'agents_chat_summary_generation_model_override'; + -- name: GetChatDesktopEnabled :one SELECT COALESCE((SELECT value = 'true' FROM site_configs WHERE key = 'agents_desktop_enabled'), false) :: boolean AS enable_desktop; diff --git a/coderd/exp_chats.go b/coderd/exp_chats.go index 24d3e6796d7ea..e93d6b0ed98c7 100644 --- a/coderd/exp_chats.go +++ b/coderd/exp_chats.go @@ -621,6 +621,12 @@ func (api *API) chatModelOverrideSiteConfig( getter: api.Database.GetChatTitleGenerationModelOverride, upsert: api.Database.UpsertChatTitleGenerationModelOverride, }, nil + case codersdk.ChatModelOverrideContextSummaryGeneration: + return chatModelOverrideSiteConfig{ + label: "summary generation", + getter: api.Database.GetChatSummaryGenerationModelOverride, + upsert: api.Database.UpsertChatSummaryGenerationModelOverride, + }, nil default: return chatModelOverrideSiteConfig{}, xerrors.Errorf( "unknown chat model override context %q", diff --git a/coderd/exp_chats_test.go b/coderd/exp_chats_test.go index 7ae02bc837be8..2306e7b0147a0 100644 --- a/coderd/exp_chats_test.go +++ b/coderd/exp_chats_test.go @@ -11667,6 +11667,16 @@ func TestChatModelOverrides(t *testing.T) { return db.UpsertChatTitleGenerationModelOverride(dbauthz.AsSystemRestricted(ctx), value) }, }, + { + name: "SummaryGeneration", + context: codersdk.ChatModelOverrideContextSummaryGeneration, + dbGet: func(ctx context.Context, db database.Store) (string, error) { + return db.GetChatSummaryGenerationModelOverride(dbauthz.AsSystemRestricted(ctx)) + }, + dbUpsert: func(ctx context.Context, db database.Store, value string) error { + return db.UpsertChatSummaryGenerationModelOverride(dbauthz.AsSystemRestricted(ctx), value) + }, + }, } for _, setting := range settings { @@ -11812,7 +11822,7 @@ func TestChatModelOverrides(t *testing.T) { require.Equal(t, "Invalid chat model override context.", sdkErr.Message) require.Equal( t, - `Expected one of general, explore, title_generation. Got "not-a-context".`, + `Expected one of general, explore, title_generation, summary_generation. Got "not-a-context".`, sdkErr.Detail, ) @@ -11821,7 +11831,7 @@ func TestChatModelOverrides(t *testing.T) { require.Equal(t, "Invalid chat model override context.", sdkErr.Message) require.Equal( t, - `Expected one of general, explore, title_generation. Got "not-a-context".`, + `Expected one of general, explore, title_generation, summary_generation. Got "not-a-context".`, sdkErr.Detail, ) }) diff --git a/coderd/x/chatd/chatd.go b/coderd/x/chatd/chatd.go index c0d2c6f275dbc..5f6afdf8c1138 100644 --- a/coderd/x/chatd/chatd.go +++ b/coderd/x/chatd/chatd.go @@ -4983,19 +4983,44 @@ func (p *Server) generateAndStoreChatSummary( p.updateChatSummary(ctx, chat, chat.HistoryVersion, summary, logger) } +// resolveChatSummaryModel prefers a usable deployment override, skips generation +// on a set-but-unusable one, and otherwise uses the chat's configured model. func (p *Server) resolveChatSummaryModel( ctx context.Context, chat database.Chat, modelOpts modelBuildOptions, logger slog.Logger, ) (fantasy.LanguageModel, database.ChatModelConfig, bool) { - //nolint:dogsled // resolveChatModel returns rich routing metadata; summary generation only needs the model and its config. - model, dbConfig, _, _, _, _, _, err := p.resolveChatModel(ctx, chat, modelOpts) + //nolint:dogsled // resolveChatModel returns rich routing metadata; summary generation only needs the model, its config, and provider keys. + model, dbConfig, keys, _, _, _, _, err := p.resolveChatModel(ctx, chat, modelOpts) if err != nil { logger.Debug(ctx, "failed to resolve chat model for summary", slog.F("chat_id", chat.ID), slog.Error(err)) return nil, database.ChatModelConfig{}, false } + + overrideConfig, overrideModel, _, _, overrideSet, overrideErr := p.resolveSummaryGenerationModelOverride( + ctx, chat, keys, modelOpts, + ) + if overrideErr != nil { + if overrideSet { + logger.Warn(ctx, "summary generation model override unavailable, skipping summary generation", + slog.F("chat_id", chat.ID), + slog.F("override_context", summaryGenerationOverrideContext), + slog.Error(overrideErr), + ) + return nil, database.ChatModelConfig{}, false + } + logger.Debug(ctx, "failed to resolve summary generation model override", + slog.F("chat_id", chat.ID), + slog.F("override_context", summaryGenerationOverrideContext), + slog.Error(overrideErr), + ) + } + if overrideSet { + return overrideModel, overrideConfig, true + } + return model, dbConfig, true } diff --git a/coderd/x/chatd/summary_override.go b/coderd/x/chatd/summary_override.go new file mode 100644 index 0000000000000..91e35d5db5ed7 --- /dev/null +++ b/coderd/x/chatd/summary_override.go @@ -0,0 +1,45 @@ +package chatd + +import ( + "context" + + "charm.land/fantasy" + "golang.org/x/xerrors" + + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/database/dbauthz" + "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" +) + +const summaryGenerationOverrideContext = "summary_generation" + +func readSummaryGenerationModelOverride( + ctx context.Context, + db database.Store, +) (string, error) { + //nolint:gocritic // Chatd is internal, not a user, so this read uses AsChatd. + chatdCtx := dbauthz.AsChatd(ctx) + raw, err := db.GetChatSummaryGenerationModelOverride(chatdCtx) + if err != nil { + return "", xerrors.Errorf( + "get chat summary generation model override: %w", + err, + ) + } + return raw, nil +} + +// resolveSummaryGenerationModelOverride resolves the deployment-wide summary +// override. overrideSet reports whether one was configured; if true, any error is +// a hard failure (skip generation), and if false the caller uses the chat's model. +func (p *Server) resolveSummaryGenerationModelOverride( + ctx context.Context, + chat database.Chat, + keys chatprovider.ProviderAPIKeys, + modelOpts modelBuildOptions, +) (database.ChatModelConfig, fantasy.LanguageModel, chatprovider.ProviderAPIKeys, resolvedModelRoute, bool, error) { + return p.resolveGenerationModelOverride( + ctx, chat, keys, modelOpts, + summaryGenerationOverrideContext, readSummaryGenerationModelOverride, + ) +} diff --git a/coderd/x/chatd/summary_override_internal_test.go b/coderd/x/chatd/summary_override_internal_test.go new file mode 100644 index 0000000000000..deacc55bfaeea --- /dev/null +++ b/coderd/x/chatd/summary_override_internal_test.go @@ -0,0 +1,118 @@ +package chatd + +import ( + "testing" + + "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" + + "cdr.dev/slog/v3/sloggers/slogtest" + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/database/dbmock" + "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" + "github.com/coder/coder/v2/testutil" +) + +func TestResolveSummaryGenerationModelOverride_Unset(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}) + chat, _ := titleOverrideTestChatAndMessages(t) + + db.EXPECT().GetChatSummaryGenerationModelOverride(gomock.Any()).Return("", nil) + + server := titleOverrideTestServer(db, logger) + config, model, _, _, overrideSet, err := server.resolveSummaryGenerationModelOverride( + ctx, + chat, + chatprovider.ProviderAPIKeys{ByProvider: map[string]string{"openai": "test-key"}}, + modelBuildOptions{}, + ) + require.NoError(t, err) + require.False(t, overrideSet) + require.Nil(t, model) + require.Equal(t, database.ChatModelConfig{}, config) +} + +func TestResolveSummaryGenerationModelOverride_MalformedFallsThrough(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}) + chat, _ := titleOverrideTestChatAndMessages(t) + + db.EXPECT().GetChatSummaryGenerationModelOverride(gomock.Any()).Return("not-a-uuid", nil) + + server := titleOverrideTestServer(db, logger) + config, model, _, _, overrideSet, err := server.resolveSummaryGenerationModelOverride( + ctx, + chat, + chatprovider.ProviderAPIKeys{ByProvider: map[string]string{"openai": "test-key"}}, + modelBuildOptions{}, + ) + require.NoError(t, err) + require.False(t, overrideSet) + require.Nil(t, model) + require.Equal(t, database.ChatModelConfig{}, config) +} + +func TestResolveSummaryGenerationModelOverride_SetUsable(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}) + chat, _ := titleOverrideTestChatAndMessages(t) + overrideConfig := titleOverrideModelConfig("gpt-4.1", true) + + db.EXPECT().GetChatSummaryGenerationModelOverride(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() + + server := titleOverrideTestServer(db, logger) + config, model, _, _, overrideSet, err := server.resolveSummaryGenerationModelOverride( + ctx, + chat, + chatprovider.ProviderAPIKeys{ByProvider: map[string]string{"openai": "test-key"}}, + modelBuildOptions{}, + ) + require.NoError(t, err) + require.True(t, overrideSet) + require.NotNil(t, model) + require.Equal(t, overrideConfig, config) +} + +func TestResolveSummaryGenerationModelOverride_SetUnusableHardFails(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}) + chat, _ := titleOverrideTestChatAndMessages(t) + overrideConfig := titleOverrideModelConfig("gpt-4.1", false) + + db.EXPECT().GetChatSummaryGenerationModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil) + db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil) + + server := titleOverrideTestServer(db, logger) + config, model, _, _, overrideSet, err := server.resolveSummaryGenerationModelOverride( + ctx, + chat, + chatprovider.ProviderAPIKeys{ByProvider: map[string]string{"openai": "test-key"}}, + modelBuildOptions{}, + ) + // overrideSet stays true on a hard failure so the caller skips generation. + require.Error(t, err) + require.True(t, overrideSet) + require.ErrorContains(t, err, "summary generation model override is unavailable") + require.Nil(t, model) + require.Equal(t, database.ChatModelConfig{}, config) +} diff --git a/codersdk/chats.go b/codersdk/chats.go index 2eafc30928ad1..87215be39a024 100644 --- a/codersdk/chats.go +++ b/codersdk/chats.go @@ -738,9 +738,10 @@ type UpdateChatPlanModeInstructionsRequest struct { type ChatModelOverrideContext string const ( - ChatModelOverrideContextGeneral ChatModelOverrideContext = "general" - ChatModelOverrideContextExplore ChatModelOverrideContext = "explore" - ChatModelOverrideContextTitleGeneration ChatModelOverrideContext = "title_generation" + ChatModelOverrideContextGeneral ChatModelOverrideContext = "general" + ChatModelOverrideContextExplore ChatModelOverrideContext = "explore" + ChatModelOverrideContextTitleGeneration ChatModelOverrideContext = "title_generation" + ChatModelOverrideContextSummaryGeneration ChatModelOverrideContext = "summary_generation" ) // Valid reports whether the override context is one of the supported values. @@ -748,7 +749,8 @@ func (c ChatModelOverrideContext) Valid() bool { switch c { case ChatModelOverrideContextGeneral, ChatModelOverrideContextExplore, - ChatModelOverrideContextTitleGeneration: + ChatModelOverrideContextTitleGeneration, + ChatModelOverrideContextSummaryGeneration: return true default: return false @@ -761,6 +763,7 @@ func AllChatModelOverrideContexts() []ChatModelOverrideContext { ChatModelOverrideContextGeneral, ChatModelOverrideContextExplore, ChatModelOverrideContextTitleGeneration, + ChatModelOverrideContextSummaryGeneration, } } diff --git a/site/src/api/typesGenerated.ts b/site/src/api/typesGenerated.ts index 97a0b34affcd2..4f48e00b3159a 100644 --- a/site/src/api/typesGenerated.ts +++ b/site/src/api/typesGenerated.ts @@ -2621,11 +2621,13 @@ export interface ChatModelOpenRouterProviderOptions { export type ChatModelOverrideContext = | "explore" | "general" + | "summary_generation" | "title_generation"; export const ChatModelOverrideContexts: ChatModelOverrideContext[] = [ "explore", "general", + "summary_generation", "title_generation", ]; diff --git a/site/src/pages/AgentsPage/AgentSettingsAgentsPage.tsx b/site/src/pages/AgentsPage/AgentSettingsAgentsPage.tsx index 5f664afb15027..ffe681299d825 100644 --- a/site/src/pages/AgentsPage/AgentSettingsAgentsPage.tsx +++ b/site/src/pages/AgentsPage/AgentSettingsAgentsPage.tsx @@ -20,6 +20,8 @@ const generalOverrideContext: TypesGen.ChatModelOverrideContext = "general"; const exploreOverrideContext: TypesGen.ChatModelOverrideContext = "explore"; const titleGenerationOverrideContext: TypesGen.ChatModelOverrideContext = "title_generation"; +const summaryGenerationOverrideContext: TypesGen.ChatModelOverrideContext = + "summary_generation"; const chatModelOverrideKey = (context: TypesGen.ChatModelOverrideContext) => ["chat-model-override", context] as const; @@ -66,6 +68,10 @@ const AgentSettingsAgentsPage: FC = () => { ...chatModelOverrideQuery(titleGenerationOverrideContext), enabled: canEditDeploymentConfig, }); + const summaryGenerationModelQuery = useQuery({ + ...chatModelOverrideQuery(summaryGenerationOverrideContext), + enabled: canEditDeploymentConfig, + }); const modelConfigsQuery = useQuery(chatModelConfigs()); const savePersonalModelOverridesAdminSettingsMutation = useMutation( updateChatPersonalModelOverridesAdminSettings(queryClient), @@ -79,6 +85,12 @@ const AgentSettingsAgentsPage: FC = () => { titleGenerationOverrideContext, ), ); + const saveSummaryGenerationModelMutation = useMutation( + updateChatModelOverrideMutation( + queryClient, + summaryGenerationOverrideContext, + ), + ); const saveExploreModelOverrideMutation = useMutation( updateChatModelOverrideMutation(queryClient, exploreOverrideContext), ); @@ -123,6 +135,14 @@ const AgentSettingsAgentsPage: FC = () => { isSaveTitleGenerationModelError={ saveTitleGenerationModelMutation.isError } + summaryGenerationModelOverrideData={summaryGenerationModelQuery.data} + onSaveSummaryGenerationModel={saveSummaryGenerationModelMutation.mutate} + isSavingSummaryGenerationModel={ + saveSummaryGenerationModelMutation.isPending + } + isSaveSummaryGenerationModelError={ + saveSummaryGenerationModelMutation.isError + } onSaveExploreModelOverride={saveExploreModelOverrideMutation.mutate} isSavingExploreModelOverride={ saveExploreModelOverrideMutation.isPending diff --git a/site/src/pages/AgentsPage/AgentSettingsAgentsPageView.stories.tsx b/site/src/pages/AgentsPage/AgentSettingsAgentsPageView.stories.tsx index 03bd43ef8f8dd..cee22d7d15534 100644 --- a/site/src/pages/AgentsPage/AgentSettingsAgentsPageView.stories.tsx +++ b/site/src/pages/AgentsPage/AgentSettingsAgentsPageView.stories.tsx @@ -42,6 +42,11 @@ const buildTitleGenerationModelOverrideData = ( ): TypesGen.ChatModelOverrideResponse => buildOverrideData("title_generation", overrides); +const buildSummaryGenerationModelOverrideData = ( + overrides: Partial = {}, +): TypesGen.ChatModelOverrideResponse => + buildOverrideData("summary_generation", overrides); + const generalModelConfig = buildModelConfig({ id: "model-general-gpt-4.1-mini", display_name: "GPT 4.1 Mini", @@ -116,6 +121,7 @@ const buildArgs = ( isSaveAdminOverridesError: false, generalModelOverrideData: buildOverrideData("general"), titleGenerationModelOverrideData: buildTitleGenerationModelOverrideData(), + summaryGenerationModelOverrideData: buildSummaryGenerationModelOverrideData(), exploreModelOverrideData: buildOverrideData("explore"), modelConfigsData: allModelConfigs, modelConfigsError: undefined, @@ -126,6 +132,9 @@ const buildArgs = ( onSaveTitleGenerationModel: fn(), isSavingTitleGenerationModel: false, isSaveTitleGenerationModelError: false, + onSaveSummaryGenerationModel: fn(), + isSavingSummaryGenerationModel: false, + isSaveSummaryGenerationModelError: false, onSaveExploreModelOverride: fn(), isSavingExploreModelOverride: false, isSaveExploreModelOverrideError: false, @@ -181,6 +190,7 @@ export const AllOverridesUnset: Story = { "Enable users to define their personal overrides", "General model", "Title generation model", + "Summary generation model", "Explore subagent model", ]); await canvas.findByText( @@ -193,6 +203,10 @@ export const AllOverridesUnset: Story = { headingName: "Title generation model", placeholder: "Use title default", }, + { + headingName: "Summary generation model", + placeholder: "Use chat model", + }, { headingName: "Explore subagent model", placeholder: "Use chat default", @@ -350,6 +364,46 @@ export const EachOverrideSetToEnabledModel: Story = { }, }; +export const SummaryGenerationModelSetToEnabledModel: Story = { + args: buildArgs({ + summaryGenerationModelOverrideData: buildSummaryGenerationModelOverrideData( + { model_config_id: titleModelConfig.id }, + ), + }), + play: async ({ canvasElement, args }) => { + const summarySection = await getSection( + canvasElement, + "Summary generation model", + ); + + expect( + within(summarySection).getByRole("combobox", { + name: /gpt 4o mini/i, + }), + ).toHaveTextContent("GPT 4o Mini"); + + await selectModelInSection( + summarySection, + canvasElement, + /gpt 4o mini/i, + "Claude Sonnet 4", + ); + const summarySaveButton = within(summarySection).getByRole("button", { + name: "Save", + }); + await waitFor(() => { + expect(summarySaveButton).toBeEnabled(); + }); + await userEvent.click(summarySaveButton); + await waitFor(() => { + expect(args.onSaveSummaryGenerationModel).toHaveBeenCalledWith( + { model_config_id: claudeSonnetModelConfig.id }, + expect.anything(), + ); + }); + }, +}; + export const MalformedOverridesRemainClearableAndSaveable: Story = { args: buildArgs({ generalModelOverrideData: buildOverrideData("general", { diff --git a/site/src/pages/AgentsPage/AgentSettingsAgentsPageView.tsx b/site/src/pages/AgentsPage/AgentSettingsAgentsPageView.tsx index 46e2300975e8b..4a69927ce5f19 100644 --- a/site/src/pages/AgentsPage/AgentSettingsAgentsPageView.tsx +++ b/site/src/pages/AgentsPage/AgentSettingsAgentsPageView.tsx @@ -25,6 +25,7 @@ export interface AgentSettingsAgentsPageViewProps { isSaveAdminOverridesError: boolean; generalModelOverrideData?: TypesGen.ChatModelOverrideResponse; titleGenerationModelOverrideData?: TypesGen.ChatModelOverrideResponse; + summaryGenerationModelOverrideData?: TypesGen.ChatModelOverrideResponse; exploreModelOverrideData?: TypesGen.ChatModelOverrideResponse; modelConfigsData: TypesGen.ChatModelConfig[] | undefined; modelConfigsError: unknown; @@ -35,6 +36,9 @@ export interface AgentSettingsAgentsPageViewProps { onSaveTitleGenerationModel: SaveModelOverride; isSavingTitleGenerationModel: boolean; isSaveTitleGenerationModelError: boolean; + onSaveSummaryGenerationModel: SaveModelOverride; + isSavingSummaryGenerationModel: boolean; + isSaveSummaryGenerationModelError: boolean; onSaveExploreModelOverride: SaveModelOverride; isSavingExploreModelOverride: boolean; isSaveExploreModelOverrideError: boolean; @@ -52,6 +56,7 @@ export const AgentSettingsAgentsPageView: FC< isSaveAdminOverridesError, generalModelOverrideData, titleGenerationModelOverrideData, + summaryGenerationModelOverrideData, exploreModelOverrideData, modelConfigsData, modelConfigsError, @@ -62,6 +67,9 @@ export const AgentSettingsAgentsPageView: FC< onSaveTitleGenerationModel, isSavingTitleGenerationModel, isSaveTitleGenerationModelError, + onSaveSummaryGenerationModel, + isSavingSummaryGenerationModel, + isSaveSummaryGenerationModelError, onSaveExploreModelOverride, isSavingExploreModelOverride, isSaveExploreModelOverrideError, @@ -137,6 +145,31 @@ export const AgentSettingsAgentsPageView: FC< showHeader={false} /> +
+ + +
( onSaveTitleGenerationModel={fn()} isSavingTitleGenerationModel={false} isSaveTitleGenerationModelError={false} + summaryGenerationModelOverrideData={{ + context: "summary_generation", + model_config_id: "", + is_malformed: false, + }} + onSaveSummaryGenerationModel={fn()} + isSavingSummaryGenerationModel={false} + isSaveSummaryGenerationModelError={false} onSaveExploreModelOverride={fn()} isSavingExploreModelOverride={false} isSaveExploreModelOverrideError={false}