diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index 10f6bcd855a..0b1baaddec9 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -3036,6 +3036,13 @@ func (q *querier) GetChatByIDForUpdate(ctx context.Context, id uuid.UUID) (datab return fetch(q.log, q.auth, q.db.GetChatByIDForUpdate)(ctx, id) } +func (q *querier) GetChatCompactionModelOverride(ctx context.Context) (string, error) { + if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceDeploymentConfig); err != nil { + return "", err + } + return q.db.GetChatCompactionModelOverride(ctx) +} + func (q *querier) GetChatComputerUseProvider(ctx context.Context) (string, error) { // The computer-use provider is a deployment-wide runtime chat setting // read by authenticated chat users and chatd. Feature and experiment @@ -8707,6 +8714,13 @@ func (q *querier) UpsertChatAutoArchiveDays(ctx context.Context, autoArchiveDays return q.db.UpsertChatAutoArchiveDays(ctx, autoArchiveDays) } +func (q *querier) UpsertChatCompactionModelOverride(ctx context.Context, value string) error { + if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil { + return err + } + return q.db.UpsertChatCompactionModelOverride(ctx, value) +} + func (q *querier) UpsertChatComputerUseProvider(ctx context.Context, provider 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 67c7cee24c6..88ad172a452 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -1209,6 +1209,10 @@ func (s *MethodTestSuite) TestChats() { dbm.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return("", nil).AnyTimes() check.Args().Asserts(rbac.ResourceDeploymentConfig, policy.ActionRead) })) + s.Run("GetChatCompactionModelOverride", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + dbm.EXPECT().GetChatCompactionModelOverride(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) @@ -1640,6 +1644,10 @@ func (s *MethodTestSuite) TestChats() { dbm.EXPECT().UpsertChatTitleGenerationModelOverride(gomock.Any(), "").Return(nil).AnyTimes() check.Args("").Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate) })) + s.Run("UpsertChatCompactionModelOverride", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) { + dbm.EXPECT().UpsertChatCompactionModelOverride(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 a2592c270ec..19efe41292b 100644 --- a/coderd/database/dbmetrics/querymetrics.go +++ b/coderd/database/dbmetrics/querymetrics.go @@ -1425,6 +1425,14 @@ func (m queryMetricsStore) GetChatByIDForUpdate(ctx context.Context, id uuid.UUI return r0, r1 } +func (m queryMetricsStore) GetChatCompactionModelOverride(ctx context.Context) (string, error) { + start := time.Now() + r0, r1 := m.s.GetChatCompactionModelOverride(ctx) + m.queryLatencies.WithLabelValues("GetChatCompactionModelOverride").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetChatCompactionModelOverride").Inc() + return r0, r1 +} + func (m queryMetricsStore) GetChatComputerUseProvider(ctx context.Context) (string, error) { start := time.Now() r0, r1 := m.s.GetChatComputerUseProvider(ctx) @@ -6145,6 +6153,14 @@ func (m queryMetricsStore) UpsertChatAutoArchiveDays(ctx context.Context, autoAr return r0 } +func (m queryMetricsStore) UpsertChatCompactionModelOverride(ctx context.Context, value string) error { + start := time.Now() + r0 := m.s.UpsertChatCompactionModelOverride(ctx, value) + m.queryLatencies.WithLabelValues("UpsertChatCompactionModelOverride").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "UpsertChatCompactionModelOverride").Inc() + return r0 +} + func (m queryMetricsStore) UpsertChatComputerUseProvider(ctx context.Context, provider string) error { start := time.Now() r0 := m.s.UpsertChatComputerUseProvider(ctx, provider) diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index ba2884a0d44..649a4dfd256 100644 --- a/coderd/database/dbmock/dbmock.go +++ b/coderd/database/dbmock/dbmock.go @@ -2623,6 +2623,21 @@ func (mr *MockStoreMockRecorder) GetChatByIDForUpdate(ctx, id any) *gomock.Call return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatByIDForUpdate", reflect.TypeOf((*MockStore)(nil).GetChatByIDForUpdate), ctx, id) } +// GetChatCompactionModelOverride mocks base method. +func (m *MockStore) GetChatCompactionModelOverride(ctx context.Context) (string, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetChatCompactionModelOverride", ctx) + ret0, _ := ret[0].(string) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetChatCompactionModelOverride indicates an expected call of GetChatCompactionModelOverride. +func (mr *MockStoreMockRecorder) GetChatCompactionModelOverride(ctx any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatCompactionModelOverride", reflect.TypeOf((*MockStore)(nil).GetChatCompactionModelOverride), ctx) +} + // GetChatComputerUseProvider mocks base method. func (m *MockStore) GetChatComputerUseProvider(ctx context.Context) (string, error) { m.ctrl.T.Helper() @@ -11514,6 +11529,20 @@ func (mr *MockStoreMockRecorder) UpsertChatAutoArchiveDays(ctx, autoArchiveDays return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpsertChatAutoArchiveDays", reflect.TypeOf((*MockStore)(nil).UpsertChatAutoArchiveDays), ctx, autoArchiveDays) } +// UpsertChatCompactionModelOverride mocks base method. +func (m *MockStore) UpsertChatCompactionModelOverride(ctx context.Context, value string) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "UpsertChatCompactionModelOverride", ctx, value) + ret0, _ := ret[0].(error) + return ret0 +} + +// UpsertChatCompactionModelOverride indicates an expected call of UpsertChatCompactionModelOverride. +func (mr *MockStoreMockRecorder) UpsertChatCompactionModelOverride(ctx, value any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpsertChatCompactionModelOverride", reflect.TypeOf((*MockStore)(nil).UpsertChatCompactionModelOverride), ctx, value) +} + // UpsertChatComputerUseProvider mocks base method. func (m *MockStore) UpsertChatComputerUseProvider(ctx context.Context, provider string) error { m.ctrl.T.Helper() diff --git a/coderd/database/querier.go b/coderd/database/querier.go index 218688f7b40..d8a170a1186 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -376,6 +376,7 @@ type sqlcQuerier interface { GetChatByID(ctx context.Context, id uuid.UUID) (Chat, error) GetChatByIDForShare(ctx context.Context, id uuid.UUID) (Chat, error) GetChatByIDForUpdate(ctx context.Context, id uuid.UUID) (Chat, error) + GetChatCompactionModelOverride(ctx context.Context) (string, error) GetChatComputerUseProvider(ctx context.Context) (string, error) // Per-root-chat cost breakdown for a single user within a date range. // Groups by root_chat_id so forked chats roll up under their root. @@ -1529,6 +1530,7 @@ type sqlcQuerier interface { // to JSON before invoking this query. UpsertChatAdvisorConfig(ctx context.Context, value string) error UpsertChatAutoArchiveDays(ctx context.Context, autoArchiveDays int32) error + UpsertChatCompactionModelOverride(ctx context.Context, value string) error UpsertChatComputerUseProvider(ctx context.Context, provider string) error // UpsertChatDebugLoggingAllowUsers updates the runtime admin setting that // allows users to opt into chat debug logging. diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index cb2663638c7..f511c4663ac 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -24261,6 +24261,18 @@ func (q *sqlQuerier) GetChatAutoArchiveDays(ctx context.Context, defaultAutoArch return auto_archive_days, err } +const getChatCompactionModelOverride = `-- name: GetChatCompactionModelOverride :one +SELECT + COALESCE((SELECT value FROM site_configs WHERE key = 'agents_chat_compaction_model_override'), '') :: text AS model_config_id +` + +func (q *sqlQuerier) GetChatCompactionModelOverride(ctx context.Context) (string, error) { + row := q.db.QueryRowContext(ctx, getChatCompactionModelOverride) + var model_config_id string + err := row.Scan(&model_config_id) + return model_config_id, err +} + const getChatComputerUseProvider = `-- name: GetChatComputerUseProvider :one SELECT COALESCE((SELECT value FROM site_configs WHERE key = 'agents_computer_use_provider'), '') :: text AS provider @@ -24702,6 +24714,16 @@ func (q *sqlQuerier) UpsertChatAutoArchiveDays(ctx context.Context, autoArchiveD return err } +const upsertChatCompactionModelOverride = `-- name: UpsertChatCompactionModelOverride :exec +INSERT INTO site_configs (key, value) VALUES ('agents_chat_compaction_model_override', $1) +ON CONFLICT (key) DO UPDATE SET value = $1 WHERE site_configs.key = 'agents_chat_compaction_model_override' +` + +func (q *sqlQuerier) UpsertChatCompactionModelOverride(ctx context.Context, value string) error { + _, err := q.db.ExecContext(ctx, upsertChatCompactionModelOverride, value) + return err +} + const upsertChatComputerUseProvider = `-- name: UpsertChatComputerUseProvider :exec INSERT INTO site_configs (key, value) VALUES ('agents_computer_use_provider', $1) ON CONFLICT (key) DO UPDATE SET value = $1 WHERE site_configs.key = 'agents_computer_use_provider' diff --git a/coderd/database/queries/siteconfig.sql b/coderd/database/queries/siteconfig.sql index 709cd287ca6..3eb3aacaf02 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: GetChatCompactionModelOverride :one +SELECT + COALESCE((SELECT value FROM site_configs WHERE key = 'agents_chat_compaction_model_override'), '') :: text AS model_config_id; + +-- name: UpsertChatCompactionModelOverride :exec +INSERT INTO site_configs (key, value) VALUES ('agents_chat_compaction_model_override', $1) +ON CONFLICT (key) DO UPDATE SET value = $1 WHERE site_configs.key = 'agents_chat_compaction_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 fa4407f880e..2f0d94939e9 100644 --- a/coderd/exp_chats.go +++ b/coderd/exp_chats.go @@ -709,6 +709,12 @@ func (api *API) chatModelOverrideSiteConfig( getter: api.Database.GetChatTitleGenerationModelOverride, upsert: api.Database.UpsertChatTitleGenerationModelOverride, }, nil + case codersdk.ChatModelOverrideContextCompaction: + return chatModelOverrideSiteConfig{ + label: "compaction", + getter: api.Database.GetChatCompactionModelOverride, + upsert: api.Database.UpsertChatCompactionModelOverride, + }, 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 4b93cfa1b8a..7bc0917400b 100644 --- a/coderd/exp_chats_test.go +++ b/coderd/exp_chats_test.go @@ -12319,6 +12319,16 @@ func TestChatModelOverrides(t *testing.T) { return db.UpsertChatTitleGenerationModelOverride(dbauthz.AsSystemRestricted(ctx), value) }, }, + { + name: "Compaction", + context: codersdk.ChatModelOverrideContextCompaction, + dbGet: func(ctx context.Context, db database.Store) (string, error) { + return db.GetChatCompactionModelOverride(dbauthz.AsSystemRestricted(ctx)) + }, + dbUpsert: func(ctx context.Context, db database.Store, value string) error { + return db.UpsertChatCompactionModelOverride(dbauthz.AsSystemRestricted(ctx), value) + }, + }, } for _, setting := range settings { @@ -12528,7 +12538,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, compaction. Got "not-a-context".`, sdkErr.Detail, ) @@ -12537,7 +12547,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, compaction. Got "not-a-context".`, sdkErr.Detail, ) }) diff --git a/coderd/x/chatd/ARCHITECTURE.md b/coderd/x/chatd/ARCHITECTURE.md index d1519231ff6..9cea5ce8475 100644 --- a/coderd/x/chatd/ARCHITECTURE.md +++ b/coderd/x/chatd/ARCHITECTURE.md @@ -831,6 +831,19 @@ Model configs may carry a `reasoning_effort` config (`{default, max}`) inside `c During generation preparation, the effective effort is resolved as the chat's `last_reasoning_effort` if set, else the config's `default`; clamped to the config's `max` on the global scale `none < minimal < low < medium < high < xhigh < max`; and passed through to the provider. The provider verifies whether the configured value is valid for that model at runtime. If the model config has no `reasoning_effort`, any user-selected value is ignored. The resolved value is injected into the provider-native options with `chatprovider.ApplyReasoningEffort` after provider option conversion. +#### Compaction model selection + +Compaction is an auxiliary LLM call: when the conversation approaches the context limit, the generation goroutine asks a model to summarize the history, commits the summary as a compressed boundary, and continues the turn on the chat model. + +By default the summary is generated with the chat model. Admins can override the compaction model deployment-wide via the `compaction` context of the chat model override API (`/api/experimental/chats/config/model-override/{context}`, stored in the `agents_chat_compaction_model_override` site config). The override affects only the summary call; thresholds, compressed-message storage, and the post-compaction assistant generation keep using the chat model. + +Details that follow from the override: + +- Context limits: the compaction trigger uses the stricter of the chat model's and the compaction model's context limits, because the history must also fit the summarizer's window. The post-compaction "still over limit" check stays against the chat model's limit, since continuation runs on the chat model. +- Failure semantics: an unset override uses the chat model. Stale or malformed stored references (deleted or disabled config or provider, missing credentials, non-UUID value) fall back to the chat model with a log. A usable override that fails at use (route or client construction, provider call failure) fails the generation visibly through the normal error path; there is no silent fallback. The override model client is constructed inside the compact generation action, not at prepare time, so a broken override cannot fail turns that finish without compacting (including turns over the threshold whose last assistant step already completed). +- Prompt safety: the prompt is built and sanitized for the chat model, so when the override points at a different provider the compaction copy of the prompt is re-sanitized: provider-executed tool history is flattened into plain text parts (keeping its content while dropping the provider-specific wire shape), file parts the compaction model rejects are replaced with text placeholders, and Anthropic provider-tool sanitization is re-run for the compaction provider. The assistant generation prompt is never mutated. +- Observability: compaction metrics and chat debug runs record the provider and model that actually generated the summary. This includes the "still over limit" terminal error, which is recorded before the override client is built: prepare-time resolution keeps the override's provider/model identity so that error lands on the same metric series as the compact action's own events. + #### Interrupt goroutine The interrupt goroutine is responsible for handling interrupts. It is spawned when the event indicates the core state machine is in `I0` or `I1` (status is `interrupting`). diff --git a/coderd/x/chatd/chatd_test.go b/coderd/x/chatd/chatd_test.go index 300ed4c647d..b7c7de1f96b 100644 --- a/coderd/x/chatd/chatd_test.go +++ b/coderd/x/chatd/chatd_test.go @@ -51,6 +51,7 @@ import ( "github.com/coder/coder/v2/coderd/workspacestats" "github.com/coder/coder/v2/coderd/x/chatd" "github.com/coder/coder/v2/coderd/x/chatd/chatadvisor" + "github.com/coder/coder/v2/coderd/x/chatd/chatdebug" "github.com/coder/coder/v2/coderd/x/chatd/chatprompt" "github.com/coder/coder/v2/coderd/x/chatd/chatsanitize" "github.com/coder/coder/v2/coderd/x/chatd/chatstate" @@ -5861,6 +5862,218 @@ func singlePartOfType(t *testing.T, msg database.ChatMessage, typ codersdk.ChatM return matches[0] } +func TestActiveServer_CompactionModelOverride(t *testing.T) { + t.Parallel() + + const ( + compactionSummary = "summary text for compaction" + chatModelName = "claude-sonnet-4-20250514" + overrideModelName = "claude-3-5-haiku-latest" + thresholdPercent = int32(70) + ) + + seedOverrideModel := func(ctx context.Context, t *testing.T, db database.Store, chatModel database.ChatModelConfig, contextLimit int64) database.ChatModelConfig { + t.Helper() + overrideModel := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{ + Model: overrideModelName, + AIProviderID: chatModel.AIProviderID, + ContextLimit: contextLimit, + }) + lowEffort := "low" + overrideModel = updateChatModelCallConfig(t, db, overrideModel, codersdk.ChatModelCallConfig{ + ReasoningEffort: &codersdk.ChatModelReasoningEffortConfig{ + Default: &lowEffort, + Max: &lowEffort, + }, + }) + require.NoError(t, db.UpsertChatCompactionModelOverride(ctx, overrideModel.ID.String())) + return overrideModel + } + + t.Run("summary routes to the override model and continuation stays on the chat model", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + db, ps := dbtestutil.NewDB(t) + reg := prometheus.NewRegistry() + var streamCount atomic.Int32 + anthropicURL := chattest.NewAnthropic(t, func(req *chattest.AnthropicRequest) chattest.AnthropicResponse { + body := anthropicRequestBody(t, *req) + if !req.Stream { + if strings.Contains(body, "You are performing a context compaction") { + require.Equal(t, overrideModelName, req.Model) + // The override config's reasoning effort must reach the + // summary request (Anthropic serializes it as + // output_config effort). + require.Contains(t, string(req.OutputConfig), `"effort":"low"`) + return anthropicCompactionResponse(compactionSummary) + } + return chattest.AnthropicNonStreamingResponse("title") + } + require.Equal(t, chatModelName, req.Model) + switch streamCount.Add(1) { + case 1: + return highUsageReadFileResponse("/tmp/a.txt") + default: + require.Contains(t, body, compactionSummary) + require.Empty(t, string(req.OutputConfig), + "the override reasoning effort must not leak into chat model generations") + return chattest.AnthropicStreamingResponse(chattest.AnthropicTextChunksWithCacheUsage(chattest.AnthropicUsage{ + InputTokens: 20, + OutputTokens: 5, + }, "continued after compaction")...) + } + }) + user, org, model := seedAnthropicChatDependencies(t, db, anthropicURL) + model = updateChatModelCompressionThreshold(t, db, model, 100, thresholdPercent) + overrideModel := seedOverrideModel(ctx, t, db, model, 1_000_000) + ws, dbAgent := seedWorkspaceWithAgent(t, db, user.ID) + + ctrl := gomock.NewController(t) + mockConn := agentconnmock.NewMockAgentConn(ctrl) + setupToolExecutionAgentConn(t, mockConn) + mockConn.EXPECT().ReadFileLines(gomock.Any(), "/tmp/a.txt", int64(1), int64(0), gomock.Any()). + Return(workspacesdk.ReadFileLinesResponse{Success: true, FileSize: 12, TotalLines: 1, LinesRead: 1, Content: "1\tpackage main"}, nil). + Times(1) + + server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) { + cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath())) + cfg.PrometheusRegistry = reg + cfg.AlwaysEnableDebugLogs = true + cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) { + require.Equal(t, dbAgent.ID, agentID) + return mockConn, func() {}, nil + } + }) + chat, err := server.CreateChat(ctx, chatd.CreateOptions{ + OrganizationID: org.ID, + OwnerID: user.ID, + APIKeyID: testAPIKeyID(t, db, user.ID), + WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, + AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, + Title: "compaction-override", + ModelConfigID: model.ID, + InitialUserContent: []codersdk.ChatMessagePart{ + codersdk.ChatMessageText("read the file and continue"), + }, + }) + require.NoError(t, err) + waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting) + + messages := chatMessages(ctx, t, db, chat.ID) + promptMessages, err := db.GetChatMessagesForPromptByChatID(ctx, chat.ID) + require.NoError(t, err) + compressed := compressedChatSummarizedMessages(t, append(promptMessages, messages...)) + require.Len(t, compressed.summaries, 1) + require.Contains(t, messageText(t, compressed.summaries[0]), compactionSummary) + requireTextPart(t, messages[len(messages)-1], "continued after compaction") + + requireChatdMetricCounter(t, reg, "coderd_chatd_compaction_total", 1, map[string]string{ + "provider": "anthropic", + "model": overrideModelName, + "result": "success", + }) + + require.NoError(t, server.Close()) + debugCtx := testutil.Context(t, testutil.WaitLong) + var compactionRun database.ChatDebugRun + testutil.Eventually(debugCtx, t, func(ctx context.Context) bool { + runs, err := db.GetChatDebugRunsByChatID(ctx, database.GetChatDebugRunsByChatIDParams{ + ChatID: chat.ID, + LimitVal: 100, + }) + if err != nil { + return false + } + for _, run := range runs { + if run.Kind == string(chatdebug.KindCompaction) { + compactionRun = run + return true + } + } + return false + }, testutil.IntervalMedium) + require.True(t, compactionRun.Provider.Valid) + require.Equal(t, "anthropic", compactionRun.Provider.String) + require.True(t, compactionRun.Model.Valid) + require.Equal(t, overrideModelName, compactionRun.Model.String) + require.True(t, compactionRun.ModelConfigID.Valid) + require.Equal(t, overrideModel.ID, compactionRun.ModelConfigID.UUID) + }) + + t.Run("compaction triggers at the stricter override context limit", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitLong) + db, ps := dbtestutil.NewDB(t) + var streamCount atomic.Int32 + anthropicURL := chattest.NewAnthropic(t, func(req *chattest.AnthropicRequest) chattest.AnthropicResponse { + body := anthropicRequestBody(t, *req) + if !req.Stream { + if strings.Contains(body, "You are performing a context compaction") { + require.Equal(t, overrideModelName, req.Model) + return anthropicCompactionResponse(compactionSummary) + } + return chattest.AnthropicNonStreamingResponse("title") + } + switch streamCount.Add(1) { + case 1: + return highUsageReadFileResponse("/tmp/a.txt") + default: + require.Contains(t, body, compactionSummary) + return chattest.AnthropicStreamingResponse(chattest.AnthropicTextChunksWithCacheUsage(chattest.AnthropicUsage{ + InputTokens: 20, + OutputTokens: 5, + }, "continued after compaction")...) + } + }) + user, org, model := seedAnthropicChatDependencies(t, db, anthropicURL) + // The chat model alone would not compact: 80 tokens of usage is + // 8% of its 1000-token limit. The override model's 100-token + // limit makes the effective threshold 70 tokens, so compaction + // must trigger. + model = updateChatModelCompressionThreshold(t, db, model, 1_000, thresholdPercent) + seedOverrideModel(ctx, t, db, model, 100) + ws, dbAgent := seedWorkspaceWithAgent(t, db, user.ID) + + ctrl := gomock.NewController(t) + mockConn := agentconnmock.NewMockAgentConn(ctrl) + setupToolExecutionAgentConn(t, mockConn) + mockConn.EXPECT().ReadFileLines(gomock.Any(), "/tmp/a.txt", int64(1), int64(0), gomock.Any()). + Return(workspacesdk.ReadFileLinesResponse{Success: true, FileSize: 12, TotalLines: 1, LinesRead: 1, Content: "1\tpackage main"}, nil). + Times(1) + + server := newActiveTestServer(t, db, ps, func(cfg *chatd.Config) { + cfg.AIBridgeTransportFactory = chatAIGatewayTransportFactoryPointer(chattest.NewMockAIBridgeTransport(t, anthropicURL, chattest.WithPreservePath())) + cfg.AgentConn = func(_ context.Context, agentID uuid.UUID) (workspacesdk.AgentConn, func(), error) { + require.Equal(t, dbAgent.ID, agentID) + return mockConn, func() {}, nil + } + }) + chat, err := server.CreateChat(ctx, chatd.CreateOptions{ + OrganizationID: org.ID, + OwnerID: user.ID, + APIKeyID: testAPIKeyID(t, db, user.ID), + WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, + AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, + Title: "compaction-override-limit", + ModelConfigID: model.ID, + InitialUserContent: []codersdk.ChatMessagePart{ + codersdk.ChatMessageText("read the file and continue"), + }, + }) + require.NoError(t, err) + waitForChatStatus(ctx, t, db, chat.ID, database.ChatStatusWaiting) + + messages := chatMessages(ctx, t, db, chat.ID) + promptMessages, err := db.GetChatMessagesForPromptByChatID(ctx, chat.ID) + require.NoError(t, err) + compressed := compressedChatSummarizedMessages(t, append(promptMessages, messages...)) + require.Len(t, compressed.summaries, 1) + requireTextPart(t, messages[len(messages)-1], "continued after compaction") + }) +} + func TestActiveServer_BasicAssistantGenerationAndPromptPreparation(t *testing.T) { t.Parallel() diff --git a/coderd/x/chatd/chatloop/chatloop.go b/coderd/x/chatd/chatloop/chatloop.go index 33505aa593b..d1166b0fc96 100644 --- a/coderd/x/chatd/chatloop/chatloop.go +++ b/coderd/x/chatd/chatloop/chatloop.go @@ -280,6 +280,17 @@ type GenerateCompactionOptions struct { ToolCallID string ToolName string + // ResolvedProvider, ResolvedModel, and ModelConfigID identify the + // summary model, which can differ from the chat model when a + // compaction override is configured. Debug runs record these. + ResolvedProvider string + ResolvedModel string + ModelConfigID uuid.UUID + + // ProviderOptions carry summary-model call options such as an + // override's reasoning effort. + ProviderOptions fantasy.ProviderOptions + PublishMessagePart func(codersdk.ChatMessageRole, codersdk.ChatMessagePart) } diff --git a/coderd/x/chatd/chatloop/compaction.go b/coderd/x/chatd/chatloop/compaction.go index 946023edb28..5ba9918c44d 100644 --- a/coderd/x/chatd/chatloop/compaction.go +++ b/coderd/x/chatd/chatloop/compaction.go @@ -64,6 +64,13 @@ type CompactionOptions struct { ChatID uuid.UUID HistoryTipMessageID int64 + // Summary model identity and call options; see + // GenerateCompactionOptions. + ResolvedProvider string + ResolvedModel string + ModelConfigID uuid.UUID + ProviderOptions fantasy.ProviderOptions + // ToolCallID and ToolName identify the synthetic tool call // used to represent compaction in the message stream. ToolCallID string @@ -169,6 +176,10 @@ func normalizedCompactionGenerateConfig(opts GenerateCompactionOptions) (Compact DebugSvc: opts.DebugSvc, ChatID: opts.ChatID, HistoryTipMessageID: opts.HistoryTipMessageID, + ResolvedProvider: opts.ResolvedProvider, + ResolvedModel: opts.ResolvedModel, + ModelConfigID: opts.ModelConfigID, + ProviderOptions: opts.ProviderOptions, ToolCallID: opts.ToolCallID, ToolName: opts.ToolName, PublishMessagePart: opts.PublishMessagePart, @@ -276,6 +287,21 @@ func startCompactionDebugRun( historyTipMessageID = parentRun.HistoryTipMessageID } + // Prefer the caller-supplied summary model identity; it can differ + // from the parent run's chat model under a compaction override. + provider := parentRun.Provider + if options.ResolvedProvider != "" { + provider = options.ResolvedProvider + } + model := parentRun.Model + if options.ResolvedModel != "" { + model = options.ResolvedModel + } + modelConfigID := parentRun.ModelConfigID + if options.ModelConfigID != uuid.Nil { + modelConfigID = options.ModelConfigID + } + // Use a separate short-lived context for the debug insert so a // slow or locked DB cannot block the model call. Detached from // the parent so cancellation of the compaction run still lets @@ -288,13 +314,13 @@ func startCompactionDebugRun( ChatID: options.ChatID, RootChatID: parentRun.RootChatID, ParentChatID: parentRun.ParentChatID, - ModelConfigID: parentRun.ModelConfigID, + ModelConfigID: modelConfigID, TriggerMessageID: parentRun.TriggerMessageID, HistoryTipMessageID: historyTipMessageID, Kind: chatdebug.KindCompaction, Status: chatdebug.StatusInProgress, - Provider: parentRun.Provider, - Model: parentRun.Model, + Provider: provider, + Model: model, }) createRunCancel() if err != nil { @@ -307,12 +333,12 @@ func startCompactionDebugRun( ChatID: options.ChatID, RootChatID: parentRun.RootChatID, ParentChatID: parentRun.ParentChatID, - ModelConfigID: parentRun.ModelConfigID, + ModelConfigID: modelConfigID, TriggerMessageID: parentRun.TriggerMessageID, HistoryTipMessageID: historyTipMessageID, Kind: chatdebug.KindCompaction, - Provider: parentRun.Provider, - Model: parentRun.Model, + Provider: provider, + Model: model, }) return compactionCtx, func(runErr error) { @@ -367,8 +393,9 @@ func generateCompactionSummary( }() response, err := model.Generate(summaryCtx, fantasy.Call{ - Prompt: summaryPrompt, - ToolChoice: &toolChoice, + Prompt: summaryPrompt, + ToolChoice: &toolChoice, + ProviderOptions: options.ProviderOptions, }) if err != nil { return "", xerrors.Errorf("generate summary text: %w", err) diff --git a/coderd/x/chatd/chattest/anthropic.go b/coderd/x/chatd/chattest/anthropic.go index ba23571b8db..c88352b8ba8 100644 --- a/coderd/x/chatd/chattest/anthropic.go +++ b/coderd/x/chatd/chattest/anthropic.go @@ -31,6 +31,7 @@ type AnthropicRequest struct { Tools []AnthropicRequestTool `json:"tools,omitempty"` Stream bool `json:"stream,omitempty"` MaxTokens int `json:"max_tokens,omitempty"` + OutputConfig json.RawMessage `json:"output_config,omitempty"` // TODO: encoding/json ignores inline tags. Add custom UnmarshalJSON to capture unknown keys. Options map[string]interface{} `json:",inline"` //nolint:revive } diff --git a/coderd/x/chatd/compaction_override.go b/coderd/x/chatd/compaction_override.go new file mode 100644 index 00000000000..0e3d4d48604 --- /dev/null +++ b/coderd/x/chatd/compaction_override.go @@ -0,0 +1,188 @@ +package chatd + +import ( + "context" + "encoding/json" + + "charm.land/fantasy" + "github.com/google/uuid" + "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" + "github.com/coder/coder/v2/codersdk" +) + +const compactionOverrideContext = "compaction" + +func readCompactionModelOverride( + 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.GetChatCompactionModelOverride(chatdCtx) + if err != nil { + return "", xerrors.Errorf( + "get chat compaction model override: %w", + err, + ) + } + return raw, nil +} + +// compactionModelOverride carries the built compaction override model plus +// the identity metadata debug runs and prompt sanitization need. +type compactionModelOverride struct { + modelConfig database.ChatModelConfig + model fantasy.LanguageModel + resolvedProvider string + resolvedModel string + // providerOptions include the override's reasoning effort for the + // summary call. + providerOptions fantasy.ProviderOptions +} + +// resolvedCompactionOverride is the compaction override resolved at +// prepare time. The provider/model identity is resolved without building +// the model client so metrics recorded before the client exists +// (still-over-limit) attribute to the same model as the compact action's. +type resolvedCompactionOverride struct { + Config database.ChatModelConfig + // ResolvedProvider and ResolvedModel match the built client's + // identity: ResolveModelWithProviderHint normalizes its hint, so the + // normalized provider name here and the route's raw provider type in + // buildCompactionOverrideModel yield the same result. + ResolvedProvider string + ResolvedModel string +} + +// resolveCompactionOverrideConfig resolves the stored deployment-wide +// compaction model override. Unset, malformed, stale, and credential-less +// overrides fall back to the chat model (nil override). This runs on every +// generation prepare because the override's context limit feeds the +// compaction trigger; the model client is built only when compaction runs. +func (p *Server) resolveCompactionOverrideConfig( + ctx context.Context, + chat database.Chat, +) (*resolvedCompactionOverride, error) { + raw, err := readCompactionModelOverride(ctx, p.db) + if err != nil { + return nil, xerrors.Errorf( + "read compaction model override: %w", + err, + ) + } + + modelConfig, providerName, overrideEffort, overrideSet, err := p.resolveConfiguredModelOverride( + ctx, + compactionOverrideContext, + raw, + chat.OwnerID, + p.resolveModelConfigAndNormalizedProvider, + func(ctx context.Context, ownerID uuid.UUID, aiProviderID uuid.UUID) (chatprovider.ProviderAPIKeys, error) { + return p.resolveUserProviderAPIKeys(ctx, ownerID, aiProviderID) + }, + modelOverrideFailureModeSoft, + ) + if err != nil || !overrideSet { + return nil, err + } + // Already validated by the shared resolver; failure is unreachable. + resolvedProvider, resolvedModel, err := chatprovider.ResolveModelWithProviderHint( + modelConfig.Model, + providerName, + ) + if err != nil { + return nil, xerrors.Errorf( + "resolve compaction model override identity: %w", + err, + ) + } + return &resolvedCompactionOverride{ + Config: withResolvedReasoningEffort(modelConfig, overrideEffort), + ResolvedProvider: resolvedProvider, + ResolvedModel: resolvedModel, + }, nil +} + +// buildCompactionOverrideModel resolves the route and constructs the model +// client for a usable override config. Errors are hard failures: a usable +// override that cannot be constructed must fail the generation visibly +// instead of silently compacting with the chat model. +func (p *Server) buildCompactionOverrideModel( + ctx context.Context, + chat database.Chat, + modelConfig database.ChatModelConfig, + modelOpts modelBuildOptions, +) (compactionModelOverride, error) { + //nolint:gocritic // Compaction overrides need chatd-scoped provider reads for user-owned chats. + route, err := p.resolveModelRouteForConfig(dbauthz.AsChatd(ctx), chat.OwnerID, modelConfig) + if err != nil { + return compactionModelOverride{}, xerrors.Errorf( + "resolve compaction model override route: %w", + err, + ) + } + resolvedProvider, resolvedModel, err := chatprovider.ResolveModelWithProviderHint( + modelConfig.Model, + route.ModelProviderHint, + ) + if err != nil { + return compactionModelOverride{}, xerrors.Errorf( + "resolve compaction model override metadata: %w", + err, + ) + } + model, _, err := p.newDebugAwareModel(ctx, modelClientRequest{ + Chat: chat, + ModelName: modelConfig.Model, + UserAgent: chatprovider.UserAgent(), + ExtraHeaders: chatprovider.CoderHeaders(chat), + }, route, modelOpts) + if err != nil { + return compactionModelOverride{}, xerrors.Errorf( + "create compaction model override: %w", + err, + ) + } + providerOptions, err := compactionOverrideProviderOptions(model, modelConfig) + if err != nil { + return compactionModelOverride{}, err + } + return compactionModelOverride{ + modelConfig: modelConfig, + model: model, + resolvedProvider: resolvedProvider, + resolvedModel: resolvedModel, + providerOptions: providerOptions, + }, nil +} + +// compactionOverrideProviderOptions converts the override config's call +// options, including the admin-resolved reasoning effort, into provider +// options for the summary call. +func compactionOverrideProviderOptions( + model fantasy.LanguageModel, + modelConfig database.ChatModelConfig, +) (fantasy.ProviderOptions, error) { + callConfig := codersdk.ChatModelCallConfig{} + if len(modelConfig.Options) > 0 { + if err := json.Unmarshal(modelConfig.Options, &callConfig); err != nil { + return nil, xerrors.Errorf( + "parse compaction model override call config: %w", + err, + ) + } + } + providerOptions := chatprovider.ProviderOptionsFromChatModelConfig( + model, + callConfig.ProviderOptions, + ) + reasoningEffort := chatprovider.ResolveReasoningEffort( + nil, + callConfig.ReasoningEffort, + ) + return chatprovider.ApplyReasoningEffort(model, providerOptions, reasoningEffort), nil +} diff --git a/coderd/x/chatd/compaction_override_internal_test.go b/coderd/x/chatd/compaction_override_internal_test.go new file mode 100644 index 00000000000..9267c84446d --- /dev/null +++ b/coderd/x/chatd/compaction_override_internal_test.go @@ -0,0 +1,216 @@ +package chatd + +import ( + "database/sql" + "encoding/json" + "testing" + + fantasyanthropic "charm.land/fantasy/providers/anthropic" + "github.com/google/uuid" + "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/chattest" + "github.com/coder/coder/v2/codersdk" + "github.com/coder/coder/v2/testutil" +) + +func TestCompactionOverrideProviderOptions(t *testing.T) { + t.Parallel() + + model := &chattest.FakeModel{ProviderName: "anthropic", ModelName: "claude-3-5-haiku"} + + t.Run("NoOptions", func(t *testing.T) { + t.Parallel() + opts, err := compactionOverrideProviderOptions(model, database.ChatModelConfig{}) + require.NoError(t, err) + require.Nil(t, opts) + }) + + t.Run("ReasoningEffort", func(t *testing.T) { + t.Parallel() + effort := "low" + options, err := json.Marshal(codersdk.ChatModelCallConfig{ + ReasoningEffort: &codersdk.ChatModelReasoningEffortConfig{ + Default: &effort, + Max: &effort, + }, + }) + require.NoError(t, err) + opts, err := compactionOverrideProviderOptions(model, database.ChatModelConfig{Options: options}) + require.NoError(t, err) + anthropicOpts, ok := opts[fantasyanthropic.Name].(*fantasyanthropic.ProviderOptions) + require.True(t, ok) + require.NotNil(t, anthropicOpts.Effort) + require.Equal(t, fantasyanthropic.Effort("low"), *anthropicOpts.Effort) + }) + + t.Run("MalformedOptions", func(t *testing.T) { + t.Parallel() + _, err := compactionOverrideProviderOptions(model, database.ChatModelConfig{Options: []byte("{")}) + require.ErrorContains(t, err, "parse compaction model override call config") + }) +} + +func TestResolveCompactionOverrideConfig_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().GetChatCompactionModelOverride(gomock.Any()).Return("", nil) + + server := titleOverrideTestServer(db, logger) + override, err := server.resolveCompactionOverrideConfig(ctx, chat) + require.NoError(t, err) + require.Nil(t, override) +} + +func TestResolveCompactionOverrideConfig_ReadDBError(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().GetChatCompactionModelOverride(gomock.Any()).Return("", sql.ErrConnDone) + + server := titleOverrideTestServer(db, logger) + override, err := server.resolveCompactionOverrideConfig(ctx, chat) + require.Error(t, err) + require.ErrorContains(t, err, "read compaction model override") + require.Nil(t, override) +} + +func TestResolveCompactionOverrideConfig_MalformedFallsBack(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().GetChatCompactionModelOverride(gomock.Any()).Return("not-a-uuid", nil) + + server := titleOverrideTestServer(db, logger) + override, err := server.resolveCompactionOverrideConfig(ctx, chat) + require.NoError(t, err) + require.Nil(t, override) +} + +func TestResolveCompactionOverrideConfig_DeletedConfigFallsBack(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) + missingID := uuid.New() + + db.EXPECT().GetChatCompactionModelOverride(gomock.Any()).Return(missingID.String(), nil) + db.EXPECT().GetChatModelConfigByID(gomock.Any(), missingID).Return(database.ChatModelConfig{}, sql.ErrNoRows) + + server := titleOverrideTestServer(db, logger) + override, err := server.resolveCompactionOverrideConfig(ctx, chat) + require.NoError(t, err) + require.Nil(t, override) +} + +func TestResolveCompactionOverrideConfig_DisabledConfigFallsBack(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().GetChatCompactionModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil) + db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil) + + server := titleOverrideTestServer(db, logger) + override, err := server.resolveCompactionOverrideConfig(ctx, chat) + require.NoError(t, err) + require.Nil(t, override) +} + +func TestResolveCompactionOverrideConfig_MissingCredentialsFallsBack(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) + providerID := uuid.New() + overrideConfig.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true} + + db.EXPECT().GetChatCompactionModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil) + db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil) + db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(database.AIProvider{ + ID: providerID, + Type: database.AIProviderTypeOpenai, + Enabled: true, + }, nil).AnyTimes() + db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return(nil, nil).AnyTimes() + + server := titleOverrideTestServer(db, logger) + override, err := server.resolveCompactionOverrideConfig(ctx, chat) + require.NoError(t, err) + require.Nil(t, override) +} + +func TestCompactionOverride_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) + providerID := uuid.New() + overrideConfig.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true} + + db.EXPECT().GetChatCompactionModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil) + db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil) + db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(aibridgeTestAIProvider(providerID, "primary-openai", database.AIProviderTypeOpenai), nil).AnyTimes() + db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ + ProviderID: providerID, + APIKey: "test-key", + }}, nil).AnyTimes() + + server := titleOverrideTestServer(db, logger) + resolved, err := server.resolveCompactionOverrideConfig(ctx, chat) + require.NoError(t, err) + require.NotNil(t, resolved) + require.Equal(t, overrideConfig.ID, resolved.Config.ID) + + override, err := server.buildCompactionOverrideModel( + ctx, + chat, + resolved.Config, + modelBuildOptions{ActiveAPIKeyID: uuid.NewString()}, + ) + require.NoError(t, err) + require.NotNil(t, override.model) + require.Equal(t, overrideConfig.ID, override.modelConfig.ID) + require.Equal(t, "openai", override.resolvedProvider) + require.Equal(t, "gpt-4.1", override.resolvedModel) + // Prepare-time identity must match the built client's so + // still-over-limit metrics land on the same series. + require.Equal(t, override.resolvedProvider, resolved.ResolvedProvider) + require.Equal(t, override.resolvedModel, resolved.ResolvedModel) +} diff --git a/coderd/x/chatd/compaction_sanitize.go b/coderd/x/chatd/compaction_sanitize.go new file mode 100644 index 00000000000..9f4676b29f5 --- /dev/null +++ b/coderd/x/chatd/compaction_sanitize.go @@ -0,0 +1,187 @@ +package chatd + +import ( + "context" + "fmt" + + "charm.land/fantasy" + + "cdr.dev/slog/v3" + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" + "github.com/coder/coder/v2/coderd/x/chatd/chatsanitize" +) + +// sameCompactionProviderIdentity reports whether the chat and compaction +// override models share a provider instance. Configs without an +// AIProviderID compare as different (fail closed). +func sameCompactionProviderIdentity(chatConfig, overrideConfig database.ChatModelConfig) bool { + return chatConfig.AIProviderID.Valid && overrideConfig.AIProviderID.Valid && + chatConfig.AIProviderID.UUID == overrideConfig.AIProviderID.UUID +} + +// sanitizeCompactionPrompt adapts a prompt built for the chat model to a +// differing compaction model. The input messages are never mutated; the +// assistant generation keeps using the original prompt. +func sanitizeCompactionPrompt( + ctx context.Context, + logger slog.Logger, + prompt []fantasy.Message, + compactionModel fantasy.LanguageModel, + chatConfig database.ChatModelConfig, + overrideConfig database.ChatModelConfig, +) []fantasy.Message { + messages := prompt + if !sameCompactionProviderIdentity(chatConfig, overrideConfig) { + messages = flattenProviderExecutedToolParts(ctx, logger, messages) + } + messages = replaceUnsupportedFileParts(ctx, logger, messages, func(mediaType string) bool { + return chatprovider.AcceptsFilePartMediaType( + compactionModel.Provider(), + compactionModel.Model(), + mediaType, + ) + }) + sanitized, stats := chatsanitize.SanitizeAnthropicProviderToolHistory( + compactionModel.Provider(), + messages, + ) + chatsanitize.LogAnthropicProviderToolSanitization( + ctx, + logger, + "compaction_prompt", + compactionModel.Provider(), + compactionModel.Model(), + stats, + ) + return sanitized +} + +// flattenProviderExecutedToolParts rewrites provider-executed tool calls +// and results in assistant messages into text parts on a copy of messages, +// keeping their content while shedding provider-specific wire shapes other +// providers reject on replay. Provider-executed parts outside assistant +// messages are anomalous and dropped, since a text part is not valid +// tool-message content everywhere; messages emptied by the drop are removed. +func flattenProviderExecutedToolParts( + ctx context.Context, + logger slog.Logger, + messages []fantasy.Message, +) []fantasy.Message { + flattened := 0 + dropped := 0 + // Tool names live on the call part only; results reference the call by ID. + toolNamesByCallID := make(map[string]string) + out := make([]fantasy.Message, 0, len(messages)) + for _, msg := range messages { + flattenToText := msg.Role == fantasy.MessageRoleAssistant + parts := make([]fantasy.MessagePart, 0, len(msg.Content)) + for _, part := range msg.Content { + switch typed := part.(type) { + case fantasy.ToolCallPart: + if typed.ProviderExecuted { + if !flattenToText { + dropped++ + continue + } + toolNamesByCallID[typed.ToolCallID] = typed.ToolName + flattened++ + parts = append(parts, fantasy.TextPart{ + Text: fmt.Sprintf("[Server tool call: %s] %s", typed.ToolName, typed.Input), + }) + continue + } + case fantasy.ToolResultPart: + if typed.ProviderExecuted { + if !flattenToText { + dropped++ + continue + } + flattened++ + parts = append(parts, fantasy.TextPart{ + Text: fmt.Sprintf( + "[Server tool result: %s] %s", + toolNamesByCallID[typed.ToolCallID], + stringifyToolResultOutput(typed.Output), + ), + }) + continue + } + } + parts = append(parts, part) + } + if len(parts) == 0 && len(msg.Content) > 0 { + continue + } + msg.Content = parts + out = append(out, msg) + } + if flattened > 0 || dropped > 0 { + logger.Debug(ctx, "flattened provider-executed tool history in compaction prompt", + slog.F("flattened_parts", flattened), + slog.F("dropped_parts", dropped), + ) + } + return out +} + +// stringifyToolResultOutput renders a tool result as prompt text. Media +// payloads are summarized so base64 data does not enter the prompt. +func stringifyToolResultOutput(output fantasy.ToolResultOutputContent) string { + switch typed := output.(type) { + case fantasy.ToolResultOutputContentText: + return typed.Text + case fantasy.ToolResultOutputContentError: + if typed.Error == nil { + return "error" + } + return typed.Error.Error() + case fantasy.ToolResultOutputContentMedia: + if typed.Text != "" { + return fmt.Sprintf("%s [media %s omitted]", typed.Text, typed.MediaType) + } + return fmt.Sprintf("[media %s omitted]", typed.MediaType) + default: + return "[unserializable tool output]" + } +} + +// replaceUnsupportedFileParts swaps file parts the compaction model does +// not accept for text placeholders in a copy of messages, so the summary +// notes the attachment existed instead of silently losing it. +func replaceUnsupportedFileParts( + ctx context.Context, + logger slog.Logger, + messages []fantasy.Message, + acceptsFilePart func(mediaType string) bool, +) []fantasy.Message { + replaced := 0 + out := make([]fantasy.Message, 0, len(messages)) + for _, msg := range messages { + parts := make([]fantasy.MessagePart, 0, len(msg.Content)) + for _, part := range msg.Content { + filePart, ok := part.(fantasy.FilePart) + if !ok || acceptsFilePart(filePart.MediaType) { + parts = append(parts, part) + continue + } + replaced++ + parts = append(parts, fantasy.TextPart{ + Text: fmt.Sprintf( + "[Attachment %q (%s) omitted: not supported by the compaction model]", + filePart.Filename, + filePart.MediaType, + ), + ProviderOptions: filePart.ProviderOptions, + }) + } + msg.Content = parts + out = append(out, msg) + } + if replaced > 0 { + logger.Debug(ctx, "replaced unsupported file parts in compaction prompt", + slog.F("replaced_parts", replaced), + ) + } + return out +} diff --git a/coderd/x/chatd/compaction_sanitize_internal_test.go b/coderd/x/chatd/compaction_sanitize_internal_test.go new file mode 100644 index 00000000000..f3b21f7bc10 --- /dev/null +++ b/coderd/x/chatd/compaction_sanitize_internal_test.go @@ -0,0 +1,197 @@ +package chatd + +import ( + "testing" + + "charm.land/fantasy" + "github.com/google/uuid" + "github.com/stretchr/testify/require" + + "cdr.dev/slog/v3/sloggers/slogtest" + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/x/chatd/chattest" + "github.com/coder/coder/v2/testutil" +) + +func TestSameCompactionProviderIdentity(t *testing.T) { + t.Parallel() + + providerID := uuid.New() + + require.True(t, sameCompactionProviderIdentity(configWithProvider(providerID), configWithProvider(providerID))) + require.False(t, sameCompactionProviderIdentity(configWithProvider(providerID), configWithProvider(uuid.New()))) + // Legacy configs without a provider FK compare as different (fail closed). + require.False(t, sameCompactionProviderIdentity(database.ChatModelConfig{}, configWithProvider(providerID))) + require.False(t, sameCompactionProviderIdentity(database.ChatModelConfig{}, database.ChatModelConfig{})) +} + +func configWithProvider(id uuid.UUID) database.ChatModelConfig { + return database.ChatModelConfig{AIProviderID: uuid.NullUUID{UUID: id, Valid: true}} +} + +func TestSanitizeCompactionPrompt_FlattensForeignProviderExecutedToolParts(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitShort) + logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) + prompt := []fantasy.Message{ + { + Role: fantasy.MessageRoleUser, + Content: []fantasy.MessagePart{ + fantasy.TextPart{Text: "search the web"}, + }, + }, + { + Role: fantasy.MessageRoleAssistant, + Content: []fantasy.MessagePart{ + fantasy.TextPart{Text: "searching"}, + fantasy.ToolCallPart{ + ToolCallID: "ws-1", + ToolName: "web_search", + Input: `{"query":"coder"}`, + ProviderExecuted: true, + }, + fantasy.ToolResultPart{ + ToolCallID: "ws-1", + Output: fantasy.ToolResultOutputContentText{Text: "results"}, + ProviderExecuted: true, + }, + }, + }, + { + Role: fantasy.MessageRoleAssistant, + Content: []fantasy.MessagePart{ + fantasy.ToolCallPart{ + ToolCallID: "local-1", + ToolName: "read_file", + Input: `{"path":"/tmp/a.txt"}`, + }, + }, + }, + } + + compactionModel := &chattest.FakeModel{ProviderName: "openai", ModelName: "gpt-4.1-mini"} + sanitized := sanitizeCompactionPrompt(ctx, logger, prompt, compactionModel, configWithProvider(uuid.New()), configWithProvider(uuid.New())) + + require.Len(t, sanitized, 3) + // Provider-executed parts are flattened to text so the summary keeps + // their content without the provider-specific wire shape. + require.Len(t, sanitized[1].Content, 3) + require.Equal(t, fantasy.TextPart{Text: "searching"}, sanitized[1].Content[0]) + require.Equal(t, fantasy.TextPart{Text: `[Server tool call: web_search] {"query":"coder"}`}, sanitized[1].Content[1]) + require.Equal(t, fantasy.TextPart{Text: "[Server tool result: web_search] results"}, sanitized[1].Content[2]) + // Local tool calls replay fine across providers and must survive. + require.Len(t, sanitized[2].Content, 1) + require.Equal(t, "read_file", sanitized[2].Content[0].(fantasy.ToolCallPart).ToolName) + + // The original prompt used for assistant generation is untouched. + require.Equal(t, fantasy.ToolCallPart{ + ToolCallID: "ws-1", + ToolName: "web_search", + Input: `{"query":"coder"}`, + ProviderExecuted: true, + }, prompt[1].Content[1]) +} + +func TestSanitizeCompactionPrompt_DropsNonAssistantProviderExecutedParts(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitShort) + logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) + // Provider-executed parts outside assistant messages are anomalous; a + // flattened text part is not valid tool-message content, so they drop. + prompt := []fantasy.Message{ + { + Role: fantasy.MessageRoleTool, + Content: []fantasy.MessagePart{ + fantasy.ToolResultPart{ + ToolCallID: "ws-1", + Output: fantasy.ToolResultOutputContentText{Text: "results"}, + ProviderExecuted: true, + }, + }, + }, + { + Role: fantasy.MessageRoleUser, + Content: []fantasy.MessagePart{ + fantasy.TextPart{Text: "hello"}, + }, + }, + } + + compactionModel := &chattest.FakeModel{ProviderName: "openai", ModelName: "gpt-4.1-mini"} + sanitized := sanitizeCompactionPrompt(ctx, logger, prompt, compactionModel, configWithProvider(uuid.New()), configWithProvider(uuid.New())) + + require.Len(t, sanitized, 1) + require.Equal(t, fantasy.MessageRoleUser, sanitized[0].Role) + require.Len(t, prompt, 2) +} + +func TestSanitizeCompactionPrompt_ReplacesUnsupportedFileParts(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitShort) + logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) + prompt := []fantasy.Message{ + { + Role: fantasy.MessageRoleUser, + Content: []fantasy.MessagePart{ + fantasy.TextPart{Text: "look at this"}, + fantasy.FilePart{ + Filename: "diagram.pdf", + Data: []byte("%PDF-"), + MediaType: "application/pdf", + }, + }, + }, + } + + // Mistral accepts images but not PDFs, so the PDF part must become a + // placeholder while the prompt stays otherwise intact. + compactionModel := &chattest.FakeModel{ProviderName: "mistral", ModelName: "mistral-large"} + sharedProviderID := uuid.New() + sanitized := sanitizeCompactionPrompt(ctx, logger, prompt, compactionModel, configWithProvider(sharedProviderID), configWithProvider(sharedProviderID)) + + require.Len(t, sanitized, 1) + require.Len(t, sanitized[0].Content, 2) + textPart, ok := sanitized[0].Content[1].(fantasy.TextPart) + require.True(t, ok) + require.Contains(t, textPart.Text, "diagram.pdf") + require.Contains(t, textPart.Text, "not supported by the compaction model") + + // The original prompt keeps its file part. + _, ok = prompt[0].Content[1].(fantasy.FilePart) + require.True(t, ok) +} + +func TestSanitizeCompactionPrompt_SameProviderKeepsProviderExecutedParts(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitShort) + logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) + prompt := []fantasy.Message{ + { + Role: fantasy.MessageRoleAssistant, + Content: []fantasy.MessagePart{ + fantasy.ToolCallPart{ + ToolCallID: "ws-1", + ToolName: "web_search", + Input: `{"query":"coder"}`, + ProviderExecuted: true, + }, + fantasy.ToolResultPart{ + ToolCallID: "ws-1", + Output: fantasy.ToolResultOutputContentText{Text: "results"}, + ProviderExecuted: true, + }, + }, + }, + } + + compactionModel := &chattest.FakeModel{ProviderName: "openai", ModelName: "gpt-4.1-mini"} + sharedProviderID := uuid.New() + sanitized := sanitizeCompactionPrompt(ctx, logger, prompt, compactionModel, configWithProvider(sharedProviderID), configWithProvider(sharedProviderID)) + + require.Len(t, sanitized, 1) + require.Len(t, sanitized[0].Content, 2) +} diff --git a/coderd/x/chatd/generation.go b/coderd/x/chatd/generation.go index ac3bcd9a128..25d553ef2cd 100644 --- a/coderd/x/chatd/generation.go +++ b/coderd/x/chatd/generation.go @@ -69,6 +69,14 @@ type generationPrepared struct { // generationCompaction contains compaction inputs prepared for generation. type generationCompaction struct { + // Override, when non-nil, is the compaction model override resolved at + // prepare time. Its model client is built in the compact action path, + // so construction failures cannot fail turns that never compact. + Override *resolvedCompactionOverride + // ChatModelConfig is the chat model's config, used to detect provider + // changes when sanitizing the compaction prompt. + ChatModelConfig database.ChatModelConfig + Required bool Options chatloop.GenerateCompactionOptions } @@ -238,6 +246,18 @@ func generationCompactionThreshold(compaction *generationCompaction) int32 { return compaction.Options.ThresholdPercent } +// generationCompactionContextLimit returns the context limit the compaction +// trigger was evaluated against at prepare time (the stricter of the chat and +// override models' limits). The still-over-limit check must compare against +// the same limit, otherwise a stricter override loops through repeated +// compactions instead of surfacing errCompactionStillOverLimit. +func generationCompactionContextLimit(compaction *generationCompaction) int64 { + if compaction == nil { + return 0 + } + return compaction.Options.ContextLimit +} + func unresolvedToolCallsFromHistory( messages []database.ChatMessage, dynamicToolNames map[string]bool, @@ -324,7 +344,7 @@ func (s *taskStarter) StartGeneration(ctx context.Context, input chatWorkerTaskS compactionEnabled: prepared.Compaction != nil, compactionNeeded: prepared.Compaction != nil && prepared.Compaction.Required, compactionThresholdPercent: generationCompactionThreshold(prepared.Compaction), - compactionContextLimit: prepared.ContextLimitFallback, + compactionContextLimit: generationCompactionContextLimit(prepared.Compaction), }) }) if err != nil { @@ -333,9 +353,10 @@ func (s *taskStarter) StartGeneration(ctx context.Context, input chatWorkerTaskS return xerrors.Errorf("decide generation: %w", err) } if errors.Is(err, errCompactionStillOverLimit) && prepared.Compaction != nil { + metricProvider, metricModel := compactionMetricIdentity(prepared.Compaction) s.server.metrics.RecordCompaction( - compactionProvider(prepared.Compaction.Options), - compactionModel(prepared.Compaction.Options), + metricProvider, + metricModel, false, errCompactionStillOverLimit, ) @@ -683,15 +704,43 @@ func (s *taskStarter) generateCompaction( return s.finishGenerationError(ctx, machine, input, xerrors.New("compaction action missing options"), requireGenerationAttempt(attempt.number)) } compactionOpts := prepared.Compaction.Options + metricProvider, metricModel := compactionMetricIdentity(prepared.Compaction) + if override := prepared.Compaction.Override; override != nil { + overrideModel, err := s.server.buildCompactionOverrideModel(ctx, prepared.Chat, override.Config, prepared.ModelBuildOptions) + if err != nil { + return xerrors.Errorf("build compaction model override: %w", err) + } + logger := s.server.logger.With( + slog.F("chat_id", prepared.Chat.ID), + slog.F("owner_id", prepared.Chat.OwnerID), + ) + compactionOpts.Model = overrideModel.model + compactionOpts.ResolvedProvider = overrideModel.resolvedProvider + compactionOpts.ResolvedModel = overrideModel.resolvedModel + compactionOpts.ModelConfigID = overrideModel.modelConfig.ID + compactionOpts.ProviderOptions = overrideModel.providerOptions + compactionOpts.Messages = sanitizeCompactionPrompt( + ctx, + logger, + compactionOpts.Messages, + overrideModel.model, + prepared.Compaction.ChatModelConfig, + overrideModel.modelConfig, + ) + } compactionOpts.PublishMessagePart = attempt.publish - outcome, err := chatloop.GenerateCompaction(ctx, compactionOpts) + // Attach the turn debug run so the compaction call records a child + // debug run; without it startCompactionDebugRun finds no parent and + // skips debug instrumentation entirely. + runCtx := input.DebugTurn.Ensure(ctx, prepared.Chat, prepared.Debug) + outcome, err := chatloop.GenerateCompaction(runCtx, compactionOpts) if err != nil { - s.server.metrics.RecordCompaction(compactionProvider(compactionOpts), compactionModel(compactionOpts), false, err) + s.server.metrics.RecordCompaction(metricProvider, metricModel, false, err) return xerrors.Errorf("generate compaction: %w", err) } if strings.TrimSpace(outcome.SystemSummary) == "" || strings.TrimSpace(outcome.SummaryReport) == "" { err := xerrors.New("compaction produced no summary") - s.server.metrics.RecordCompaction(compactionProvider(compactionOpts), compactionModel(compactionOpts), false, err) + s.server.metrics.RecordCompaction(metricProvider, metricModel, false, err) return s.finishGenerationError(ctx, machine, input, err, requireGenerationAttempt(attempt.number)) } messages, err := buildCompactionMessages(buildCompactionMessagesInput{ @@ -703,20 +752,31 @@ func (s *taskStarter) generateCompaction( contentVersion: chatprompt.CurrentContentVersion, }) if err != nil { - s.server.metrics.RecordCompaction(compactionProvider(compactionOpts), compactionModel(compactionOpts), false, err) + s.server.metrics.RecordCompaction(metricProvider, metricModel, false, err) return s.finishGenerationError(ctx, machine, input, err, requireGenerationAttempt(attempt.number)) } err = s.commitGenerationStep(ctx, machine, input, attempt.number, generationActionCompact, stepMessagesForCommit{ Messages: messages.Messages, VisibleIndexes: visibleMessageIndexes(messages.Messages), }) - s.server.metrics.RecordCompaction(compactionProvider(compactionOpts), compactionModel(compactionOpts), err == nil, err) + s.server.metrics.RecordCompaction(metricProvider, metricModel, err == nil, err) if err != nil { return xerrors.Errorf("commit generation step: %w", err) } return nil } +// compactionMetricIdentity returns the provider/model labels for compaction +// metrics. Override labels come from prepare-time resolution so events +// recorded before the override client is built (still-over-limit) match +// the compact action's own events. +func compactionMetricIdentity(compaction *generationCompaction) (provider, model string) { + if compaction.Override != nil { + return compaction.Override.ResolvedProvider, compaction.Override.ResolvedModel + } + return compactionProvider(compaction.Options), compactionModel(compaction.Options) +} + func compactionProvider(opts chatloop.GenerateCompactionOptions) string { if opts.Model == nil { return "" diff --git a/coderd/x/chatd/generation_internal_test.go b/coderd/x/chatd/generation_internal_test.go index 9090506b671..aa8e93a3d96 100644 --- a/coderd/x/chatd/generation_internal_test.go +++ b/coderd/x/chatd/generation_internal_test.go @@ -7,10 +7,50 @@ import ( "golang.org/x/xerrors" "github.com/coder/coder/v2/coderd/x/chatd/chatdebug" + "github.com/coder/coder/v2/coderd/x/chatd/chatloop" "github.com/coder/coder/v2/coderd/x/chatd/chatstate" + "github.com/coder/coder/v2/coderd/x/chatd/chattest" "github.com/coder/coder/v2/testutil" ) +func TestCompactionMetricIdentity(t *testing.T) { + t.Parallel() + + compaction := &generationCompaction{ + Options: chatloop.GenerateCompactionOptions{ + Model: &chattest.FakeModel{ProviderName: "anthropic", ModelName: "claude-sonnet-4-5"}, + }, + } + + provider, model := compactionMetricIdentity(compaction) + require.Equal(t, "anthropic", provider) + require.Equal(t, "claude-sonnet-4-5", model) + + // With an override, metrics use the prepare-time identity, not the + // chat model carried by the options. + compaction.Override = &resolvedCompactionOverride{ + ResolvedProvider: "openai", + ResolvedModel: "gpt-4.1-mini", + } + provider, model = compactionMetricIdentity(compaction) + require.Equal(t, "openai", provider) + require.Equal(t, "gpt-4.1-mini", model) +} + +func TestGenerationCompactionContextLimit(t *testing.T) { + t.Parallel() + + require.EqualValues(t, 0, generationCompactionContextLimit(nil)) + + // The decision path must see the prepare-time compaction limit (the + // stricter of the chat and override models' limits), not the chat + // model's limit. + compaction := &generationCompaction{ + Options: chatloop.GenerateCompactionOptions{ContextLimit: 50_000}, + } + require.EqualValues(t, 50_000, generationCompactionContextLimit(compaction)) +} + func TestRecordGenerationFinishFailure(t *testing.T) { t.Parallel() diff --git a/coderd/x/chatd/generation_preparer.go b/coderd/x/chatd/generation_preparer.go index 6ca947b1f79..ffee255d126 100644 --- a/coderd/x/chatd/generation_preparer.go +++ b/coderd/x/chatd/generation_preparer.go @@ -580,20 +580,41 @@ func (server *Server) prepareGeneration( if override, ok := server.resolveUserCompactionThreshold(ctx, chat.OwnerID, modelConfig.ID); ok { effectiveThreshold = override } + // The compaction trigger uses the stricter of the chat and override + // models' context limits: the history must also fit the summarizer's + // window. + compactionContextLimit := modelConfig.ContextLimit + compactionOverride, err := server.resolveCompactionOverrideConfig(ctx, chat) + if err != nil { + cleanup() + return generationPrepared{}, err + } + if compactionOverride != nil { + if overrideLimit := compactionOverride.Config.ContextLimit; overrideLimit > 0 && + (compactionContextLimit <= 0 || overrideLimit < compactionContextLimit) { + compactionContextLimit = overrideLimit + } + } + compactionStepUsage := latestPromptUsage(promptRows) + compactionNeeded := shouldCompactPromptUsage(compactionStepUsage, compactionContextLimit, effectiveThreshold) + // The options carry the chat model; generateCompaction swaps in the + // override client when one is configured. compactionOptions := chatloop.GenerateCompactionOptions{ Model: model, Messages: prompt, ThresholdPercent: effectiveThreshold, - ContextLimit: modelConfig.ContextLimit, - ContextLimitFallback: modelConfig.ContextLimit, + ContextLimit: compactionContextLimit, + ContextLimitFallback: compactionContextLimit, ToolCallID: compactionToolCallID, ToolName: "chat_summarized", DebugSvc: debugSvc, ChatID: chat.ID, HistoryTipMessageID: historyTipMessageID, + ResolvedProvider: resolvedProvider, + ResolvedModel: debugModel, + ModelConfigID: modelConfig.ID, + StepUsage: compactionStepUsage, } - compactionOptions.StepUsage = latestPromptUsage(promptRows) - compactionNeeded := shouldCompactPromptUsage(compactionOptions.StepUsage, modelConfig.ContextLimit, effectiveThreshold) // workspaceCtx.currentChatSnapshot may carry a freshly persisted // AgentID/BuildID binding from the getWorkspaceAgent call above. @@ -626,8 +647,10 @@ func (server *Server) prepareGeneration( ToolNameToConfigID: toolNameToConfigID, MaxSteps: maxChatSteps, Compaction: &generationCompaction{ - Required: compactionNeeded, - Options: compactionOptions, + Override: compactionOverride, + ChatModelConfig: modelConfig, + Required: compactionNeeded, + Options: compactionOptions, }, Cleanup: cleanup, Debug: debug, diff --git a/coderd/x/chatd/subagent.go b/coderd/x/chatd/subagent.go index ccb7e1a26d3..9ad563247c5 100644 --- a/coderd/x/chatd/subagent.go +++ b/coderd/x/chatd/subagent.go @@ -184,6 +184,7 @@ func modelOverrideErrorLabel(overrideContext string) string { // resolveConfiguredModelOverride returns ok when a usable override is // resolved. In hard failure mode, ok is also true for configured but unusable // overrides so callers can distinguish them from unset or malformed values. +// The normalized provider name is only meaningful for a usable override. func (p *Server) resolveConfiguredModelOverride( ctx context.Context, overrideContext string, @@ -192,7 +193,7 @@ func (p *Server) resolveConfiguredModelOverride( resolveModelConfig modelOverrideConfigResolver, resolveProviderKeys modelOverrideProviderKeysResolver, failureMode modelOverrideFailureMode, -) (database.ChatModelConfig, *string, bool, error) { +) (database.ChatModelConfig, string, *string, bool, error) { parsed, ok := parseModelOverride(raw) if !ok { p.logger.Info(ctx, @@ -200,10 +201,10 @@ func (p *Server) resolveConfiguredModelOverride( slog.F("override_context", overrideContext), slog.F("raw_model_config_id", strings.TrimSpace(raw)), ) - return database.ChatModelConfig{}, nil, false, nil + return database.ChatModelConfig{}, "", nil, false, nil } if parsed.modelConfigID == uuid.Nil { - return database.ChatModelConfig{}, nil, false, nil + return database.ChatModelConfig{}, "", nil, false, nil } modelConfig, providerName, err := resolveModelConfig( @@ -215,20 +216,20 @@ func (p *Server) resolveConfiguredModelOverride( label := modelOverrideErrorLabel(overrideContext) switch { case errors.Is(err, sql.ErrNoRows): - return database.ChatModelConfig{}, parsed.reasoningEffort, true, xerrors.Errorf( + return database.ChatModelConfig{}, "", parsed.reasoningEffort, true, xerrors.Errorf( "%s model override is unavailable: %s", label, parsed.modelConfigID, ) case errors.Is(err, errInvalidModelOverrideMetadata): - return database.ChatModelConfig{}, parsed.reasoningEffort, true, xerrors.Errorf( + return database.ChatModelConfig{}, "", parsed.reasoningEffort, true, xerrors.Errorf( "%s model override metadata is invalid for %s: %w", label, parsed.modelConfigID, err, ) default: - return database.ChatModelConfig{}, parsed.reasoningEffort, true, xerrors.Errorf( + return database.ChatModelConfig{}, "", parsed.reasoningEffort, true, xerrors.Errorf( "resolve %s model override %s: %w", label, parsed.modelConfigID, @@ -259,19 +260,19 @@ func (p *Server) resolveConfiguredModelOverride( slog.Error(err), ) } - return database.ChatModelConfig{}, nil, false, nil + return database.ChatModelConfig{}, "", nil, false, nil } providerKeys, err := resolveProviderKeys(ctx, ownerID, modelConfigAIProviderID(modelConfig)) if err != nil { - return database.ChatModelConfig{}, nil, false, xerrors.Errorf( + return database.ChatModelConfig{}, "", nil, false, xerrors.Errorf( "resolve provider API keys: %w", err, ) } if !userCanUseProviderKeys(providerKeys, providerName) { if failureMode == modelOverrideFailureModeHard { - return database.ChatModelConfig{}, parsed.reasoningEffort, true, xerrors.Errorf( + return database.ChatModelConfig{}, "", parsed.reasoningEffort, true, xerrors.Errorf( "%s model override credentials are unavailable for provider %q", modelOverrideErrorLabel(overrideContext), providerName, @@ -284,9 +285,9 @@ func (p *Server) resolveConfiguredModelOverride( slog.F("model_config_id", parsed.modelConfigID), slog.F("provider", providerName), ) - return database.ChatModelConfig{}, nil, false, nil + return database.ChatModelConfig{}, "", nil, false, nil } - return modelConfig, parsed.reasoningEffort, true, nil + return modelConfig, providerName, parsed.reasoningEffort, true, nil } func (p *Server) resolvePersonalSubagentModelConfigID( @@ -483,7 +484,7 @@ func (p *Server) resolveSubagentModelConfigID( err, ) } - modelConfig, reasoningEffort, ok, err := p.resolveConfiguredModelOverride( + modelConfig, _, reasoningEffort, ok, err := p.resolveConfiguredModelOverride( chatdCtx, string(overrideContext), raw, diff --git a/coderd/x/chatd/subagent_internal_test.go b/coderd/x/chatd/subagent_internal_test.go index 9ea4d057ad7..ada9354c3da 100644 --- a/coderd/x/chatd/subagent_internal_test.go +++ b/coderd/x/chatd/subagent_internal_test.go @@ -1490,7 +1490,7 @@ func TestResolveConfiguredModelOverride_AcceptsAmbientCredentialsProvider( Enabled: true, } - resolvedModelConfig, reasoningEffort, ok, err := server.resolveConfiguredModelOverride( + resolvedModelConfig, _, reasoningEffort, ok, err := server.resolveConfiguredModelOverride( ctx, "plan", modelConfig.ID.String(), diff --git a/coderd/x/chatd/title_override.go b/coderd/x/chatd/title_override.go index 0072587498a..a4cfb4de8f7 100644 --- a/coderd/x/chatd/title_override.go +++ b/coderd/x/chatd/title_override.go @@ -70,7 +70,7 @@ func (p *Server) resolveTitleGenerationModelOverride( ) } - modelConfig, overrideEffort, overrideSet, err := p.resolveConfiguredModelOverride( + modelConfig, _, overrideEffort, overrideSet, err := p.resolveConfiguredModelOverride( ctx, titleGenerationOverrideContext, raw, diff --git a/codersdk/chats.go b/codersdk/chats.go index 7a02d86aaa5..27d02fac8a8 100644 --- a/codersdk/chats.go +++ b/codersdk/chats.go @@ -754,6 +754,7 @@ const ( ChatModelOverrideContextGeneral ChatModelOverrideContext = "general" ChatModelOverrideContextExplore ChatModelOverrideContext = "explore" ChatModelOverrideContextTitleGeneration ChatModelOverrideContext = "title_generation" + ChatModelOverrideContextCompaction ChatModelOverrideContext = "compaction" ) // Valid reports whether the override context is one of the supported values. @@ -761,7 +762,8 @@ func (c ChatModelOverrideContext) Valid() bool { switch c { case ChatModelOverrideContextGeneral, ChatModelOverrideContextExplore, - ChatModelOverrideContextTitleGeneration: + ChatModelOverrideContextTitleGeneration, + ChatModelOverrideContextCompaction: return true default: return false @@ -774,6 +776,7 @@ func AllChatModelOverrideContexts() []ChatModelOverrideContext { ChatModelOverrideContextGeneral, ChatModelOverrideContextExplore, ChatModelOverrideContextTitleGeneration, + ChatModelOverrideContextCompaction, } } diff --git a/site/src/api/typesGenerated.ts b/site/src/api/typesGenerated.ts index 740a5e729fe..d9330790849 100644 --- a/site/src/api/typesGenerated.ts +++ b/site/src/api/typesGenerated.ts @@ -2647,11 +2647,13 @@ export interface ChatModelOpenRouterProviderOptions { // From codersdk/chats.go export type ChatModelOverrideContext = + | "compaction" | "explore" | "general" | "title_generation"; export const ChatModelOverrideContexts: ChatModelOverrideContext[] = [ + "compaction", "explore", "general", "title_generation", diff --git a/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPage.tsx b/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPage.tsx index 4d549028c54..d794dd74f3b 100644 --- a/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPage.tsx +++ b/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPage.tsx @@ -28,6 +28,8 @@ const generalOverrideContext: TypesGen.ChatModelOverrideContext = "general"; const exploreOverrideContext: TypesGen.ChatModelOverrideContext = "explore"; const titleGenerationOverrideContext: TypesGen.ChatModelOverrideContext = "title_generation"; +const compactionOverrideContext: TypesGen.ChatModelOverrideContext = + "compaction"; const chatModelOverrideKey = (context: TypesGen.ChatModelOverrideContext) => ["chat-model-override", context] as const; @@ -79,6 +81,10 @@ const CoderAgentsPage: FC = () => { ...chatModelOverrideQuery(titleGenerationOverrideContext), enabled: canEditDeploymentConfig, }); + const compactionModelQuery = useQuery({ + ...chatModelOverrideQuery(compactionOverrideContext), + enabled: canEditDeploymentConfig, + }); const modelConfigsQuery = useQuery(chatModelConfigs()); const advisorConfigQuery = useQuery({ ...chatAdvisorConfig(), @@ -104,6 +110,9 @@ const CoderAgentsPage: FC = () => { titleGenerationOverrideContext, ), ); + const saveCompactionModelMutation = useMutation( + updateChatModelOverrideMutation(queryClient, compactionOverrideContext), + ); const saveExploreModelOverrideMutation = useMutation( updateChatModelOverrideMutation(queryClient, exploreOverrideContext), ); @@ -141,6 +150,7 @@ const CoderAgentsPage: FC = () => { } generalModelOverrideData={generalModelOverrideQuery.data} titleGenerationModelOverrideData={titleGenerationModelQuery.data} + compactionModelOverrideData={compactionModelQuery.data} exploreModelOverrideData={exploreModelOverrideQuery.data} modelConfigsData={modelConfigsQuery.data} providerInfoByID={providerInfoByID} @@ -161,6 +171,9 @@ const CoderAgentsPage: FC = () => { isSaveTitleGenerationModelError={ saveTitleGenerationModelMutation.isError } + onSaveCompactionModel={saveCompactionModelMutation.mutate} + isSavingCompactionModel={saveCompactionModelMutation.isPending} + isSaveCompactionModelError={saveCompactionModelMutation.isError} onSaveExploreModelOverride={saveExploreModelOverrideMutation.mutate} isSavingExploreModelOverride={ saveExploreModelOverrideMutation.isPending diff --git a/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.stories.tsx b/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.stories.tsx index 04e818b5ed6..5ce4eec6427 100644 --- a/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.stories.tsx +++ b/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.stories.tsx @@ -94,6 +94,14 @@ const exploreDisabledModelConfig = buildModelConfig({ context_limit: 200_000, }); +const compactionDisabledModelConfig = buildModelConfig({ + id: "model-compaction-disabled", + model: "gpt-4.1-nano-legacy", + display_name: "GPT 4.1 Nano Legacy", + enabled: false, + context_limit: 128_000, +}); + const allModelConfigs: TypesGen.ChatModelConfig[] = [ generalModelConfig, claudeSonnetModelConfig, @@ -102,6 +110,7 @@ const allModelConfigs: TypesGen.ChatModelConfig[] = [ generalDisabledModelConfig, titleDisabledModelConfig, exploreDisabledModelConfig, + compactionDisabledModelConfig, ]; const providerInfoByID = new Map([ @@ -124,6 +133,7 @@ const buildArgs = ( isSaveAdminOverridesError: false, generalModelOverrideData: buildOverrideData("general"), titleGenerationModelOverrideData: buildTitleGenerationModelOverrideData(), + compactionModelOverrideData: buildOverrideData("compaction"), exploreModelOverrideData: buildOverrideData("explore"), modelConfigsData: allModelConfigs, providerInfoByID, @@ -136,6 +146,9 @@ const buildArgs = ( onSaveTitleGenerationModel: fn(), isSavingTitleGenerationModel: false, isSaveTitleGenerationModelError: false, + onSaveCompactionModel: fn(), + isSavingCompactionModel: false, + isSaveCompactionModelError: false, onSaveExploreModelOverride: fn(), isSavingExploreModelOverride: false, isSaveExploreModelOverrideError: false, @@ -211,6 +224,7 @@ export const AllOverridesUnset: Story = { expect(headings.map((heading) => heading.textContent?.trim())).toEqual([ "General model", "Title generation model", + "Compaction model", "Explore subagent model", ]); await canvas.findByText( @@ -223,6 +237,10 @@ export const AllOverridesUnset: Story = { headingName: "Title generation model", placeholder: "Use title default", }, + { + headingName: "Compaction model", + placeholder: "Use chat model", + }, { headingName: "Explore subagent model", placeholder: "Use chat default", @@ -293,6 +311,9 @@ export const EachOverrideSetToEnabledModel: Story = { titleGenerationModelOverrideData: buildTitleGenerationModelOverrideData({ model_config_id: titleModelConfig.id, }), + compactionModelOverrideData: buildOverrideData("compaction", { + model_config_id: claudeSonnetModelConfig.id, + }), exploreModelOverrideData: buildOverrideData("explore", { model_config_id: exploreFallbackModelConfig.id, }), @@ -303,6 +324,10 @@ export const EachOverrideSetToEnabledModel: Story = { canvasElement, "Title generation model", ); + const compactionSection = await getSection( + canvasElement, + "Compaction model", + ); const exploreSection = await getSection( canvasElement, "Explore subagent model", @@ -360,6 +385,26 @@ export const EachOverrideSetToEnabledModel: Story = { ); }); + await selectModelInSection( + compactionSection, + canvasElement, + /claude sonnet 4/i, + "GPT 4o Mini", + ); + const compactionSaveButton = within(compactionSection).getByRole("button", { + name: "Save", + }); + await waitFor(() => { + expect(compactionSaveButton).toBeEnabled(); + }); + await userEvent.click(compactionSaveButton); + await waitFor(() => { + expect(args.onSaveCompactionModel).toHaveBeenCalledWith( + { model_config_id: titleModelConfig.id }, + expect.anything(), + ); + }); + const exploreClearButton = within(exploreSection).getByRole("button", { name: "Clear", }); @@ -388,6 +433,9 @@ export const MalformedOverridesRemainClearableAndSaveable: Story = { titleGenerationModelOverrideData: buildTitleGenerationModelOverrideData({ is_malformed: true, }), + compactionModelOverrideData: buildOverrideData("compaction", { + is_malformed: true, + }), exploreModelOverrideData: buildOverrideData("explore", { is_malformed: true, }), @@ -398,12 +446,21 @@ export const MalformedOverridesRemainClearableAndSaveable: Story = { canvasElement, "Title generation model", ); + const compactionSection = await getSection( + canvasElement, + "Compaction model", + ); const exploreSection = await getSection( canvasElement, "Explore subagent model", ); - for (const section of [generalSection, titleSection, exploreSection]) { + for (const section of [ + generalSection, + titleSection, + compactionSection, + exploreSection, + ]) { await within(section).findByText(OVERRIDE_MALFORMED_WARNING); } @@ -448,6 +505,20 @@ export const MalformedOverridesRemainClearableAndSaveable: Story = { expect.anything(), ); }); + + const compactionSaveButton = within(compactionSection).getByRole("button", { + name: "Save", + }); + await waitFor(() => { + expect(compactionSaveButton).toBeEnabled(); + }); + await userEvent.click(compactionSaveButton); + await waitFor(() => { + expect(args.onSaveCompactionModel).toHaveBeenCalledWith( + { model_config_id: "" }, + expect.anything(), + ); + }); }, }; @@ -459,6 +530,9 @@ export const UnavailableSavedModels: Story = { titleGenerationModelOverrideData: buildTitleGenerationModelOverrideData({ model_config_id: titleDisabledModelConfig.id, }), + compactionModelOverrideData: buildOverrideData("compaction", { + model_config_id: compactionDisabledModelConfig.id, + }), exploreModelOverrideData: buildOverrideData("explore", { model_config_id: exploreDisabledModelConfig.id, }), @@ -469,12 +543,16 @@ export const UnavailableSavedModels: Story = { canvasElement, "Title generation model", ); + const compactionSection = await getSection( + canvasElement, + "Compaction model", + ); const exploreSection = await getSection( canvasElement, "Explore subagent model", ); - for (const section of [generalSection, exploreSection]) { + for (const section of [generalSection, compactionSection, exploreSection]) { await within(section).findByText(UNAVAILABLE_SAVED_MODEL_WARNING); expect( within(section).getByRole("combobox", { name: "Unavailable model" }), diff --git a/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.tsx b/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.tsx index 20181508618..ab305ef67b1 100644 --- a/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.tsx +++ b/site/src/pages/AISettingsPage/CoderAgentsPage/CoderAgentsPageView.tsx @@ -33,6 +33,7 @@ export interface CoderAgentsPageViewProps { isSaveAdminOverridesError: boolean; generalModelOverrideData?: TypesGen.ChatModelOverrideResponse; titleGenerationModelOverrideData?: TypesGen.ChatModelOverrideResponse; + compactionModelOverrideData?: TypesGen.ChatModelOverrideResponse; exploreModelOverrideData?: TypesGen.ChatModelOverrideResponse; modelConfigsData: TypesGen.ChatModelConfig[] | undefined; providerInfoByID: ReadonlyMap; @@ -45,6 +46,9 @@ export interface CoderAgentsPageViewProps { onSaveTitleGenerationModel: SaveModelOverride; isSavingTitleGenerationModel: boolean; isSaveTitleGenerationModelError: boolean; + onSaveCompactionModel: SaveModelOverride; + isSavingCompactionModel: boolean; + isSaveCompactionModelError: boolean; onSaveExploreModelOverride: SaveModelOverride; isSavingExploreModelOverride: boolean; isSaveExploreModelOverrideError: boolean; @@ -83,6 +87,7 @@ export const CoderAgentsPageView: FC = ({ isSaveAdminOverridesError, generalModelOverrideData, titleGenerationModelOverrideData, + compactionModelOverrideData, exploreModelOverrideData, modelConfigsData, providerInfoByID, @@ -95,6 +100,9 @@ export const CoderAgentsPageView: FC = ({ onSaveTitleGenerationModel, isSavingTitleGenerationModel, isSaveTitleGenerationModelError, + onSaveCompactionModel, + isSavingCompactionModel, + isSaveCompactionModelError, onSaveExploreModelOverride, isSavingExploreModelOverride, isSaveExploreModelOverrideError, @@ -172,6 +180,20 @@ export const CoderAgentsPageView: FC = ({ unsetPlaceholder="Use title default" unavailableModelWarning="The selected model is currently unavailable. Title generation will be skipped until you choose another model or clear this setting." /> + ( onSaveTitleGenerationModel={fn()} isSavingTitleGenerationModel={false} isSaveTitleGenerationModelError={false} + onSaveCompactionModel={fn()} + isSavingCompactionModel={false} + isSaveCompactionModelError={false} onSaveExploreModelOverride={fn()} isSavingExploreModelOverride={false} isSaveExploreModelOverrideError={false}