mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: per-user per-model chat compaction threshold overrides (#23412)
## What
Adds per-user per-model auto-compaction threshold overrides. Users can
now customize the percentage of context window usage that triggers chat
compaction, independently for each enabled model.
## Why
The compaction threshold was previously only configurable at the
deployment level (`chat_model_configs.compression_threshold`). Different
users have different preferences — some want aggressive compaction to
keep costs low, others prefer higher thresholds to retain more context.
This gives users control without requiring admin intervention.
## Architecture
**Storage:** Reuses the existing `user_configs` table (no migration
needed). Overrides are stored as key/value pairs with keys shaped
`chat_compaction_threshold:<modelConfigID>` and integer percent values.
**API:** Three new experimental endpoints under
`/api/experimental/chats/config/`:
- `GET /user-compaction-thresholds` — list all overrides for the current
user
- `PUT /user-compaction-thresholds/{modelConfig}` — upsert an override
(validates model exists and is enabled, validates 0–100 range)
- `DELETE /user-compaction-thresholds/{modelConfig}` — clear an override
(idempotent)
**Runtime resolution:** In `coderd/chatd/chatd.go`, a new
`resolveUserCompactionThreshold()` helper runs at the start of each chat
turn (inside `runChat()`), after the model config is resolved but before
`CompactionOptions` is built. If a valid override exists, it replaces
`modelConfig.CompressionThreshold`. The threshold source
(`user_override` vs `model_default`) is logged with each compaction
event.
**Precedence:** `effectiveThreshold = userOverride ??
modelConfig.CompressionThreshold`
**UI:** New "Context Compaction" subsection in the Agents → Settings →
Behavior tab, placed after Personal Instructions. Shows one row per
enabled model with the system default, a number input for the override,
and Save/Reset controls.
## Testing
- 9 API subtests covering CRUD, validation (boundary values 0/100,
out-of-range rejection), upsert behavior, idempotent delete, user
isolation, and non-existent model config
- 4 dbauthz tests (16 scenarios) verifying `ActionReadPersonal` /
`ActionUpdatePersonal` on all query methods
- 4 Storybook stories with play functions (Default, WithOverrides,
Loading, Error)
<details>
<summary>Implementation plan</summary>
### Phase 1 — Tests
- Backend API tests in `coderd/chats_test.go` (9 subtests)
- Database auth wrapper tests in
`coderd/database/dbauthz/dbauthz_test.go` (4 methods)
- Frontend stories in `UserCompactionThresholdSettings.stories.tsx` (4
stories)
### Phase 2 — Backend preference surface
- 4 SQL queries in `coderd/database/queries/users.sql` (list, get,
upsert, delete)
- `make gen` to propagate into generated artifacts
- Auth/metrics wrappers in dbauthz and dbmetrics
- SDK types and client methods in `codersdk/chats.go`
- HTTP handlers and routes in `coderd/chats.go` and `coderd/coderd.go`
- Key prefix constant shared between handlers and runtime
### Phase 3 — Runtime override
- `resolveUserCompactionThreshold()` helper in `coderd/chatd/chatd.go`
- Override injection in `runChat()` before building `CompactionOptions`
- `threshold_source` field added to compaction log
### Phase 4 — Settings UI
- API client methods and React Query hooks in `site/src/api/`
- `UserCompactionThresholdSettings` component extracted from
`SettingsPageContent`
- Per-model mutation tracking (only the active row disables during save)
- 100% warning, "System default" label, helpful empty state copy
### Phase 5 — Refactor and review fixes
- Consolidated key prefix constant in `codersdk`
- Explicit PUT range validation (not just struct tags)
- GET handler gracefully skips malformed rows instead of 500
- Boundary value, upsert, and non-existent model config tests
- UX improvements: per-model mutation state, aria-live on errors
</details>
This commit is contained in:
@@ -1180,6 +1180,9 @@ func New(options *Options) *API {
|
||||
r.Put("/desktop-enabled", api.putChatDesktopEnabled)
|
||||
r.Get("/user-prompt", api.getUserChatCustomPrompt)
|
||||
r.Put("/user-prompt", api.putUserChatCustomPrompt)
|
||||
r.Get("/user-compaction-thresholds", api.getUserChatCompactionThresholds)
|
||||
r.Put("/user-compaction-thresholds/{modelConfig}", api.putUserChatCompactionThreshold)
|
||||
r.Delete("/user-compaction-thresholds/{modelConfig}", api.deleteUserChatCompactionThreshold)
|
||||
r.Get("/workspace-ttl", api.getChatWorkspaceTTL)
|
||||
r.Put("/workspace-ttl", api.putChatWorkspaceTTL)
|
||||
})
|
||||
|
||||
@@ -2118,6 +2118,17 @@ func (q *querier) DeleteTask(ctx context.Context, arg database.DeleteTaskParams)
|
||||
return q.db.DeleteTask(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) DeleteUserChatCompactionThreshold(ctx context.Context, arg database.DeleteUserChatCompactionThresholdParams) error {
|
||||
u, err := q.db.GetUserByID(ctx, arg.UserID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := q.authorizeContext(ctx, policy.ActionUpdatePersonal, u); err != nil {
|
||||
return err
|
||||
}
|
||||
return q.db.DeleteUserChatCompactionThreshold(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) DeleteUserSecret(ctx context.Context, id uuid.UUID) error {
|
||||
// First get the secret to check ownership
|
||||
secret, err := q.GetUserSecret(ctx, id)
|
||||
@@ -3921,6 +3932,17 @@ func (q *querier) GetUserByID(ctx context.Context, id uuid.UUID) (database.User,
|
||||
return fetch(q.log, q.auth, q.db.GetUserByID)(ctx, id)
|
||||
}
|
||||
|
||||
func (q *querier) GetUserChatCompactionThreshold(ctx context.Context, arg database.GetUserChatCompactionThresholdParams) (string, error) {
|
||||
u, err := q.db.GetUserByID(ctx, arg.UserID)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := q.authorizeContext(ctx, policy.ActionReadPersonal, u); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return q.db.GetUserChatCompactionThreshold(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) GetUserChatCustomPrompt(ctx context.Context, userID uuid.UUID) (string, error) {
|
||||
u, err := q.db.GetUserByID(ctx, userID)
|
||||
if err != nil {
|
||||
@@ -5352,6 +5374,17 @@ func (q *querier) ListTasks(ctx context.Context, arg database.ListTasksParams) (
|
||||
return fetchWithPostFilter(q.auth, policy.ActionRead, q.db.ListTasks)(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) ListUserChatCompactionThresholds(ctx context.Context, userID uuid.UUID) ([]database.UserConfig, error) {
|
||||
u, err := q.db.GetUserByID(ctx, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := q.authorizeContext(ctx, policy.ActionReadPersonal, u); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return q.db.ListUserChatCompactionThresholds(ctx, userID)
|
||||
}
|
||||
|
||||
func (q *querier) ListUserSecrets(ctx context.Context, userID uuid.UUID) ([]database.UserSecret, error) {
|
||||
obj := rbac.ResourceUserSecret.WithOwner(userID.String())
|
||||
if err := q.authorizeContext(ctx, policy.ActionRead, obj); err != nil {
|
||||
@@ -6212,6 +6245,17 @@ func (q *querier) UpdateUsageEventsPostPublish(ctx context.Context, arg database
|
||||
return q.db.UpdateUsageEventsPostPublish(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) UpdateUserChatCompactionThreshold(ctx context.Context, arg database.UpdateUserChatCompactionThresholdParams) (database.UserConfig, error) {
|
||||
u, err := q.db.GetUserByID(ctx, arg.UserID)
|
||||
if err != nil {
|
||||
return database.UserConfig{}, err
|
||||
}
|
||||
if err := q.authorizeContext(ctx, policy.ActionUpdatePersonal, u); err != nil {
|
||||
return database.UserConfig{}, err
|
||||
}
|
||||
return q.db.UpdateUserChatCompactionThreshold(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) UpdateUserChatCustomPrompt(ctx context.Context, arg database.UpdateUserChatCustomPromptParams) (database.UserConfig, error) {
|
||||
u, err := q.db.GetUserByID(ctx, arg.UserID)
|
||||
if err != nil {
|
||||
|
||||
@@ -2278,6 +2278,35 @@ func (s *MethodTestSuite) TestUser() {
|
||||
dbm.EXPECT().UpdateUserChatCustomPrompt(gomock.Any(), arg).Return(uc, nil).AnyTimes()
|
||||
check.Args(arg).Asserts(u, policy.ActionUpdatePersonal).Returns(uc)
|
||||
}))
|
||||
s.Run("ListUserChatCompactionThresholds", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
u := testutil.Fake(s.T(), faker, database.User{})
|
||||
uc := database.UserConfig{UserID: u.ID, Key: codersdk.ChatCompactionThresholdKeyPrefix + "00000000-0000-0000-0000-000000000001", Value: "75"}
|
||||
dbm.EXPECT().GetUserByID(gomock.Any(), u.ID).Return(u, nil).AnyTimes()
|
||||
dbm.EXPECT().ListUserChatCompactionThresholds(gomock.Any(), u.ID).Return([]database.UserConfig{uc}, nil).AnyTimes()
|
||||
check.Args(u.ID).Asserts(u, policy.ActionReadPersonal).Returns([]database.UserConfig{uc})
|
||||
}))
|
||||
s.Run("GetUserChatCompactionThreshold", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
u := testutil.Fake(s.T(), faker, database.User{})
|
||||
arg := database.GetUserChatCompactionThresholdParams{UserID: u.ID, Key: codersdk.ChatCompactionThresholdKeyPrefix + "00000000-0000-0000-0000-000000000001"}
|
||||
dbm.EXPECT().GetUserByID(gomock.Any(), u.ID).Return(u, nil).AnyTimes()
|
||||
dbm.EXPECT().GetUserChatCompactionThreshold(gomock.Any(), arg).Return("75", nil).AnyTimes()
|
||||
check.Args(arg).Asserts(u, policy.ActionReadPersonal).Returns("75")
|
||||
}))
|
||||
s.Run("UpdateUserChatCompactionThreshold", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
u := testutil.Fake(s.T(), faker, database.User{})
|
||||
uc := database.UserConfig{UserID: u.ID, Key: codersdk.ChatCompactionThresholdKeyPrefix + "00000000-0000-0000-0000-000000000001", Value: "75"}
|
||||
arg := database.UpdateUserChatCompactionThresholdParams{UserID: u.ID, Key: uc.Key, ThresholdPercent: 75}
|
||||
dbm.EXPECT().GetUserByID(gomock.Any(), u.ID).Return(u, nil).AnyTimes()
|
||||
dbm.EXPECT().UpdateUserChatCompactionThreshold(gomock.Any(), arg).Return(uc, nil).AnyTimes()
|
||||
check.Args(arg).Asserts(u, policy.ActionUpdatePersonal).Returns(uc)
|
||||
}))
|
||||
s.Run("DeleteUserChatCompactionThreshold", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
u := testutil.Fake(s.T(), faker, database.User{})
|
||||
arg := database.DeleteUserChatCompactionThresholdParams{UserID: u.ID, Key: codersdk.ChatCompactionThresholdKeyPrefix + "00000000-0000-0000-0000-000000000001"}
|
||||
dbm.EXPECT().GetUserByID(gomock.Any(), u.ID).Return(u, nil).AnyTimes()
|
||||
dbm.EXPECT().DeleteUserChatCompactionThreshold(gomock.Any(), arg).Return(nil).AnyTimes()
|
||||
check.Args(arg).Asserts(u, policy.ActionUpdatePersonal)
|
||||
}))
|
||||
s.Run("UpdateUserTaskNotificationAlertDismissed", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
user := testutil.Fake(s.T(), faker, database.User{})
|
||||
userConfig := database.UserConfig{UserID: user.ID, Key: "task_notification_alert_dismissed", Value: "false"}
|
||||
|
||||
@@ -680,6 +680,14 @@ func (m queryMetricsStore) DeleteTask(ctx context.Context, arg database.DeleteTa
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) DeleteUserChatCompactionThreshold(ctx context.Context, arg database.DeleteUserChatCompactionThresholdParams) error {
|
||||
start := time.Now()
|
||||
r0 := m.s.DeleteUserChatCompactionThreshold(ctx, arg)
|
||||
m.queryLatencies.WithLabelValues("DeleteUserChatCompactionThreshold").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "DeleteUserChatCompactionThreshold").Inc()
|
||||
return r0
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) DeleteUserSecret(ctx context.Context, id uuid.UUID) error {
|
||||
start := time.Now()
|
||||
r0 := m.s.DeleteUserSecret(ctx, id)
|
||||
@@ -2448,6 +2456,14 @@ func (m queryMetricsStore) GetUserByID(ctx context.Context, id uuid.UUID) (datab
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetUserChatCompactionThreshold(ctx context.Context, arg database.GetUserChatCompactionThresholdParams) (string, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetUserChatCompactionThreshold(ctx, arg)
|
||||
m.queryLatencies.WithLabelValues("GetUserChatCompactionThreshold").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetUserChatCompactionThreshold").Inc()
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetUserChatCustomPrompt(ctx context.Context, userID uuid.UUID) (string, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetUserChatCustomPrompt(ctx, userID)
|
||||
@@ -3768,6 +3784,14 @@ func (m queryMetricsStore) ListTasks(ctx context.Context, arg database.ListTasks
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) ListUserChatCompactionThresholds(ctx context.Context, userID uuid.UUID) ([]database.UserConfig, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.ListUserChatCompactionThresholds(ctx, userID)
|
||||
m.queryLatencies.WithLabelValues("ListUserChatCompactionThresholds").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "ListUserChatCompactionThresholds").Inc()
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) ListUserSecrets(ctx context.Context, userID uuid.UUID) ([]database.UserSecret, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.ListUserSecrets(ctx, userID)
|
||||
@@ -4360,6 +4384,14 @@ func (m queryMetricsStore) UpdateUsageEventsPostPublish(ctx context.Context, arg
|
||||
return r0
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) UpdateUserChatCompactionThreshold(ctx context.Context, arg database.UpdateUserChatCompactionThresholdParams) (database.UserConfig, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.UpdateUserChatCompactionThreshold(ctx, arg)
|
||||
m.queryLatencies.WithLabelValues("UpdateUserChatCompactionThreshold").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "UpdateUserChatCompactionThreshold").Inc()
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) UpdateUserChatCustomPrompt(ctx context.Context, arg database.UpdateUserChatCustomPromptParams) (database.UserConfig, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.UpdateUserChatCustomPrompt(ctx, arg)
|
||||
|
||||
@@ -1126,6 +1126,20 @@ func (mr *MockStoreMockRecorder) DeleteTask(ctx, arg any) *gomock.Call {
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteTask", reflect.TypeOf((*MockStore)(nil).DeleteTask), ctx, arg)
|
||||
}
|
||||
|
||||
// DeleteUserChatCompactionThreshold mocks base method.
|
||||
func (m *MockStore) DeleteUserChatCompactionThreshold(ctx context.Context, arg database.DeleteUserChatCompactionThresholdParams) error {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "DeleteUserChatCompactionThreshold", ctx, arg)
|
||||
ret0, _ := ret[0].(error)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// DeleteUserChatCompactionThreshold indicates an expected call of DeleteUserChatCompactionThreshold.
|
||||
func (mr *MockStoreMockRecorder) DeleteUserChatCompactionThreshold(ctx, arg any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteUserChatCompactionThreshold", reflect.TypeOf((*MockStore)(nil).DeleteUserChatCompactionThreshold), ctx, arg)
|
||||
}
|
||||
|
||||
// DeleteUserSecret mocks base method.
|
||||
func (m *MockStore) DeleteUserSecret(ctx context.Context, id uuid.UUID) error {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -4564,6 +4578,21 @@ func (mr *MockStoreMockRecorder) GetUserByID(ctx, id any) *gomock.Call {
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetUserByID", reflect.TypeOf((*MockStore)(nil).GetUserByID), ctx, id)
|
||||
}
|
||||
|
||||
// GetUserChatCompactionThreshold mocks base method.
|
||||
func (m *MockStore) GetUserChatCompactionThreshold(ctx context.Context, arg database.GetUserChatCompactionThresholdParams) (string, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetUserChatCompactionThreshold", ctx, arg)
|
||||
ret0, _ := ret[0].(string)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetUserChatCompactionThreshold indicates an expected call of GetUserChatCompactionThreshold.
|
||||
func (mr *MockStoreMockRecorder) GetUserChatCompactionThreshold(ctx, arg any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetUserChatCompactionThreshold", reflect.TypeOf((*MockStore)(nil).GetUserChatCompactionThreshold), ctx, arg)
|
||||
}
|
||||
|
||||
// GetUserChatCustomPrompt mocks base method.
|
||||
func (m *MockStore) GetUserChatCustomPrompt(ctx context.Context, userID uuid.UUID) (string, error) {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -7068,6 +7097,21 @@ func (mr *MockStoreMockRecorder) ListTasks(ctx, arg any) *gomock.Call {
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ListTasks", reflect.TypeOf((*MockStore)(nil).ListTasks), ctx, arg)
|
||||
}
|
||||
|
||||
// ListUserChatCompactionThresholds mocks base method.
|
||||
func (m *MockStore) ListUserChatCompactionThresholds(ctx context.Context, userID uuid.UUID) ([]database.UserConfig, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "ListUserChatCompactionThresholds", ctx, userID)
|
||||
ret0, _ := ret[0].([]database.UserConfig)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// ListUserChatCompactionThresholds indicates an expected call of ListUserChatCompactionThresholds.
|
||||
func (mr *MockStoreMockRecorder) ListUserChatCompactionThresholds(ctx, userID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ListUserChatCompactionThresholds", reflect.TypeOf((*MockStore)(nil).ListUserChatCompactionThresholds), ctx, userID)
|
||||
}
|
||||
|
||||
// ListUserSecrets mocks base method.
|
||||
func (m *MockStore) ListUserSecrets(ctx context.Context, userID uuid.UUID) ([]database.UserSecret, error) {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -8173,6 +8217,21 @@ func (mr *MockStoreMockRecorder) UpdateUsageEventsPostPublish(ctx, arg any) *gom
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateUsageEventsPostPublish", reflect.TypeOf((*MockStore)(nil).UpdateUsageEventsPostPublish), ctx, arg)
|
||||
}
|
||||
|
||||
// UpdateUserChatCompactionThreshold mocks base method.
|
||||
func (m *MockStore) UpdateUserChatCompactionThreshold(ctx context.Context, arg database.UpdateUserChatCompactionThresholdParams) (database.UserConfig, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "UpdateUserChatCompactionThreshold", ctx, arg)
|
||||
ret0, _ := ret[0].(database.UserConfig)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// UpdateUserChatCompactionThreshold indicates an expected call of UpdateUserChatCompactionThreshold.
|
||||
func (mr *MockStoreMockRecorder) UpdateUserChatCompactionThreshold(ctx, arg any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateUserChatCompactionThreshold", reflect.TypeOf((*MockStore)(nil).UpdateUserChatCompactionThreshold), ctx, arg)
|
||||
}
|
||||
|
||||
// UpdateUserChatCustomPrompt mocks base method.
|
||||
func (m *MockStore) UpdateUserChatCustomPrompt(ctx context.Context, arg database.UpdateUserChatCustomPromptParams) (database.UserConfig, error) {
|
||||
m.ctrl.T.Helper()
|
||||
|
||||
@@ -148,6 +148,7 @@ type sqlcQuerier interface {
|
||||
DeleteTailnetPeer(ctx context.Context, arg DeleteTailnetPeerParams) (DeleteTailnetPeerRow, error)
|
||||
DeleteTailnetTunnel(ctx context.Context, arg DeleteTailnetTunnelParams) (DeleteTailnetTunnelRow, error)
|
||||
DeleteTask(ctx context.Context, arg DeleteTaskParams) (uuid.UUID, error)
|
||||
DeleteUserChatCompactionThreshold(ctx context.Context, arg DeleteUserChatCompactionThresholdParams) error
|
||||
DeleteUserSecret(ctx context.Context, id uuid.UUID) error
|
||||
DeleteWebpushSubscriptionByUserIDAndEndpoint(ctx context.Context, arg DeleteWebpushSubscriptionByUserIDAndEndpointParams) error
|
||||
DeleteWebpushSubscriptions(ctx context.Context, ids []uuid.UUID) error
|
||||
@@ -553,6 +554,7 @@ type sqlcQuerier interface {
|
||||
GetUserActivityInsights(ctx context.Context, arg GetUserActivityInsightsParams) ([]GetUserActivityInsightsRow, error)
|
||||
GetUserByEmailOrUsername(ctx context.Context, arg GetUserByEmailOrUsernameParams) (User, error)
|
||||
GetUserByID(ctx context.Context, id uuid.UUID) (User, error)
|
||||
GetUserChatCompactionThreshold(ctx context.Context, arg GetUserChatCompactionThresholdParams) (string, error)
|
||||
GetUserChatCustomPrompt(ctx context.Context, userID uuid.UUID) (string, error)
|
||||
GetUserChatSpendInPeriod(ctx context.Context, arg GetUserChatSpendInPeriodParams) (int64, error)
|
||||
GetUserCount(ctx context.Context, includeSystem bool) (int64, error)
|
||||
@@ -765,6 +767,7 @@ type sqlcQuerier interface {
|
||||
ListProvisionerKeysByOrganization(ctx context.Context, organizationID uuid.UUID) ([]ProvisionerKey, error)
|
||||
ListProvisionerKeysByOrganizationExcludeReserved(ctx context.Context, organizationID uuid.UUID) ([]ProvisionerKey, error)
|
||||
ListTasks(ctx context.Context, arg ListTasksParams) ([]Task, error)
|
||||
ListUserChatCompactionThresholds(ctx context.Context, userID uuid.UUID) ([]UserConfig, error)
|
||||
ListUserSecrets(ctx context.Context, userID uuid.UUID) ([]UserSecret, error)
|
||||
ListWorkspaceAgentPortShares(ctx context.Context, workspaceID uuid.UUID) ([]WorkspaceAgentPortShare, error)
|
||||
MarkAllInboxNotificationsAsRead(ctx context.Context, arg MarkAllInboxNotificationsAsReadParams) error
|
||||
@@ -868,6 +871,7 @@ type sqlcQuerier interface {
|
||||
UpdateTemplateVersionFlagsByJobID(ctx context.Context, arg UpdateTemplateVersionFlagsByJobIDParams) error
|
||||
UpdateTemplateWorkspacesLastUsedAt(ctx context.Context, arg UpdateTemplateWorkspacesLastUsedAtParams) error
|
||||
UpdateUsageEventsPostPublish(ctx context.Context, arg UpdateUsageEventsPostPublishParams) error
|
||||
UpdateUserChatCompactionThreshold(ctx context.Context, arg UpdateUserChatCompactionThresholdParams) (UserConfig, error)
|
||||
UpdateUserChatCustomPrompt(ctx context.Context, arg UpdateUserChatCustomPromptParams) (UserConfig, error)
|
||||
UpdateUserDeletedByID(ctx context.Context, id uuid.UUID) error
|
||||
UpdateUserGithubComUserID(ctx context.Context, arg UpdateUserGithubComUserIDParams) error
|
||||
|
||||
@@ -21157,6 +21157,20 @@ func (q *sqlQuerier) AllUserIDs(ctx context.Context, includeSystem bool) ([]uuid
|
||||
return items, nil
|
||||
}
|
||||
|
||||
const deleteUserChatCompactionThreshold = `-- name: DeleteUserChatCompactionThreshold :exec
|
||||
DELETE FROM user_configs WHERE user_id = $1 AND key = $2
|
||||
`
|
||||
|
||||
type DeleteUserChatCompactionThresholdParams struct {
|
||||
UserID uuid.UUID `db:"user_id" json:"user_id"`
|
||||
Key string `db:"key" json:"key"`
|
||||
}
|
||||
|
||||
func (q *sqlQuerier) DeleteUserChatCompactionThreshold(ctx context.Context, arg DeleteUserChatCompactionThresholdParams) error {
|
||||
_, err := q.db.ExecContext(ctx, deleteUserChatCompactionThreshold, arg.UserID, arg.Key)
|
||||
return err
|
||||
}
|
||||
|
||||
const getActiveUserCount = `-- name: GetActiveUserCount :one
|
||||
SELECT
|
||||
COUNT(*)
|
||||
@@ -21337,6 +21351,23 @@ func (q *sqlQuerier) GetUserByID(ctx context.Context, id uuid.UUID) (User, error
|
||||
return i, err
|
||||
}
|
||||
|
||||
const getUserChatCompactionThreshold = `-- name: GetUserChatCompactionThreshold :one
|
||||
SELECT value AS threshold_percent FROM user_configs
|
||||
WHERE user_id = $1 AND key = $2
|
||||
`
|
||||
|
||||
type GetUserChatCompactionThresholdParams struct {
|
||||
UserID uuid.UUID `db:"user_id" json:"user_id"`
|
||||
Key string `db:"key" json:"key"`
|
||||
}
|
||||
|
||||
func (q *sqlQuerier) GetUserChatCompactionThreshold(ctx context.Context, arg GetUserChatCompactionThresholdParams) (string, error) {
|
||||
row := q.db.QueryRowContext(ctx, getUserChatCompactionThreshold, arg.UserID, arg.Key)
|
||||
var threshold_percent string
|
||||
err := row.Scan(&threshold_percent)
|
||||
return threshold_percent, err
|
||||
}
|
||||
|
||||
const getUserChatCustomPrompt = `-- name: GetUserChatCustomPrompt :one
|
||||
SELECT
|
||||
value as chat_custom_prompt
|
||||
@@ -21760,6 +21791,36 @@ func (q *sqlQuerier) InsertUser(ctx context.Context, arg InsertUserParams) (User
|
||||
return i, err
|
||||
}
|
||||
|
||||
const listUserChatCompactionThresholds = `-- name: ListUserChatCompactionThresholds :many
|
||||
SELECT user_id, key, value FROM user_configs
|
||||
WHERE user_id = $1
|
||||
AND key LIKE 'chat\_compaction\_threshold\_pct:%'
|
||||
ORDER BY key
|
||||
`
|
||||
|
||||
func (q *sqlQuerier) ListUserChatCompactionThresholds(ctx context.Context, userID uuid.UUID) ([]UserConfig, error) {
|
||||
rows, err := q.db.QueryContext(ctx, listUserChatCompactionThresholds, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var items []UserConfig
|
||||
for rows.Next() {
|
||||
var i UserConfig
|
||||
if err := rows.Scan(&i.UserID, &i.Key, &i.Value); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items = append(items, i)
|
||||
}
|
||||
if err := rows.Close(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
const updateInactiveUsersToDormant = `-- name: UpdateInactiveUsersToDormant :many
|
||||
UPDATE
|
||||
users
|
||||
@@ -21813,6 +21874,27 @@ func (q *sqlQuerier) UpdateInactiveUsersToDormant(ctx context.Context, arg Updat
|
||||
return items, nil
|
||||
}
|
||||
|
||||
const updateUserChatCompactionThreshold = `-- name: UpdateUserChatCompactionThreshold :one
|
||||
INSERT INTO user_configs (user_id, key, value)
|
||||
VALUES ($1, $2, ($3::int)::text)
|
||||
ON CONFLICT ON CONSTRAINT user_configs_pkey
|
||||
DO UPDATE SET value = ($3::int)::text
|
||||
RETURNING user_id, key, value
|
||||
`
|
||||
|
||||
type UpdateUserChatCompactionThresholdParams struct {
|
||||
UserID uuid.UUID `db:"user_id" json:"user_id"`
|
||||
Key string `db:"key" json:"key"`
|
||||
ThresholdPercent int32 `db:"threshold_percent" json:"threshold_percent"`
|
||||
}
|
||||
|
||||
func (q *sqlQuerier) UpdateUserChatCompactionThreshold(ctx context.Context, arg UpdateUserChatCompactionThresholdParams) (UserConfig, error) {
|
||||
row := q.db.QueryRowContext(ctx, updateUserChatCompactionThreshold, arg.UserID, arg.Key, arg.ThresholdPercent)
|
||||
var i UserConfig
|
||||
err := row.Scan(&i.UserID, &i.Key, &i.Value)
|
||||
return i, err
|
||||
}
|
||||
|
||||
const updateUserChatCustomPrompt = `-- name: UpdateUserChatCustomPrompt :one
|
||||
INSERT INTO
|
||||
user_configs (user_id, key, value)
|
||||
|
||||
@@ -193,6 +193,26 @@ WHERE user_configs.user_id = @user_id
|
||||
AND user_configs.key = 'chat_custom_prompt'
|
||||
RETURNING *;
|
||||
|
||||
-- name: ListUserChatCompactionThresholds :many
|
||||
SELECT user_id, key, value FROM user_configs
|
||||
WHERE user_id = @user_id
|
||||
AND key LIKE 'chat\_compaction\_threshold\_pct:%'
|
||||
ORDER BY key;
|
||||
|
||||
-- name: GetUserChatCompactionThreshold :one
|
||||
SELECT value AS threshold_percent FROM user_configs
|
||||
WHERE user_id = @user_id AND key = @key;
|
||||
|
||||
-- name: UpdateUserChatCompactionThreshold :one
|
||||
INSERT INTO user_configs (user_id, key, value)
|
||||
VALUES (@user_id, @key, (@threshold_percent::int)::text)
|
||||
ON CONFLICT ON CONSTRAINT user_configs_pkey
|
||||
DO UPDATE SET value = (@threshold_percent::int)::text
|
||||
RETURNING *;
|
||||
|
||||
-- name: DeleteUserChatCompactionThreshold :exec
|
||||
DELETE FROM user_configs WHERE user_id = @user_id AND key = @key;
|
||||
|
||||
-- name: GetUserTaskNotificationAlertDismissed :one
|
||||
SELECT
|
||||
value::boolean as task_notification_alert_dismissed
|
||||
|
||||
@@ -2542,6 +2542,17 @@ func normalizeChatCompressionThreshold(
|
||||
return threshold, nil
|
||||
}
|
||||
|
||||
func parseCompactionThresholdKey(key string) (uuid.UUID, error) {
|
||||
if !strings.HasPrefix(key, codersdk.ChatCompactionThresholdKeyPrefix) {
|
||||
return uuid.Nil, xerrors.Errorf("invalid compaction threshold key: %q", key)
|
||||
}
|
||||
id, err := uuid.Parse(key[len(codersdk.ChatCompactionThresholdKeyPrefix):])
|
||||
if err != nil {
|
||||
return uuid.Nil, xerrors.Errorf("invalid model config ID in key %q: %w", key, err)
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
|
||||
const (
|
||||
// maxChatFileSize is the maximum size of a chat file upload (10 MB).
|
||||
maxChatFileSize = 10 << 20
|
||||
@@ -2816,6 +2827,170 @@ func (api *API) putUserChatCustomPrompt(rw http.ResponseWriter, r *http.Request)
|
||||
})
|
||||
}
|
||||
|
||||
// @Summary Get user chat compaction thresholds
|
||||
// @x-apidocgen {"skip": true}
|
||||
// EXPERIMENTAL: this endpoint is experimental and is subject to change.
|
||||
//
|
||||
//nolint:revive // get-return: revive assumes get* must be a getter, but this is an HTTP handler.
|
||||
func (api *API) getUserChatCompactionThresholds(rw http.ResponseWriter, r *http.Request) {
|
||||
var (
|
||||
ctx = r.Context()
|
||||
apiKey = httpmw.APIKey(r)
|
||||
)
|
||||
|
||||
rows, err := api.Database.ListUserChatCompactionThresholds(ctx, apiKey.UserID)
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Error listing user chat compaction thresholds.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
resp := codersdk.UserChatCompactionThresholds{
|
||||
Thresholds: make([]codersdk.UserChatCompactionThreshold, 0, len(rows)),
|
||||
}
|
||||
for _, row := range rows {
|
||||
modelConfigID, err := parseCompactionThresholdKey(row.Key)
|
||||
if err != nil {
|
||||
api.Logger.Warn(ctx, "skipping malformed user chat compaction threshold key",
|
||||
slog.F("key", row.Key),
|
||||
slog.F("value", row.Value),
|
||||
slog.Error(err),
|
||||
)
|
||||
continue
|
||||
}
|
||||
|
||||
thresholdPercent, err := strconv.ParseInt(row.Value, 10, 32)
|
||||
if err != nil {
|
||||
api.Logger.Warn(ctx, "skipping malformed user chat compaction threshold value",
|
||||
slog.F("key", row.Key),
|
||||
slog.F("value", row.Value),
|
||||
slog.Error(err),
|
||||
)
|
||||
continue
|
||||
}
|
||||
if thresholdPercent < int64(minChatContextCompressionThreshold) ||
|
||||
thresholdPercent > int64(maxChatContextCompressionThreshold) {
|
||||
api.Logger.Warn(ctx, "skipping out-of-range user chat compaction threshold",
|
||||
slog.F("key", row.Key),
|
||||
slog.F("value", row.Value),
|
||||
)
|
||||
continue
|
||||
}
|
||||
|
||||
resp.Thresholds = append(resp.Thresholds, codersdk.UserChatCompactionThreshold{
|
||||
ModelConfigID: modelConfigID,
|
||||
ThresholdPercent: int32(thresholdPercent),
|
||||
})
|
||||
}
|
||||
|
||||
httpapi.Write(ctx, rw, http.StatusOK, resp)
|
||||
}
|
||||
|
||||
// @Summary Set user chat compaction threshold for a model config
|
||||
// @x-apidocgen {"skip": true}
|
||||
// EXPERIMENTAL: this endpoint is experimental and is subject to change.
|
||||
func (api *API) putUserChatCompactionThreshold(rw http.ResponseWriter, r *http.Request) {
|
||||
var (
|
||||
ctx = r.Context()
|
||||
apiKey = httpmw.APIKey(r)
|
||||
)
|
||||
|
||||
modelConfigID, ok := parseChatModelConfigID(rw, r)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
var req codersdk.UpdateUserChatCompactionThresholdRequest
|
||||
if !httpapi.Read(ctx, rw, r, &req) {
|
||||
return
|
||||
}
|
||||
if req.ThresholdPercent < minChatContextCompressionThreshold ||
|
||||
req.ThresholdPercent > maxChatContextCompressionThreshold {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "threshold_percent is out of range.",
|
||||
Detail: fmt.Sprintf(
|
||||
"threshold_percent must be between %d and %d, got %d.",
|
||||
minChatContextCompressionThreshold,
|
||||
maxChatContextCompressionThreshold,
|
||||
req.ThresholdPercent,
|
||||
),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Use system context because GetChatModelConfigByID requires
|
||||
// deployment-config read access, which non-admin users lack.
|
||||
// The user is only checking if the model exists and is enabled
|
||||
// before writing their own personal preference.
|
||||
//nolint:gocritic // Non-admin users need this lookup to save their own setting.
|
||||
modelConfig, err := api.Database.GetChatModelConfigByID(dbauthz.AsSystemRestricted(ctx), modelConfigID)
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) || httpapi.Is404Error(err) {
|
||||
httpapi.ResourceNotFound(rw)
|
||||
return
|
||||
}
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Failed to get chat model config.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
if !modelConfig.Enabled {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "Model config is disabled.",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
_, err = api.Database.UpdateUserChatCompactionThreshold(ctx, database.UpdateUserChatCompactionThresholdParams{
|
||||
UserID: apiKey.UserID,
|
||||
Key: codersdk.CompactionThresholdKey(modelConfigID),
|
||||
ThresholdPercent: req.ThresholdPercent,
|
||||
})
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Error updating user chat compaction threshold.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
httpapi.Write(ctx, rw, http.StatusOK, codersdk.UserChatCompactionThreshold{
|
||||
ModelConfigID: modelConfigID,
|
||||
ThresholdPercent: req.ThresholdPercent,
|
||||
})
|
||||
}
|
||||
|
||||
// @Summary Delete user chat compaction threshold for a model config
|
||||
// @x-apidocgen {"skip": true}
|
||||
// EXPERIMENTAL: this endpoint is experimental and is subject to change.
|
||||
func (api *API) deleteUserChatCompactionThreshold(rw http.ResponseWriter, r *http.Request) {
|
||||
var (
|
||||
ctx = r.Context()
|
||||
apiKey = httpmw.APIKey(r)
|
||||
)
|
||||
|
||||
modelConfigID, ok := parseChatModelConfigID(rw, r)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
if err := api.Database.DeleteUserChatCompactionThreshold(ctx, database.DeleteUserChatCompactionThresholdParams{
|
||||
UserID: apiKey.UserID,
|
||||
Key: codersdk.CompactionThresholdKey(modelConfigID),
|
||||
}); err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Error deleting user chat compaction threshold.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
rw.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
func (api *API) resolvedChatSystemPrompt(ctx context.Context) string {
|
||||
custom, err := api.Database.GetChatSystemPrompt(ctx)
|
||||
if err != nil {
|
||||
|
||||
@@ -4950,6 +4950,152 @@ func TestChatWorkspaceTTL(t *testing.T) {
|
||||
requireSDKError(t, err, http.StatusBadRequest)
|
||||
}
|
||||
|
||||
//nolint:tparallel,paralleltest // Subtests share a single coderdtest instance.
|
||||
func TestUserChatCompactionThresholds(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
client, _ := newChatClientWithDatabase(t)
|
||||
firstUser := coderdtest.CreateFirstUser(t, client.Client)
|
||||
modelConfig := createChatModelConfig(t, client)
|
||||
|
||||
t.Run("EmptyByDefault", func(t *testing.T) {
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
thresholds, err := client.GetUserChatCompactionThresholds(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, thresholds.Thresholds)
|
||||
})
|
||||
|
||||
t.Run("PutAndGet", func(t *testing.T) {
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
override, err := client.UpdateUserChatCompactionThreshold(ctx, modelConfig.ID, codersdk.UpdateUserChatCompactionThresholdRequest{
|
||||
ThresholdPercent: 75,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, modelConfig.ID, override.ModelConfigID)
|
||||
require.EqualValues(t, 75, override.ThresholdPercent)
|
||||
|
||||
thresholds, err := client.GetUserChatCompactionThresholds(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, thresholds.Thresholds, 1)
|
||||
require.Equal(t, modelConfig.ID, thresholds.Thresholds[0].ModelConfigID)
|
||||
require.EqualValues(t, 75, thresholds.Thresholds[0].ThresholdPercent)
|
||||
})
|
||||
|
||||
t.Run("UpsertChangesValue", func(t *testing.T) {
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
_, err := client.UpdateUserChatCompactionThreshold(ctx, modelConfig.ID, codersdk.UpdateUserChatCompactionThresholdRequest{
|
||||
ThresholdPercent: 50,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
override, err := client.UpdateUserChatCompactionThreshold(ctx, modelConfig.ID, codersdk.UpdateUserChatCompactionThresholdRequest{
|
||||
ThresholdPercent: 75,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 75, override.ThresholdPercent)
|
||||
|
||||
thresholds, err := client.GetUserChatCompactionThresholds(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, thresholds.Thresholds, 1)
|
||||
require.EqualValues(t, 75, thresholds.Thresholds[0].ThresholdPercent)
|
||||
})
|
||||
|
||||
t.Run("BoundaryValues", func(t *testing.T) {
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
override, err := client.UpdateUserChatCompactionThreshold(ctx, modelConfig.ID, codersdk.UpdateUserChatCompactionThresholdRequest{
|
||||
ThresholdPercent: 0,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 0, override.ThresholdPercent)
|
||||
|
||||
thresholds, err := client.GetUserChatCompactionThresholds(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, thresholds.Thresholds, 1)
|
||||
require.EqualValues(t, 0, thresholds.Thresholds[0].ThresholdPercent)
|
||||
|
||||
override, err = client.UpdateUserChatCompactionThreshold(ctx, modelConfig.ID, codersdk.UpdateUserChatCompactionThresholdRequest{
|
||||
ThresholdPercent: 100,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.EqualValues(t, 100, override.ThresholdPercent)
|
||||
|
||||
thresholds, err = client.GetUserChatCompactionThresholds(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, thresholds.Thresholds, 1)
|
||||
require.EqualValues(t, 100, thresholds.Thresholds[0].ThresholdPercent)
|
||||
})
|
||||
|
||||
t.Run("ValidationRejectsInvalid", func(t *testing.T) {
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
_, err := client.UpdateUserChatCompactionThreshold(ctx, modelConfig.ID, codersdk.UpdateUserChatCompactionThresholdRequest{
|
||||
ThresholdPercent: -1,
|
||||
})
|
||||
requireSDKError(t, err, http.StatusBadRequest)
|
||||
|
||||
_, err = client.UpdateUserChatCompactionThreshold(ctx, modelConfig.ID, codersdk.UpdateUserChatCompactionThresholdRequest{
|
||||
ThresholdPercent: 101,
|
||||
})
|
||||
requireSDKError(t, err, http.StatusBadRequest)
|
||||
})
|
||||
|
||||
t.Run("Delete", func(t *testing.T) {
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
err := client.DeleteUserChatCompactionThreshold(ctx, modelConfig.ID)
|
||||
require.NoError(t, err)
|
||||
|
||||
thresholds, err := client.GetUserChatCompactionThresholds(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, thresholds.Thresholds)
|
||||
})
|
||||
|
||||
t.Run("DeleteIdempotent", func(t *testing.T) {
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
err := client.DeleteUserChatCompactionThreshold(ctx, modelConfig.ID)
|
||||
require.NoError(t, err)
|
||||
})
|
||||
|
||||
t.Run("NonExistentModelConfig", func(t *testing.T) {
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
fakeID := uuid.New()
|
||||
_, err := client.UpdateUserChatCompactionThreshold(ctx, fakeID, codersdk.UpdateUserChatCompactionThresholdRequest{
|
||||
ThresholdPercent: 50,
|
||||
})
|
||||
requireSDKError(t, err, http.StatusNotFound)
|
||||
})
|
||||
|
||||
t.Run("IsolatedPerUser", func(t *testing.T) {
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
memberClientRaw, _ := coderdtest.CreateAnotherUser(t, client.Client, firstUser.OrganizationID)
|
||||
memberClient := codersdk.NewExperimentalClient(memberClientRaw)
|
||||
|
||||
override, err := client.UpdateUserChatCompactionThreshold(ctx, modelConfig.ID, codersdk.UpdateUserChatCompactionThresholdRequest{
|
||||
ThresholdPercent: 75,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, modelConfig.ID, override.ModelConfigID)
|
||||
require.EqualValues(t, 75, override.ThresholdPercent)
|
||||
|
||||
adminThresholds, err := client.GetUserChatCompactionThresholds(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, adminThresholds.Thresholds, 1)
|
||||
require.Equal(t, modelConfig.ID, adminThresholds.Thresholds[0].ModelConfigID)
|
||||
require.EqualValues(t, 75, adminThresholds.Thresholds[0].ThresholdPercent)
|
||||
|
||||
memberThresholds, err := memberClient.GetUserChatCompactionThresholds(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, memberThresholds.Thresholds)
|
||||
})
|
||||
}
|
||||
|
||||
func requireSDKError(t *testing.T, err error, expectedStatus int) *codersdk.Error {
|
||||
t.Helper()
|
||||
|
||||
|
||||
+37
-1
@@ -7,6 +7,7 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -3150,8 +3151,14 @@ func (p *Server) runChat(
|
||||
// "Summarizing..." tool call with the "Summarized" tool
|
||||
// result.
|
||||
compactionToolCallID := "chat_summarized_" + uuid.NewString()
|
||||
effectiveThreshold := modelConfig.CompressionThreshold
|
||||
thresholdSource := "model_default"
|
||||
if override, ok := p.resolveUserCompactionThreshold(ctx, chat.OwnerID, modelConfig.ID); ok {
|
||||
effectiveThreshold = override
|
||||
thresholdSource = "user_override"
|
||||
}
|
||||
compactionOptions := &chatloop.CompactionOptions{
|
||||
ThresholdPercent: modelConfig.CompressionThreshold,
|
||||
ThresholdPercent: effectiveThreshold,
|
||||
ContextLimit: modelConfig.ContextLimit,
|
||||
Persist: func(
|
||||
persistCtx context.Context,
|
||||
@@ -3168,6 +3175,7 @@ func (p *Server) runChat(
|
||||
}
|
||||
logger.Info(persistCtx, "chat context summarized",
|
||||
slog.F("chat_id", chat.ID),
|
||||
slog.F("threshold_source", thresholdSource),
|
||||
slog.F("threshold_percent", result.ThresholdPercent),
|
||||
slog.F("usage_percent", result.UsagePercent),
|
||||
slog.F("context_tokens", result.ContextTokens),
|
||||
@@ -3718,6 +3726,34 @@ func (p *Server) resolveInstructions(
|
||||
return instruction
|
||||
}
|
||||
|
||||
// resolveUserCompactionThreshold looks up the user's per-model
|
||||
// compaction threshold override. Returns the override value and
|
||||
// true if one exists and is valid, or 0 and false otherwise.
|
||||
func (p *Server) resolveUserCompactionThreshold(ctx context.Context, userID uuid.UUID, modelConfigID uuid.UUID) (int32, bool) {
|
||||
raw, err := p.db.GetUserChatCompactionThreshold(ctx, database.GetUserChatCompactionThresholdParams{
|
||||
UserID: userID,
|
||||
Key: codersdk.CompactionThresholdKey(modelConfigID),
|
||||
})
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return 0, false
|
||||
}
|
||||
if err != nil {
|
||||
p.logger.Warn(ctx, "failed to fetch compaction threshold override",
|
||||
slog.F("user_id", userID),
|
||||
slog.F("model_config_id", modelConfigID),
|
||||
slog.Error(err),
|
||||
)
|
||||
return 0, false
|
||||
}
|
||||
// Range 0..100 must stay in sync with handler validation in
|
||||
// coderd/chats.go.
|
||||
val, err := strconv.ParseInt(raw, 10, 32)
|
||||
if err != nil || val < 0 || val > 100 {
|
||||
return 0, false
|
||||
}
|
||||
return int32(val), true
|
||||
}
|
||||
|
||||
// resolveUserPrompt fetches the user's custom chat prompt from the
|
||||
// database and wraps it in <user-instructions> tags. Returns empty
|
||||
// string if no prompt is set.
|
||||
|
||||
@@ -2,6 +2,7 @@ package chatd
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -606,6 +607,85 @@ func TestPublishToStream_DropWarnRateLimiting(t *testing.T) {
|
||||
requireFieldValue(t, subWarn[2], "dropped_count", int64(1))
|
||||
}
|
||||
|
||||
func TestResolveUserCompactionThreshold(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
userID := uuid.New()
|
||||
modelConfigID := uuid.New()
|
||||
expectedKey := codersdk.CompactionThresholdKey(modelConfigID)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
dbReturn string
|
||||
dbErr error
|
||||
wantVal int32
|
||||
wantOK bool
|
||||
wantWarnLog bool
|
||||
}{
|
||||
{
|
||||
name: "NoRowsReturnsDefault",
|
||||
dbErr: sql.ErrNoRows,
|
||||
wantOK: false,
|
||||
},
|
||||
{
|
||||
name: "ValidOverride",
|
||||
dbReturn: "75",
|
||||
wantVal: 75,
|
||||
wantOK: true,
|
||||
},
|
||||
{
|
||||
name: "OutOfRangeValue",
|
||||
dbReturn: "101",
|
||||
wantOK: false,
|
||||
},
|
||||
{
|
||||
name: "NonIntegerValue",
|
||||
dbReturn: "abc",
|
||||
wantOK: false,
|
||||
},
|
||||
{
|
||||
name: "UnexpectedDBError",
|
||||
dbErr: xerrors.New("connection refused"),
|
||||
wantOK: false,
|
||||
wantWarnLog: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
tc := tc
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
mockDB := dbmock.NewMockStore(ctrl)
|
||||
sink := testutil.NewFakeSink(t)
|
||||
|
||||
srv := &Server{
|
||||
db: mockDB,
|
||||
logger: sink.Logger(),
|
||||
}
|
||||
|
||||
mockDB.EXPECT().GetUserChatCompactionThreshold(gomock.Any(), database.GetUserChatCompactionThresholdParams{
|
||||
UserID: userID,
|
||||
Key: expectedKey,
|
||||
}).Return(tc.dbReturn, tc.dbErr)
|
||||
|
||||
val, ok := srv.resolveUserCompactionThreshold(context.Background(), userID, modelConfigID)
|
||||
require.Equal(t, tc.wantVal, val)
|
||||
require.Equal(t, tc.wantOK, ok)
|
||||
|
||||
warns := sink.Entries(func(e slog.SinkEntry) bool {
|
||||
return e.Level == slog.LevelWarn
|
||||
})
|
||||
if tc.wantWarnLog {
|
||||
require.NotEmpty(t, warns, "expected a warning log entry")
|
||||
return
|
||||
}
|
||||
require.Empty(t, warns, "unexpected warning log entry")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// requireFieldValue asserts that a SinkEntry contains a field with
|
||||
// the given name and value.
|
||||
func requireFieldValue(t *testing.T, entry slog.SinkEntry, name string, expected interface{}) {
|
||||
|
||||
@@ -22,6 +22,16 @@ import (
|
||||
"github.com/coder/websocket/wsjson"
|
||||
)
|
||||
|
||||
// ChatCompactionThresholdKeyPrefix scopes per-model chat compaction
|
||||
// threshold settings.
|
||||
const ChatCompactionThresholdKeyPrefix = "chat_compaction_threshold_pct:"
|
||||
|
||||
// CompactionThresholdKey returns the user-config key for a specific
|
||||
// model configuration's compaction threshold.
|
||||
func CompactionThresholdKey(modelConfigID uuid.UUID) string {
|
||||
return ChatCompactionThresholdKeyPrefix + modelConfigID.String()
|
||||
}
|
||||
|
||||
// ChatStatus represents the status of a chat.
|
||||
type ChatStatus string
|
||||
|
||||
@@ -349,6 +359,25 @@ type UserChatCustomPrompt struct {
|
||||
CustomPrompt string `json:"custom_prompt"`
|
||||
}
|
||||
|
||||
// UserChatCompactionThreshold is a user's per-model chat compaction
|
||||
// threshold override.
|
||||
type UserChatCompactionThreshold struct {
|
||||
ModelConfigID uuid.UUID `json:"model_config_id" format:"uuid"`
|
||||
ThresholdPercent int32 `json:"threshold_percent"`
|
||||
}
|
||||
|
||||
// UserChatCompactionThresholds wraps the user's per-model chat
|
||||
// compaction threshold overrides.
|
||||
type UserChatCompactionThresholds struct {
|
||||
Thresholds []UserChatCompactionThreshold `json:"thresholds"`
|
||||
}
|
||||
|
||||
// UpdateUserChatCompactionThresholdRequest sets a user's per-model
|
||||
// chat compaction threshold override.
|
||||
type UpdateUserChatCompactionThresholdRequest struct {
|
||||
ThresholdPercent int32 `json:"threshold_percent" validate:"min=0,max=100"`
|
||||
}
|
||||
|
||||
// ChatDesktopEnabledResponse is the response for getting the desktop setting.
|
||||
type ChatDesktopEnabledResponse struct {
|
||||
EnableDesktop bool `json:"enable_desktop"`
|
||||
@@ -1413,6 +1442,50 @@ func (c *ExperimentalClient) UpdateUserChatCustomPrompt(ctx context.Context, req
|
||||
return resp, json.NewDecoder(res.Body).Decode(&resp)
|
||||
}
|
||||
|
||||
// GetUserChatCompactionThresholds fetches the user's per-model chat
|
||||
// compaction thresholds.
|
||||
func (c *ExperimentalClient) GetUserChatCompactionThresholds(ctx context.Context) (UserChatCompactionThresholds, error) {
|
||||
res, err := c.Request(ctx, http.MethodGet, "/api/experimental/chats/config/user-compaction-thresholds", nil)
|
||||
if err != nil {
|
||||
return UserChatCompactionThresholds{}, err
|
||||
}
|
||||
defer res.Body.Close()
|
||||
if res.StatusCode != http.StatusOK {
|
||||
return UserChatCompactionThresholds{}, ReadBodyAsError(res)
|
||||
}
|
||||
var thresholds UserChatCompactionThresholds
|
||||
return thresholds, json.NewDecoder(res.Body).Decode(&thresholds)
|
||||
}
|
||||
|
||||
// UpdateUserChatCompactionThreshold updates the user's per-model chat
|
||||
// compaction threshold.
|
||||
func (c *ExperimentalClient) UpdateUserChatCompactionThreshold(ctx context.Context, modelConfigID uuid.UUID, req UpdateUserChatCompactionThresholdRequest) (UserChatCompactionThreshold, error) {
|
||||
res, err := c.Request(ctx, http.MethodPut, fmt.Sprintf("/api/experimental/chats/config/user-compaction-thresholds/%s", modelConfigID), req)
|
||||
if err != nil {
|
||||
return UserChatCompactionThreshold{}, err
|
||||
}
|
||||
defer res.Body.Close()
|
||||
if res.StatusCode != http.StatusOK {
|
||||
return UserChatCompactionThreshold{}, ReadBodyAsError(res)
|
||||
}
|
||||
var threshold UserChatCompactionThreshold
|
||||
return threshold, json.NewDecoder(res.Body).Decode(&threshold)
|
||||
}
|
||||
|
||||
// DeleteUserChatCompactionThreshold deletes the user's per-model chat
|
||||
// compaction threshold override.
|
||||
func (c *ExperimentalClient) DeleteUserChatCompactionThreshold(ctx context.Context, modelConfigID uuid.UUID) error {
|
||||
res, err := c.Request(ctx, http.MethodDelete, fmt.Sprintf("/api/experimental/chats/config/user-compaction-thresholds/%s", modelConfigID), nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer res.Body.Close()
|
||||
if res.StatusCode != http.StatusNoContent {
|
||||
return ReadBodyAsError(res)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// CreateChat creates a new chat.
|
||||
func (c *ExperimentalClient) CreateChat(ctx context.Context, req CreateChatRequest) (Chat, error) {
|
||||
res, err := c.Request(ctx, http.MethodPost, "/api/experimental/chats", req)
|
||||
|
||||
@@ -3221,6 +3221,32 @@ class ExperimentalApiMethods {
|
||||
return response.data;
|
||||
};
|
||||
|
||||
getUserChatCompactionThresholds =
|
||||
async (): Promise<TypesGen.UserChatCompactionThresholds> => {
|
||||
const response =
|
||||
await this.axios.get<TypesGen.UserChatCompactionThresholds>(
|
||||
"/api/experimental/chats/config/user-compaction-thresholds",
|
||||
);
|
||||
return response.data;
|
||||
};
|
||||
updateUserChatCompactionThreshold = async (
|
||||
modelConfigId: string,
|
||||
req: TypesGen.UpdateUserChatCompactionThresholdRequest,
|
||||
): Promise<TypesGen.UserChatCompactionThreshold> => {
|
||||
const response = await this.axios.put<TypesGen.UserChatCompactionThreshold>(
|
||||
`/api/experimental/chats/config/user-compaction-thresholds/${encodeURIComponent(modelConfigId)}`,
|
||||
req,
|
||||
);
|
||||
return response.data;
|
||||
};
|
||||
deleteUserChatCompactionThreshold = async (
|
||||
modelConfigId: string,
|
||||
): Promise<void> => {
|
||||
await this.axios.delete(
|
||||
`/api/experimental/chats/config/user-compaction-thresholds/${encodeURIComponent(modelConfigId)}`,
|
||||
);
|
||||
};
|
||||
|
||||
getChatProviderConfigs = async (): Promise<TypesGen.ChatProviderConfig[]> => {
|
||||
const response = await this.axios.get<TypesGen.ChatProviderConfig[]>(
|
||||
chatProviderConfigsPath,
|
||||
|
||||
@@ -441,6 +441,41 @@ export const updateUserChatCustomPrompt = (queryClient: QueryClient) => ({
|
||||
},
|
||||
});
|
||||
|
||||
const userCompactionThresholdsKey = [
|
||||
"chat-user-compaction-thresholds",
|
||||
] as const;
|
||||
|
||||
export const userCompactionThresholds = () => ({
|
||||
queryKey: userCompactionThresholdsKey,
|
||||
queryFn: () => API.experimental.getUserChatCompactionThresholds(),
|
||||
});
|
||||
|
||||
export const updateUserCompactionThreshold = (queryClient: QueryClient) => ({
|
||||
mutationFn: (vars: {
|
||||
modelConfigId: string;
|
||||
req: TypesGen.UpdateUserChatCompactionThresholdRequest;
|
||||
}) =>
|
||||
API.experimental.updateUserChatCompactionThreshold(
|
||||
vars.modelConfigId,
|
||||
vars.req,
|
||||
),
|
||||
onSuccess: async () => {
|
||||
await queryClient.invalidateQueries({
|
||||
queryKey: userCompactionThresholdsKey,
|
||||
});
|
||||
},
|
||||
});
|
||||
|
||||
export const deleteUserCompactionThreshold = (queryClient: QueryClient) => ({
|
||||
mutationFn: (modelConfigId: string) =>
|
||||
API.experimental.deleteUserChatCompactionThreshold(modelConfigId),
|
||||
onSuccess: async () => {
|
||||
await queryClient.invalidateQueries({
|
||||
queryKey: userCompactionThresholdsKey,
|
||||
});
|
||||
},
|
||||
});
|
||||
|
||||
export const chatModelsKey = ["chat-models"] as const;
|
||||
|
||||
export const chatModels = () => ({
|
||||
|
||||
Generated
+36
@@ -1074,6 +1074,14 @@ export interface Chat {
|
||||
readonly mcp_server_ids: readonly string[];
|
||||
}
|
||||
|
||||
// From codersdk/chats.go
|
||||
/**
|
||||
* ChatCompactionThresholdKeyPrefix scopes per-model chat compaction
|
||||
* threshold settings.
|
||||
*/
|
||||
export const ChatCompactionThresholdKeyPrefix =
|
||||
"chat_compaction_threshold_pct:";
|
||||
|
||||
// From codersdk/deployment.go
|
||||
export interface ChatConfig {
|
||||
readonly acquire_batch_size: number;
|
||||
@@ -7119,6 +7127,15 @@ export interface UpdateUserAppearanceSettingsRequest {
|
||||
readonly terminal_font: TerminalFontName;
|
||||
}
|
||||
|
||||
// From codersdk/chats.go
|
||||
/**
|
||||
* UpdateUserChatCompactionThresholdRequest sets a user's per-model
|
||||
* chat compaction threshold override.
|
||||
*/
|
||||
export interface UpdateUserChatCompactionThresholdRequest {
|
||||
readonly threshold_percent: number;
|
||||
}
|
||||
|
||||
// From codersdk/notifications.go
|
||||
export interface UpdateUserNotificationPreferences {
|
||||
readonly template_disabled_map: Record<string, boolean>;
|
||||
@@ -7371,6 +7388,25 @@ export interface UserAppearanceSettings {
|
||||
readonly terminal_font: TerminalFontName;
|
||||
}
|
||||
|
||||
// From codersdk/chats.go
|
||||
/**
|
||||
* UserChatCompactionThreshold is a user's per-model chat compaction
|
||||
* threshold override.
|
||||
*/
|
||||
export interface UserChatCompactionThreshold {
|
||||
readonly model_config_id: string;
|
||||
readonly threshold_percent: number;
|
||||
}
|
||||
|
||||
// From codersdk/chats.go
|
||||
/**
|
||||
* UserChatCompactionThresholds wraps the user's per-model chat
|
||||
* compaction threshold overrides.
|
||||
*/
|
||||
export interface UserChatCompactionThresholds {
|
||||
readonly thresholds: readonly UserChatCompactionThreshold[];
|
||||
}
|
||||
|
||||
// From codersdk/chats.go
|
||||
/**
|
||||
* UserChatCustomPrompt is the request and response body for the
|
||||
|
||||
@@ -12,6 +12,7 @@ import {
|
||||
editChatMessage,
|
||||
interruptChat,
|
||||
promoteChatQueuedMessage,
|
||||
userCompactionThresholds,
|
||||
} from "api/queries/chats";
|
||||
import { deploymentSSHConfig } from "api/queries/deployment";
|
||||
import { workspaceById, workspaceByIdKey } from "api/queries/workspaces";
|
||||
@@ -231,6 +232,27 @@ export function useConversationEditingState(deps: {
|
||||
};
|
||||
}
|
||||
|
||||
/**
|
||||
* Resolves the effective compaction threshold for a model configuration,
|
||||
* preferring the user's override when set.
|
||||
*/
|
||||
function resolveCompactionThreshold(
|
||||
modelConfigID: string | undefined,
|
||||
userThresholds: readonly TypesGen.UserChatCompactionThreshold[] | undefined,
|
||||
modelConfigs: readonly TypesGen.ChatModelConfig[],
|
||||
): number | undefined {
|
||||
if (!modelConfigID) return undefined;
|
||||
const config = modelConfigs.find((c) => c.id === modelConfigID);
|
||||
if (!config) return undefined;
|
||||
const userOverride = userThresholds?.find(
|
||||
(t) => t.model_config_id === modelConfigID,
|
||||
);
|
||||
if (userOverride) {
|
||||
return userOverride.threshold_percent;
|
||||
}
|
||||
return config.compression_threshold;
|
||||
}
|
||||
|
||||
const AgentDetail: FC = () => {
|
||||
const { agentId } = useParams<{ agentId: string }>();
|
||||
const {
|
||||
@@ -296,6 +318,7 @@ const AgentDetail: FC = () => {
|
||||
|
||||
const chatModelsQuery = useQuery(chatModels());
|
||||
const chatModelConfigsQuery = useQuery(chatModelConfigs());
|
||||
const userThresholdsQuery = useQuery(userCompactionThresholds());
|
||||
const desktopEnabledQuery = useQuery(chatDesktopEnabled());
|
||||
const desktopEnabled = desktopEnabledQuery.data?.enable_desktop ?? false;
|
||||
|
||||
@@ -492,10 +515,11 @@ const AgentDetail: FC = () => {
|
||||
return modelOptions[0]?.id ?? "";
|
||||
})();
|
||||
|
||||
const compressionThreshold = chatLastModelConfigID
|
||||
? modelConfigs.find((c) => c.id === chatLastModelConfigID)
|
||||
?.compression_threshold
|
||||
: undefined;
|
||||
const compressionThreshold = resolveCompactionThreshold(
|
||||
chatLastModelConfigID,
|
||||
userThresholdsQuery.data?.thresholds,
|
||||
modelConfigs,
|
||||
);
|
||||
const hasModelOptions = modelOptions.length > 0;
|
||||
const hasConfiguredModels = hasConfiguredModelsInCatalog(modelCatalog);
|
||||
const modelSelectorPlaceholder = getModelSelectorPlaceholder(
|
||||
|
||||
@@ -159,6 +159,13 @@ const meta = {
|
||||
spyOn(API.experimental, "updateUserChatCustomPrompt").mockResolvedValue({
|
||||
custom_prompt: "",
|
||||
});
|
||||
spyOn(API.experimental, "getChatModelConfigs").mockResolvedValue([]);
|
||||
spyOn(
|
||||
API.experimental,
|
||||
"getUserChatCompactionThresholds",
|
||||
).mockResolvedValue({
|
||||
thresholds: [],
|
||||
});
|
||||
spyOn(API.experimental, "getChatWorkspaceTTL").mockResolvedValue({
|
||||
workspace_ttl_ms: 0,
|
||||
});
|
||||
|
||||
@@ -3,6 +3,7 @@ import {
|
||||
chatCostSummary,
|
||||
chatCostUsers,
|
||||
chatDesktopEnabled,
|
||||
chatModelConfigs,
|
||||
chatSystemPrompt,
|
||||
chatUserCustomPrompt,
|
||||
chatWorkspaceTTL,
|
||||
@@ -62,6 +63,7 @@ import { InsightsContent } from "./components/InsightsContent";
|
||||
import { LimitsTab } from "./components/LimitsTab";
|
||||
import { MCPServerAdminPanel } from "./components/MCPServerAdminPanel";
|
||||
import { SectionHeader } from "./components/SectionHeader";
|
||||
import { UserCompactionThresholdSettings } from "./UserCompactionThresholdSettings";
|
||||
|
||||
const AdminBadge: FC = () => (
|
||||
<TooltipProvider delayDuration={0}>
|
||||
@@ -526,6 +528,10 @@ export const AgentSettingsPageView: FC<AgentSettingsPageViewProps> = ({
|
||||
} = useMutation(updateChatDesktopEnabled(queryClient));
|
||||
|
||||
const workspaceTTLQuery = useQuery(chatWorkspaceTTL());
|
||||
const modelConfigsQuery = useQuery({
|
||||
...chatModelConfigs(),
|
||||
enabled: activeSection === "behavior",
|
||||
});
|
||||
const {
|
||||
mutate: saveWorkspaceTTL,
|
||||
isPending: isSavingWorkspaceTTL,
|
||||
@@ -636,6 +642,13 @@ export const AgentSettingsPageView: FC<AgentSettingsPageViewProps> = ({
|
||||
)}
|
||||
</form>
|
||||
|
||||
<hr className="my-5 border-0 border-t border-solid border-border" />
|
||||
<UserCompactionThresholdSettings
|
||||
modelConfigs={modelConfigsQuery.data ?? []}
|
||||
modelConfigsError={modelConfigsQuery.error}
|
||||
isLoadingModelConfigs={modelConfigsQuery.isLoading}
|
||||
/>
|
||||
|
||||
{/* ── Admin system prompt (admin only) ── */}
|
||||
{canSetSystemPrompt && (
|
||||
<>
|
||||
|
||||
@@ -0,0 +1,201 @@
|
||||
import { MockUserOwner } from "testHelpers/entities";
|
||||
import { withAuthProvider, withDashboardProvider } from "testHelpers/storybook";
|
||||
import type { Meta, StoryObj } from "@storybook/react-vite";
|
||||
import { API } from "api/api";
|
||||
import type * as TypesGen from "api/typesGenerated";
|
||||
import { expect, spyOn, userEvent, waitFor, within } from "storybook/test";
|
||||
import { UserCompactionThresholdSettings } from "./UserCompactionThresholdSettings";
|
||||
|
||||
const mockModelConfigs: TypesGen.ChatModelConfig[] = [
|
||||
{
|
||||
id: "model-1",
|
||||
provider: "openai",
|
||||
model: "gpt-4o",
|
||||
display_name: "GPT-4o",
|
||||
enabled: true,
|
||||
is_default: true,
|
||||
context_limit: 128000,
|
||||
compression_threshold: 80,
|
||||
created_at: "2025-01-01T00:00:00Z",
|
||||
updated_at: "2025-01-01T00:00:00Z",
|
||||
},
|
||||
{
|
||||
id: "model-2",
|
||||
provider: "anthropic",
|
||||
model: "claude-sonnet",
|
||||
display_name: "Claude Sonnet",
|
||||
enabled: true,
|
||||
is_default: false,
|
||||
context_limit: 200000,
|
||||
compression_threshold: 70,
|
||||
created_at: "2025-01-01T00:00:00Z",
|
||||
updated_at: "2025-01-01T00:00:00Z",
|
||||
},
|
||||
{
|
||||
id: "model-3",
|
||||
provider: "openai",
|
||||
model: "gpt-3.5",
|
||||
display_name: "GPT-3.5 (Disabled)",
|
||||
enabled: false,
|
||||
is_default: false,
|
||||
context_limit: 16000,
|
||||
compression_threshold: 60,
|
||||
created_at: "2025-01-01T00:00:00Z",
|
||||
updated_at: "2025-01-01T00:00:00Z",
|
||||
},
|
||||
];
|
||||
|
||||
const meta = {
|
||||
title: "pages/AgentsPage/UserCompactionThresholdSettings",
|
||||
component: UserCompactionThresholdSettings,
|
||||
decorators: [withAuthProvider, withDashboardProvider],
|
||||
args: {
|
||||
modelConfigs: mockModelConfigs,
|
||||
},
|
||||
parameters: {
|
||||
user: MockUserOwner,
|
||||
},
|
||||
} satisfies Meta<typeof UserCompactionThresholdSettings>;
|
||||
|
||||
export default meta;
|
||||
type Story = StoryObj<typeof meta>;
|
||||
|
||||
export const Default: Story = {
|
||||
beforeEach: () => {
|
||||
spyOn(
|
||||
API.experimental,
|
||||
"getUserChatCompactionThresholds",
|
||||
).mockResolvedValue({
|
||||
thresholds: [],
|
||||
});
|
||||
spyOn(
|
||||
API.experimental,
|
||||
"updateUserChatCompactionThreshold",
|
||||
).mockResolvedValue({
|
||||
model_config_id: "model-1",
|
||||
threshold_percent: 90,
|
||||
});
|
||||
spyOn(
|
||||
API.experimental,
|
||||
"deleteUserChatCompactionThreshold",
|
||||
).mockResolvedValue(undefined);
|
||||
},
|
||||
play: async ({ canvasElement }) => {
|
||||
const canvas = within(canvasElement);
|
||||
const gpt4oInput = await canvas.findByRole("spinbutton", {
|
||||
name: /GPT-4o compaction threshold/i,
|
||||
});
|
||||
|
||||
expect(canvas.getByText("GPT-4o")).toBeInTheDocument();
|
||||
expect(canvas.getByText("Claude Sonnet")).toBeInTheDocument();
|
||||
expect(canvas.getByText("System default: 80%")).toBeInTheDocument();
|
||||
expect(canvas.getByText("System default: 70%")).toBeInTheDocument();
|
||||
expect(canvas.queryByText("GPT-3.5 (Disabled)")).not.toBeInTheDocument();
|
||||
|
||||
await userEvent.type(gpt4oInput, "100");
|
||||
expect(
|
||||
canvas.getByText(
|
||||
"⚠ Setting 100% will disable auto-compaction for this model.",
|
||||
),
|
||||
).toBeInTheDocument();
|
||||
await userEvent.clear(gpt4oInput);
|
||||
await userEvent.type(gpt4oInput, "95");
|
||||
|
||||
const saveButtons = canvas.getAllByRole("button", { name: "Save" });
|
||||
await waitFor(() => {
|
||||
expect(saveButtons[0]).toBeEnabled();
|
||||
});
|
||||
|
||||
await userEvent.click(saveButtons[0]);
|
||||
await waitFor(() => {
|
||||
expect(
|
||||
API.experimental.updateUserChatCompactionThreshold,
|
||||
).toHaveBeenCalledWith("model-1", { threshold_percent: 95 });
|
||||
});
|
||||
},
|
||||
};
|
||||
|
||||
export const WithOverrides: Story = {
|
||||
beforeEach: () => {
|
||||
spyOn(
|
||||
API.experimental,
|
||||
"getUserChatCompactionThresholds",
|
||||
).mockResolvedValue({
|
||||
thresholds: [
|
||||
{ model_config_id: "model-1", threshold_percent: 90 },
|
||||
{ model_config_id: "model-2", threshold_percent: 50 },
|
||||
],
|
||||
});
|
||||
spyOn(
|
||||
API.experimental,
|
||||
"updateUserChatCompactionThreshold",
|
||||
).mockResolvedValue({
|
||||
model_config_id: "model-1",
|
||||
threshold_percent: 90,
|
||||
});
|
||||
spyOn(
|
||||
API.experimental,
|
||||
"deleteUserChatCompactionThreshold",
|
||||
).mockResolvedValue(undefined);
|
||||
},
|
||||
play: async ({ canvasElement }) => {
|
||||
const canvas = within(canvasElement);
|
||||
const gpt4oInput = await canvas.findByRole("spinbutton", {
|
||||
name: /GPT-4o compaction threshold/i,
|
||||
});
|
||||
const claudeInput = await canvas.findByRole("spinbutton", {
|
||||
name: /Claude Sonnet compaction threshold/i,
|
||||
});
|
||||
|
||||
expect(gpt4oInput).toHaveValue(90);
|
||||
expect(claudeInput).toHaveValue(50);
|
||||
|
||||
const resetButtons = canvas.getAllByRole("button", { name: "Reset" });
|
||||
await userEvent.click(resetButtons[0]);
|
||||
await waitFor(() => {
|
||||
expect(
|
||||
API.experimental.deleteUserChatCompactionThreshold,
|
||||
).toHaveBeenCalledWith("model-1");
|
||||
});
|
||||
},
|
||||
};
|
||||
|
||||
export const Loading: Story = {
|
||||
beforeEach: () => {
|
||||
spyOn(API.experimental, "getUserChatCompactionThresholds").mockReturnValue(
|
||||
new Promise(() => {}),
|
||||
);
|
||||
spyOn(
|
||||
API.experimental,
|
||||
"updateUserChatCompactionThreshold",
|
||||
).mockResolvedValue({
|
||||
model_config_id: "model-1",
|
||||
threshold_percent: 90,
|
||||
});
|
||||
spyOn(
|
||||
API.experimental,
|
||||
"deleteUserChatCompactionThreshold",
|
||||
).mockResolvedValue(undefined);
|
||||
},
|
||||
};
|
||||
|
||||
export const ErrorState: Story = {
|
||||
name: "Error",
|
||||
beforeEach: () => {
|
||||
spyOn(
|
||||
API.experimental,
|
||||
"getUserChatCompactionThresholds",
|
||||
).mockRejectedValue(new globalThis.Error("Failed to load thresholds"));
|
||||
spyOn(
|
||||
API.experimental,
|
||||
"updateUserChatCompactionThreshold",
|
||||
).mockResolvedValue({
|
||||
model_config_id: "model-1",
|
||||
threshold_percent: 90,
|
||||
});
|
||||
spyOn(
|
||||
API.experimental,
|
||||
"deleteUserChatCompactionThreshold",
|
||||
).mockResolvedValue(undefined);
|
||||
},
|
||||
};
|
||||
@@ -0,0 +1,302 @@
|
||||
import { getErrorMessage } from "api/errors";
|
||||
import {
|
||||
deleteUserCompactionThreshold,
|
||||
updateUserCompactionThreshold,
|
||||
userCompactionThresholds,
|
||||
} from "api/queries/chats";
|
||||
import type * as TypesGen from "api/typesGenerated";
|
||||
import { Button } from "components/Button/Button";
|
||||
import { Input } from "components/Input/Input";
|
||||
import { Spinner } from "components/Spinner/Spinner";
|
||||
import { type FC, useState } from "react";
|
||||
import { useMutation, useQuery, useQueryClient } from "react-query";
|
||||
|
||||
interface UserCompactionThresholdSettingsProps {
|
||||
modelConfigs: readonly TypesGen.ChatModelConfig[];
|
||||
modelConfigsError?: unknown;
|
||||
isLoadingModelConfigs?: boolean;
|
||||
}
|
||||
|
||||
const parseThresholdDraft = (value: string): number | null => {
|
||||
const trimmedValue = value.trim();
|
||||
if (!/^\d+$/.test(trimmedValue)) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const parsedValue = Number(trimmedValue);
|
||||
if (!Number.isInteger(parsedValue) || parsedValue < 0 || parsedValue > 100) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return parsedValue;
|
||||
};
|
||||
|
||||
export const UserCompactionThresholdSettings: FC<
|
||||
UserCompactionThresholdSettingsProps
|
||||
> = ({ modelConfigs, modelConfigsError, isLoadingModelConfigs }) => {
|
||||
const queryClient = useQueryClient();
|
||||
const thresholdsQuery = useQuery(userCompactionThresholds());
|
||||
const [drafts, setDrafts] = useState<Record<string, string>>({});
|
||||
const [rowErrors, setRowErrors] = useState<Record<string, string>>({});
|
||||
const [pendingModels, setPendingModels] = useState<Set<string>>(new Set());
|
||||
|
||||
const clearDraft = (modelConfigID: string) => {
|
||||
setDrafts((currentDrafts) => {
|
||||
const nextDrafts = { ...currentDrafts };
|
||||
delete nextDrafts[modelConfigID];
|
||||
return nextDrafts;
|
||||
});
|
||||
};
|
||||
|
||||
const clearRowError = (modelConfigID: string) => {
|
||||
setRowErrors((currentErrors) => {
|
||||
if (!(modelConfigID in currentErrors)) {
|
||||
return currentErrors;
|
||||
}
|
||||
const nextErrors = { ...currentErrors };
|
||||
delete nextErrors[modelConfigID];
|
||||
return nextErrors;
|
||||
});
|
||||
};
|
||||
|
||||
const saveOpts = updateUserCompactionThreshold(queryClient);
|
||||
const saveThresholdMutation = useMutation({
|
||||
...saveOpts,
|
||||
onSuccess: async (_data, variables) => {
|
||||
await saveOpts.onSuccess?.();
|
||||
clearDraft(variables.modelConfigId);
|
||||
clearRowError(variables.modelConfigId);
|
||||
},
|
||||
onError: (error, variables) => {
|
||||
setRowErrors((currentErrors) => ({
|
||||
...currentErrors,
|
||||
[variables.modelConfigId]: getErrorMessage(
|
||||
error,
|
||||
"Failed to save compaction threshold.",
|
||||
),
|
||||
}));
|
||||
},
|
||||
onSettled: async (_data, _error, variables) => {
|
||||
setPendingModels((currentPendingModels) => {
|
||||
const nextPendingModels = new Set(currentPendingModels);
|
||||
nextPendingModels.delete(variables.modelConfigId);
|
||||
return nextPendingModels;
|
||||
});
|
||||
},
|
||||
});
|
||||
const resetOpts = deleteUserCompactionThreshold(queryClient);
|
||||
const resetThresholdMutation = useMutation({
|
||||
...resetOpts,
|
||||
onSuccess: async (_data, variables) => {
|
||||
await resetOpts.onSuccess?.();
|
||||
clearDraft(variables);
|
||||
clearRowError(variables);
|
||||
},
|
||||
onError: (error, variables) => {
|
||||
setRowErrors((currentErrors) => ({
|
||||
...currentErrors,
|
||||
[variables]: getErrorMessage(
|
||||
error,
|
||||
"Failed to reset compaction threshold.",
|
||||
),
|
||||
}));
|
||||
},
|
||||
onSettled: async (_data, _error, variables) => {
|
||||
setPendingModels((currentPendingModels) => {
|
||||
const nextPendingModels = new Set(currentPendingModels);
|
||||
nextPendingModels.delete(variables);
|
||||
return nextPendingModels;
|
||||
});
|
||||
},
|
||||
});
|
||||
|
||||
const enabledModelConfigs = modelConfigs.filter((config) => config.enabled);
|
||||
const overridesByModelID = new Map(
|
||||
(thresholdsQuery.data?.thresholds ?? []).map(
|
||||
(threshold: TypesGen.UserChatCompactionThreshold) => [
|
||||
threshold.model_config_id,
|
||||
threshold.threshold_percent,
|
||||
],
|
||||
),
|
||||
);
|
||||
if (thresholdsQuery.isLoading) {
|
||||
return (
|
||||
<div className="space-y-2">
|
||||
<h3 className="m-0 text-[13px] font-semibold text-content-primary">
|
||||
Context Compaction
|
||||
</h3>
|
||||
<p className="!mt-0.5 m-0 text-xs text-content-secondary">
|
||||
Control when chat context is automatically summarized for each model.
|
||||
Setting 100% means the chat will never auto-compact.
|
||||
</p>
|
||||
<div className="flex items-center gap-2 text-sm text-content-secondary">
|
||||
<Spinner loading className="h-4 w-4" />
|
||||
Loading thresholds...
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
if (thresholdsQuery.isError) {
|
||||
return (
|
||||
<div className="space-y-2">
|
||||
<h3 className="m-0 text-[13px] font-semibold text-content-primary">
|
||||
Context Compaction
|
||||
</h3>
|
||||
<p className="!mt-0.5 m-0 text-xs text-content-secondary">
|
||||
Control when chat context is automatically summarized for each model.
|
||||
Setting 100% means the chat will never auto-compact.
|
||||
</p>
|
||||
<p className="m-0 text-xs text-content-destructive">
|
||||
{getErrorMessage(
|
||||
thresholdsQuery.error,
|
||||
"Failed to load compaction thresholds.",
|
||||
)}
|
||||
</p>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
return (
|
||||
<div className="space-y-2">
|
||||
<h3 className="m-0 text-[13px] font-semibold text-content-primary">
|
||||
Context Compaction
|
||||
</h3>
|
||||
<p className="!mt-0.5 m-0 text-xs text-content-secondary">
|
||||
Control when chat context is automatically summarized for each model.
|
||||
Setting 100% means the chat will never auto-compact.
|
||||
</p>
|
||||
{isLoadingModelConfigs ? (
|
||||
<div className="flex items-center gap-2 text-sm text-content-secondary">
|
||||
<Spinner loading className="h-4 w-4" />
|
||||
Loading models...
|
||||
</div>
|
||||
) : modelConfigsError ? (
|
||||
<p className="m-0 text-xs text-content-destructive">
|
||||
{getErrorMessage(
|
||||
modelConfigsError,
|
||||
"Failed to load model configurations.",
|
||||
)}
|
||||
</p>
|
||||
) : enabledModelConfigs.length === 0 ? (
|
||||
<p className="m-0 text-xs text-content-secondary">
|
||||
No enabled chat models available. An administrator must configure chat
|
||||
models before compaction thresholds can be set.
|
||||
</p>
|
||||
) : (
|
||||
<div className="space-y-3">
|
||||
{enabledModelConfigs.map((modelConfig) => {
|
||||
const existingOverride = overridesByModelID.get(modelConfig.id);
|
||||
const hasOverride = overridesByModelID.has(modelConfig.id);
|
||||
const draftValue =
|
||||
drafts[modelConfig.id] ??
|
||||
(existingOverride !== undefined ? String(existingOverride) : "");
|
||||
const parsedDraftValue = parseThresholdDraft(draftValue);
|
||||
const isThisModelMutating = pendingModels.has(modelConfig.id);
|
||||
const isSaveDisabled =
|
||||
draftValue.length === 0 ||
|
||||
parsedDraftValue === null ||
|
||||
parsedDraftValue === existingOverride ||
|
||||
isThisModelMutating;
|
||||
|
||||
return (
|
||||
<div
|
||||
key={modelConfig.id}
|
||||
className="space-y-2 rounded-lg border border-border bg-surface-secondary/40 p-4"
|
||||
>
|
||||
<div className="flex flex-col gap-3 sm:flex-row sm:items-center sm:justify-between">
|
||||
<div className="flex-1">
|
||||
<span className="text-sm font-medium text-content-primary">
|
||||
{modelConfig.display_name || modelConfig.model}
|
||||
</span>
|
||||
<span className="ml-2 text-xs text-content-secondary">
|
||||
System default: {modelConfig.compression_threshold}%
|
||||
</span>
|
||||
</div>
|
||||
<div className="flex flex-wrap items-center gap-2 sm:justify-end">
|
||||
<Input
|
||||
aria-label={`${modelConfig.display_name || modelConfig.model} compaction threshold`}
|
||||
type="number"
|
||||
min={0}
|
||||
max={100}
|
||||
inputMode="numeric"
|
||||
className="h-9 w-20 text-[13px]"
|
||||
value={draftValue}
|
||||
placeholder={String(modelConfig.compression_threshold)}
|
||||
onChange={(event) => {
|
||||
setDrafts((currentDrafts) => ({
|
||||
...currentDrafts,
|
||||
[modelConfig.id]: event.target.value,
|
||||
}));
|
||||
clearRowError(modelConfig.id);
|
||||
}}
|
||||
disabled={isThisModelMutating}
|
||||
/>
|
||||
<span className="text-xs text-content-secondary">%</span>
|
||||
<Button
|
||||
size="sm"
|
||||
type="button"
|
||||
disabled={isSaveDisabled}
|
||||
onClick={() => {
|
||||
if (parsedDraftValue === null) {
|
||||
return;
|
||||
}
|
||||
clearRowError(modelConfig.id);
|
||||
setPendingModels((currentPendingModels) =>
|
||||
new Set(currentPendingModels).add(modelConfig.id),
|
||||
);
|
||||
saveThresholdMutation.mutate({
|
||||
modelConfigId: modelConfig.id,
|
||||
req: {
|
||||
threshold_percent: parsedDraftValue,
|
||||
},
|
||||
});
|
||||
}}
|
||||
>
|
||||
Save
|
||||
</Button>
|
||||
{hasOverride && (
|
||||
<Button
|
||||
size="sm"
|
||||
variant="outline"
|
||||
type="button"
|
||||
disabled={isThisModelMutating}
|
||||
onClick={() => {
|
||||
clearRowError(modelConfig.id);
|
||||
setPendingModels((currentPendingModels) =>
|
||||
new Set(currentPendingModels).add(modelConfig.id),
|
||||
);
|
||||
resetThresholdMutation.mutate(modelConfig.id);
|
||||
}}
|
||||
>
|
||||
Reset
|
||||
</Button>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
{draftValue.length > 0 && parsedDraftValue === null && (
|
||||
<p className="m-0 text-xs text-content-destructive">
|
||||
Enter a whole number between 0 and 100.
|
||||
</p>
|
||||
)}
|
||||
{rowErrors[modelConfig.id] && (
|
||||
<p
|
||||
aria-live="polite"
|
||||
className="m-0 text-xs text-content-destructive"
|
||||
>
|
||||
{rowErrors[modelConfig.id]}
|
||||
</p>
|
||||
)}
|
||||
{draftValue === "100" && (
|
||||
<p className="m-0 text-xs text-content-secondary">
|
||||
⚠ Setting 100% will disable auto-compaction for this model.
|
||||
</p>
|
||||
)}
|
||||
</div>
|
||||
);
|
||||
})}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
);
|
||||
};
|
||||
Reference in New Issue
Block a user