mirror of
https://github.com/coder/coder.git
synced 2026-09-21 20:51:01 +08:00
feat: add personal chat model overrides (#24715)
This commit is contained in:
@@ -1193,6 +1193,10 @@ func New(options *Options) *API {
|
||||
r.Put("/plan-mode-instructions", api.putChatPlanModeInstructions)
|
||||
r.Get("/model-override/{context}", api.getChatModelOverride)
|
||||
r.Put("/model-override/{context}", api.putChatModelOverride)
|
||||
r.Get("/personal-model-overrides", api.getChatPersonalModelOverridesAdminSettings)
|
||||
r.Put("/personal-model-overrides", api.putChatPersonalModelOverridesAdminSettings)
|
||||
r.Get("/user-personal-model-overrides", api.getUserChatPersonalModelOverrides)
|
||||
r.Put("/user-personal-model-overrides/{context}", api.putUserChatPersonalModelOverride)
|
||||
r.Get("/desktop-enabled", api.getChatDesktopEnabled)
|
||||
r.Put("/desktop-enabled", api.putChatDesktopEnabled)
|
||||
r.Get("/computer-use-provider", api.getChatComputerUseProvider)
|
||||
|
||||
@@ -2885,6 +2885,16 @@ func (q *querier) GetChatModelConfigsForTelemetry(ctx context.Context) ([]databa
|
||||
return q.db.GetChatModelConfigsForTelemetry(ctx)
|
||||
}
|
||||
|
||||
func (q *querier) GetChatPersonalModelOverridesEnabled(ctx context.Context) (bool, error) {
|
||||
// The personal model overrides flag is a deployment-wide setting read by
|
||||
// authenticated chat users. We only require that an explicit actor is
|
||||
// present in the context so unauthenticated calls fail closed.
|
||||
if _, ok := ActorFromContext(ctx); !ok {
|
||||
return false, ErrNoActor
|
||||
}
|
||||
return q.db.GetChatPersonalModelOverridesEnabled(ctx)
|
||||
}
|
||||
|
||||
func (q *querier) GetChatPlanModeInstructions(ctx context.Context) (string, error) {
|
||||
if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil {
|
||||
return "", err
|
||||
@@ -4329,6 +4339,17 @@ func (q *querier) GetUserChatDebugLoggingEnabled(ctx context.Context, userID uui
|
||||
return q.db.GetUserChatDebugLoggingEnabled(ctx, userID)
|
||||
}
|
||||
|
||||
func (q *querier) GetUserChatPersonalModelOverride(ctx context.Context, arg database.GetUserChatPersonalModelOverrideParams) (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.GetUserChatPersonalModelOverride(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) GetUserChatProviderKeys(ctx context.Context, userID uuid.UUID) ([]database.UserChatProviderKey, error) {
|
||||
u, err := q.db.GetUserByID(ctx, userID)
|
||||
if err != nil {
|
||||
@@ -5847,6 +5868,17 @@ func (q *querier) ListUserChatCompactionThresholds(ctx context.Context, userID u
|
||||
return q.db.ListUserChatCompactionThresholds(ctx, userID)
|
||||
}
|
||||
|
||||
func (q *querier) ListUserChatPersonalModelOverrides(ctx context.Context, userID uuid.UUID) ([]database.ListUserChatPersonalModelOverridesRow, 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.ListUserChatPersonalModelOverrides(ctx, userID)
|
||||
}
|
||||
|
||||
func (q *querier) ListUserSecrets(ctx context.Context, userID uuid.UUID) ([]database.ListUserSecretsRow, error) {
|
||||
obj := rbac.ResourceUserSecret.WithOwner(userID.String())
|
||||
if err := q.authorizeContext(ctx, policy.ActionRead, obj); err != nil {
|
||||
@@ -7515,6 +7547,13 @@ func (q *querier) UpsertChatIncludeDefaultSystemPrompt(ctx context.Context, incl
|
||||
return q.db.UpsertChatIncludeDefaultSystemPrompt(ctx, includeDefaultSystemPrompt)
|
||||
}
|
||||
|
||||
func (q *querier) UpsertChatPersonalModelOverridesEnabled(ctx context.Context, enabled bool) error {
|
||||
if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil {
|
||||
return err
|
||||
}
|
||||
return q.db.UpsertChatPersonalModelOverridesEnabled(ctx, enabled)
|
||||
}
|
||||
|
||||
func (q *querier) UpsertChatPlanModeInstructions(ctx context.Context, value string) error {
|
||||
if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceDeploymentConfig); err != nil {
|
||||
return err
|
||||
@@ -7732,6 +7771,17 @@ func (q *querier) UpsertUserChatDebugLoggingEnabled(ctx context.Context, arg dat
|
||||
return q.db.UpsertUserChatDebugLoggingEnabled(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) UpsertUserChatPersonalModelOverride(ctx context.Context, arg database.UpsertUserChatPersonalModelOverrideParams) 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.UpsertUserChatPersonalModelOverride(ctx, arg)
|
||||
}
|
||||
|
||||
func (q *querier) UpsertUserChatProviderKey(ctx context.Context, arg database.UpsertUserChatProviderKeyParams) (database.UserChatProviderKey, error) {
|
||||
u, err := q.db.GetUserByID(ctx, arg.UserID)
|
||||
if err != nil {
|
||||
|
||||
@@ -32,6 +32,7 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/rbac"
|
||||
"github.com/coder/coder/v2/coderd/rbac/policy"
|
||||
"github.com/coder/coder/v2/coderd/util/slice"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/provisionersdk"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
@@ -494,6 +495,10 @@ func (s *MethodTestSuite) TestChats() {
|
||||
dbm.EXPECT().GetChatDebugLoggingAllowUsers(gomock.Any()).Return(true, nil).AnyTimes()
|
||||
check.Args().Asserts().Returns(true)
|
||||
}))
|
||||
s.Run("GetChatPersonalModelOverridesEnabled", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
|
||||
dbm.EXPECT().GetChatPersonalModelOverridesEnabled(gomock.Any()).Return(true, nil).AnyTimes()
|
||||
check.Args().Asserts().Returns(true)
|
||||
}))
|
||||
s.Run("GetChatDebugRunByID", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
chat := testutil.Fake(s.T(), faker, database.Chat{})
|
||||
run := database.ChatDebugRun{ID: uuid.New(), ChatID: chat.ID}
|
||||
@@ -576,6 +581,10 @@ func (s *MethodTestSuite) TestChats() {
|
||||
dbm.EXPECT().UpsertChatAdvisorConfig(gomock.Any(), "{}").Return(nil).AnyTimes()
|
||||
check.Args("{}").Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate)
|
||||
}))
|
||||
s.Run("UpsertChatPersonalModelOverridesEnabled", s.Mocked(func(dbm *dbmock.MockStore, _ *gofakeit.Faker, check *expects) {
|
||||
dbm.EXPECT().UpsertChatPersonalModelOverridesEnabled(gomock.Any(), true).Return(nil).AnyTimes()
|
||||
check.Args(true).Asserts(rbac.ResourceDeploymentConfig, policy.ActionUpdate)
|
||||
}))
|
||||
s.Run("GetChatByID", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
chat := testutil.Fake(s.T(), faker, database.Chat{})
|
||||
dbm.EXPECT().GetChatByID(gomock.Any(), chat.ID).Return(chat, nil).AnyTimes()
|
||||
@@ -2745,6 +2754,30 @@ func (s *MethodTestSuite) TestUser() {
|
||||
dbm.EXPECT().UpsertUserChatDebugLoggingEnabled(gomock.Any(), arg).Return(nil).AnyTimes()
|
||||
check.Args(arg).Asserts(u, policy.ActionUpdatePersonal)
|
||||
}))
|
||||
s.Run("ListUserChatPersonalModelOverrides", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
u := testutil.Fake(s.T(), faker, database.User{})
|
||||
key := chatd.ChatPersonalModelOverrideKey(codersdk.ChatPersonalModelOverrideContextRoot)
|
||||
row := database.ListUserChatPersonalModelOverridesRow{Key: key, Value: "chat_default"}
|
||||
dbm.EXPECT().GetUserByID(gomock.Any(), u.ID).Return(u, nil).AnyTimes()
|
||||
dbm.EXPECT().ListUserChatPersonalModelOverrides(gomock.Any(), u.ID).Return([]database.ListUserChatPersonalModelOverridesRow{row}, nil).AnyTimes()
|
||||
check.Args(u.ID).Asserts(u, policy.ActionReadPersonal).Returns([]database.ListUserChatPersonalModelOverridesRow{row})
|
||||
}))
|
||||
s.Run("GetUserChatPersonalModelOverride", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
u := testutil.Fake(s.T(), faker, database.User{})
|
||||
key := chatd.ChatPersonalModelOverrideKey(codersdk.ChatPersonalModelOverrideContextRoot)
|
||||
arg := database.GetUserChatPersonalModelOverrideParams{UserID: u.ID, Key: key}
|
||||
dbm.EXPECT().GetUserByID(gomock.Any(), u.ID).Return(u, nil).AnyTimes()
|
||||
dbm.EXPECT().GetUserChatPersonalModelOverride(gomock.Any(), arg).Return("chat_default", nil).AnyTimes()
|
||||
check.Args(arg).Asserts(u, policy.ActionReadPersonal).Returns("chat_default")
|
||||
}))
|
||||
s.Run("UpsertUserChatPersonalModelOverride", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
|
||||
u := testutil.Fake(s.T(), faker, database.User{})
|
||||
key := chatd.ChatPersonalModelOverrideKey(codersdk.ChatPersonalModelOverrideContextRoot)
|
||||
arg := database.UpsertUserChatPersonalModelOverrideParams{UserID: u.ID, Key: key, Value: "chat_default"}
|
||||
dbm.EXPECT().GetUserByID(gomock.Any(), u.ID).Return(u, nil).AnyTimes()
|
||||
dbm.EXPECT().UpsertUserChatPersonalModelOverride(gomock.Any(), arg).Return(nil).AnyTimes()
|
||||
check.Args(arg).Asserts(u, policy.ActionUpdatePersonal)
|
||||
}))
|
||||
s.Run("UpdateUserChatCustomPrompt", 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: "chat_custom_prompt", Value: "my custom prompt"}
|
||||
|
||||
@@ -1376,6 +1376,14 @@ func (m queryMetricsStore) GetChatModelConfigsForTelemetry(ctx context.Context)
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetChatPersonalModelOverridesEnabled(ctx context.Context) (bool, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetChatPersonalModelOverridesEnabled(ctx)
|
||||
m.queryLatencies.WithLabelValues("GetChatPersonalModelOverridesEnabled").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetChatPersonalModelOverridesEnabled").Inc()
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetChatPlanModeInstructions(ctx context.Context) (string, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetChatPlanModeInstructions(ctx)
|
||||
@@ -2800,6 +2808,14 @@ func (m queryMetricsStore) GetUserChatDebugLoggingEnabled(ctx context.Context, u
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetUserChatPersonalModelOverride(ctx context.Context, arg database.GetUserChatPersonalModelOverrideParams) (string, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetUserChatPersonalModelOverride(ctx, arg)
|
||||
m.queryLatencies.WithLabelValues("GetUserChatPersonalModelOverride").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetUserChatPersonalModelOverride").Inc()
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) GetUserChatProviderKeys(ctx context.Context, userID uuid.UUID) ([]database.UserChatProviderKey, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.GetUserChatProviderKeys(ctx, userID)
|
||||
@@ -4192,6 +4208,14 @@ func (m queryMetricsStore) ListUserChatCompactionThresholds(ctx context.Context,
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) ListUserChatPersonalModelOverrides(ctx context.Context, userID uuid.UUID) ([]database.ListUserChatPersonalModelOverridesRow, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.ListUserChatPersonalModelOverrides(ctx, userID)
|
||||
m.queryLatencies.WithLabelValues("ListUserChatPersonalModelOverrides").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "ListUserChatPersonalModelOverrides").Inc()
|
||||
return r0, r1
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) ListUserSecrets(ctx context.Context, userID uuid.UUID) ([]database.ListUserSecretsRow, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.ListUserSecrets(ctx, userID)
|
||||
@@ -5400,6 +5424,14 @@ func (m queryMetricsStore) UpsertChatIncludeDefaultSystemPrompt(ctx context.Cont
|
||||
return r0
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) UpsertChatPersonalModelOverridesEnabled(ctx context.Context, enabled bool) error {
|
||||
start := time.Now()
|
||||
r0 := m.s.UpsertChatPersonalModelOverridesEnabled(ctx, enabled)
|
||||
m.queryLatencies.WithLabelValues("UpsertChatPersonalModelOverridesEnabled").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "UpsertChatPersonalModelOverridesEnabled").Inc()
|
||||
return r0
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) UpsertChatPlanModeInstructions(ctx context.Context, value string) error {
|
||||
start := time.Now()
|
||||
r0 := m.s.UpsertChatPlanModeInstructions(ctx, value)
|
||||
@@ -5624,6 +5656,14 @@ func (m queryMetricsStore) UpsertUserChatDebugLoggingEnabled(ctx context.Context
|
||||
return r0
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) UpsertUserChatPersonalModelOverride(ctx context.Context, arg database.UpsertUserChatPersonalModelOverrideParams) error {
|
||||
start := time.Now()
|
||||
r0 := m.s.UpsertUserChatPersonalModelOverride(ctx, arg)
|
||||
m.queryLatencies.WithLabelValues("UpsertUserChatPersonalModelOverride").Observe(time.Since(start).Seconds())
|
||||
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "UpsertUserChatPersonalModelOverride").Inc()
|
||||
return r0
|
||||
}
|
||||
|
||||
func (m queryMetricsStore) UpsertUserChatProviderKey(ctx context.Context, arg database.UpsertUserChatProviderKeyParams) (database.UserChatProviderKey, error) {
|
||||
start := time.Now()
|
||||
r0, r1 := m.s.UpsertUserChatProviderKey(ctx, arg)
|
||||
|
||||
@@ -2537,6 +2537,21 @@ func (mr *MockStoreMockRecorder) GetChatModelConfigsForTelemetry(ctx any) *gomoc
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatModelConfigsForTelemetry", reflect.TypeOf((*MockStore)(nil).GetChatModelConfigsForTelemetry), ctx)
|
||||
}
|
||||
|
||||
// GetChatPersonalModelOverridesEnabled mocks base method.
|
||||
func (m *MockStore) GetChatPersonalModelOverridesEnabled(ctx context.Context) (bool, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetChatPersonalModelOverridesEnabled", ctx)
|
||||
ret0, _ := ret[0].(bool)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetChatPersonalModelOverridesEnabled indicates an expected call of GetChatPersonalModelOverridesEnabled.
|
||||
func (mr *MockStoreMockRecorder) GetChatPersonalModelOverridesEnabled(ctx any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatPersonalModelOverridesEnabled", reflect.TypeOf((*MockStore)(nil).GetChatPersonalModelOverridesEnabled), ctx)
|
||||
}
|
||||
|
||||
// GetChatPlanModeInstructions mocks base method.
|
||||
func (m *MockStore) GetChatPlanModeInstructions(ctx context.Context) (string, error) {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -5237,6 +5252,21 @@ func (mr *MockStoreMockRecorder) GetUserChatDebugLoggingEnabled(ctx, userID any)
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetUserChatDebugLoggingEnabled", reflect.TypeOf((*MockStore)(nil).GetUserChatDebugLoggingEnabled), ctx, userID)
|
||||
}
|
||||
|
||||
// GetUserChatPersonalModelOverride mocks base method.
|
||||
func (m *MockStore) GetUserChatPersonalModelOverride(ctx context.Context, arg database.GetUserChatPersonalModelOverrideParams) (string, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetUserChatPersonalModelOverride", ctx, arg)
|
||||
ret0, _ := ret[0].(string)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetUserChatPersonalModelOverride indicates an expected call of GetUserChatPersonalModelOverride.
|
||||
func (mr *MockStoreMockRecorder) GetUserChatPersonalModelOverride(ctx, arg any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetUserChatPersonalModelOverride", reflect.TypeOf((*MockStore)(nil).GetUserChatPersonalModelOverride), ctx, arg)
|
||||
}
|
||||
|
||||
// GetUserChatProviderKeys mocks base method.
|
||||
func (m *MockStore) GetUserChatProviderKeys(ctx context.Context, userID uuid.UUID) ([]database.UserChatProviderKey, error) {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -7921,6 +7951,21 @@ func (mr *MockStoreMockRecorder) ListUserChatCompactionThresholds(ctx, userID an
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ListUserChatCompactionThresholds", reflect.TypeOf((*MockStore)(nil).ListUserChatCompactionThresholds), ctx, userID)
|
||||
}
|
||||
|
||||
// ListUserChatPersonalModelOverrides mocks base method.
|
||||
func (m *MockStore) ListUserChatPersonalModelOverrides(ctx context.Context, userID uuid.UUID) ([]database.ListUserChatPersonalModelOverridesRow, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "ListUserChatPersonalModelOverrides", ctx, userID)
|
||||
ret0, _ := ret[0].([]database.ListUserChatPersonalModelOverridesRow)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// ListUserChatPersonalModelOverrides indicates an expected call of ListUserChatPersonalModelOverrides.
|
||||
func (mr *MockStoreMockRecorder) ListUserChatPersonalModelOverrides(ctx, userID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ListUserChatPersonalModelOverrides", reflect.TypeOf((*MockStore)(nil).ListUserChatPersonalModelOverrides), ctx, userID)
|
||||
}
|
||||
|
||||
// ListUserSecrets mocks base method.
|
||||
func (m *MockStore) ListUserSecrets(ctx context.Context, userID uuid.UUID) ([]database.ListUserSecretsRow, error) {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -10140,6 +10185,20 @@ func (mr *MockStoreMockRecorder) UpsertChatIncludeDefaultSystemPrompt(ctx, inclu
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpsertChatIncludeDefaultSystemPrompt", reflect.TypeOf((*MockStore)(nil).UpsertChatIncludeDefaultSystemPrompt), ctx, includeDefaultSystemPrompt)
|
||||
}
|
||||
|
||||
// UpsertChatPersonalModelOverridesEnabled mocks base method.
|
||||
func (m *MockStore) UpsertChatPersonalModelOverridesEnabled(ctx context.Context, enabled bool) error {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "UpsertChatPersonalModelOverridesEnabled", ctx, enabled)
|
||||
ret0, _ := ret[0].(error)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// UpsertChatPersonalModelOverridesEnabled indicates an expected call of UpsertChatPersonalModelOverridesEnabled.
|
||||
func (mr *MockStoreMockRecorder) UpsertChatPersonalModelOverridesEnabled(ctx, enabled any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpsertChatPersonalModelOverridesEnabled", reflect.TypeOf((*MockStore)(nil).UpsertChatPersonalModelOverridesEnabled), ctx, enabled)
|
||||
}
|
||||
|
||||
// UpsertChatPlanModeInstructions mocks base method.
|
||||
func (m *MockStore) UpsertChatPlanModeInstructions(ctx context.Context, value string) error {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -10541,6 +10600,20 @@ func (mr *MockStoreMockRecorder) UpsertUserChatDebugLoggingEnabled(ctx, arg any)
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpsertUserChatDebugLoggingEnabled", reflect.TypeOf((*MockStore)(nil).UpsertUserChatDebugLoggingEnabled), ctx, arg)
|
||||
}
|
||||
|
||||
// UpsertUserChatPersonalModelOverride mocks base method.
|
||||
func (m *MockStore) UpsertUserChatPersonalModelOverride(ctx context.Context, arg database.UpsertUserChatPersonalModelOverrideParams) error {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "UpsertUserChatPersonalModelOverride", ctx, arg)
|
||||
ret0, _ := ret[0].(error)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// UpsertUserChatPersonalModelOverride indicates an expected call of UpsertUserChatPersonalModelOverride.
|
||||
func (mr *MockStoreMockRecorder) UpsertUserChatPersonalModelOverride(ctx, arg any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpsertUserChatPersonalModelOverride", reflect.TypeOf((*MockStore)(nil).UpsertUserChatPersonalModelOverride), ctx, arg)
|
||||
}
|
||||
|
||||
// UpsertUserChatProviderKey mocks base method.
|
||||
func (m *MockStore) UpsertUserChatProviderKey(ctx context.Context, arg database.UpsertUserChatProviderKeyParams) (database.UserChatProviderKey, error) {
|
||||
m.ctrl.T.Helper()
|
||||
|
||||
@@ -341,6 +341,9 @@ type sqlcQuerier interface {
|
||||
GetChatModelConfigs(ctx context.Context) ([]ChatModelConfig, error)
|
||||
// Returns all model configurations for telemetry snapshot collection.
|
||||
GetChatModelConfigsForTelemetry(ctx context.Context) ([]GetChatModelConfigsForTelemetryRow, error)
|
||||
// GetChatPersonalModelOverridesEnabled returns whether users may configure
|
||||
// personal chat model overrides. It defaults to false when unset.
|
||||
GetChatPersonalModelOverridesEnabled(ctx context.Context) (bool, error)
|
||||
GetChatPlanModeInstructions(ctx context.Context) (string, error)
|
||||
GetChatProviderByID(ctx context.Context, id uuid.UUID) (ChatProvider, error)
|
||||
GetChatProviderByIDForUpdate(ctx context.Context, id uuid.UUID) (ChatProvider, error)
|
||||
@@ -689,6 +692,7 @@ type sqlcQuerier interface {
|
||||
GetUserChatCompactionThreshold(ctx context.Context, arg GetUserChatCompactionThresholdParams) (string, error)
|
||||
GetUserChatCustomPrompt(ctx context.Context, userID uuid.UUID) (string, error)
|
||||
GetUserChatDebugLoggingEnabled(ctx context.Context, userID uuid.UUID) (bool, error)
|
||||
GetUserChatPersonalModelOverride(ctx context.Context, arg GetUserChatPersonalModelOverrideParams) (string, error)
|
||||
GetUserChatProviderKeys(ctx context.Context, userID uuid.UUID) ([]UserChatProviderKey, error)
|
||||
// Returns the total spend for a user in the given period.
|
||||
// When organization_id is NULL, spend across all organizations is
|
||||
@@ -944,6 +948,7 @@ type sqlcQuerier interface {
|
||||
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)
|
||||
ListUserChatPersonalModelOverrides(ctx context.Context, userID uuid.UUID) ([]ListUserChatPersonalModelOverridesRow, error)
|
||||
// Returns metadata only (no value or value_key_id) for the
|
||||
// REST API list and get endpoints.
|
||||
ListUserSecrets(ctx context.Context, userID uuid.UUID) ([]ListUserSecretsRow, error)
|
||||
@@ -1205,6 +1210,9 @@ type sqlcQuerier interface {
|
||||
UpsertChatExploreModelOverride(ctx context.Context, value string) error
|
||||
UpsertChatGeneralModelOverride(ctx context.Context, value string) error
|
||||
UpsertChatIncludeDefaultSystemPrompt(ctx context.Context, includeDefaultSystemPrompt bool) error
|
||||
// UpsertChatPersonalModelOverridesEnabled updates whether users may configure
|
||||
// personal chat model overrides.
|
||||
UpsertChatPersonalModelOverridesEnabled(ctx context.Context, enabled bool) error
|
||||
UpsertChatPlanModeInstructions(ctx context.Context, value string) error
|
||||
UpsertChatRetentionDays(ctx context.Context, retentionDays int32) error
|
||||
UpsertChatSystemPrompt(ctx context.Context, value string) error
|
||||
@@ -1241,6 +1249,7 @@ type sqlcQuerier interface {
|
||||
// combination. The result is stored in the template_usage_stats table.
|
||||
UpsertTemplateUsageStats(ctx context.Context) error
|
||||
UpsertUserChatDebugLoggingEnabled(ctx context.Context, arg UpsertUserChatDebugLoggingEnabledParams) error
|
||||
UpsertUserChatPersonalModelOverride(ctx context.Context, arg UpsertUserChatPersonalModelOverrideParams) error
|
||||
UpsertUserChatProviderKey(ctx context.Context, arg UpsertUserChatProviderKeyParams) (UserChatProviderKey, error)
|
||||
UpsertWebpushVAPIDKeys(ctx context.Context, arg UpsertWebpushVAPIDKeysParams) error
|
||||
UpsertWorkspaceAgentPortShare(ctx context.Context, arg UpsertWorkspaceAgentPortShareParams) (WorkspaceAgentPortShare, error)
|
||||
|
||||
@@ -20655,6 +20655,20 @@ func (q *sqlQuerier) GetChatIncludeDefaultSystemPrompt(ctx context.Context) (boo
|
||||
return include_default_system_prompt, err
|
||||
}
|
||||
|
||||
const getChatPersonalModelOverridesEnabled = `-- name: GetChatPersonalModelOverridesEnabled :one
|
||||
SELECT
|
||||
COALESCE((SELECT value = 'true' FROM site_configs WHERE key = 'agents_chat_personal_model_overrides_enabled'), false) :: boolean AS enabled
|
||||
`
|
||||
|
||||
// GetChatPersonalModelOverridesEnabled returns whether users may configure
|
||||
// personal chat model overrides. It defaults to false when unset.
|
||||
func (q *sqlQuerier) GetChatPersonalModelOverridesEnabled(ctx context.Context) (bool, error) {
|
||||
row := q.db.QueryRowContext(ctx, getChatPersonalModelOverridesEnabled)
|
||||
var enabled bool
|
||||
err := row.Scan(&enabled)
|
||||
return enabled, err
|
||||
}
|
||||
|
||||
const getChatPlanModeInstructions = `-- name: GetChatPlanModeInstructions :one
|
||||
SELECT
|
||||
COALESCE((SELECT value FROM site_configs WHERE key = 'agents_chat_plan_mode_instructions'), '') :: text AS plan_mode_instructions
|
||||
@@ -21077,6 +21091,30 @@ func (q *sqlQuerier) UpsertChatIncludeDefaultSystemPrompt(ctx context.Context, i
|
||||
return err
|
||||
}
|
||||
|
||||
const upsertChatPersonalModelOverridesEnabled = `-- name: UpsertChatPersonalModelOverridesEnabled :exec
|
||||
INSERT INTO site_configs (key, value)
|
||||
VALUES (
|
||||
'agents_chat_personal_model_overrides_enabled',
|
||||
CASE
|
||||
WHEN $1::bool THEN 'true'
|
||||
ELSE 'false'
|
||||
END
|
||||
)
|
||||
ON CONFLICT (key) DO UPDATE
|
||||
SET value = CASE
|
||||
WHEN $1::bool THEN 'true'
|
||||
ELSE 'false'
|
||||
END
|
||||
WHERE site_configs.key = 'agents_chat_personal_model_overrides_enabled'
|
||||
`
|
||||
|
||||
// UpsertChatPersonalModelOverridesEnabled updates whether users may configure
|
||||
// personal chat model overrides.
|
||||
func (q *sqlQuerier) UpsertChatPersonalModelOverridesEnabled(ctx context.Context, enabled bool) error {
|
||||
_, err := q.db.ExecContext(ctx, upsertChatPersonalModelOverridesEnabled, enabled)
|
||||
return err
|
||||
}
|
||||
|
||||
const upsertChatPlanModeInstructions = `-- name: UpsertChatPlanModeInstructions :exec
|
||||
INSERT INTO site_configs (key, value) VALUES ('agents_chat_plan_mode_instructions', $1)
|
||||
ON CONFLICT (key) DO UPDATE SET value = $1 WHERE site_configs.key = 'agents_chat_plan_mode_instructions'
|
||||
@@ -25402,6 +25440,24 @@ func (q *sqlQuerier) GetUserChatDebugLoggingEnabled(ctx context.Context, userID
|
||||
return debug_logging_enabled, err
|
||||
}
|
||||
|
||||
const getUserChatPersonalModelOverride = `-- name: GetUserChatPersonalModelOverride :one
|
||||
SELECT value AS personal_model_override FROM user_configs
|
||||
WHERE user_id = $1
|
||||
AND key = $2
|
||||
`
|
||||
|
||||
type GetUserChatPersonalModelOverrideParams struct {
|
||||
UserID uuid.UUID `db:"user_id" json:"user_id"`
|
||||
Key string `db:"key" json:"key"`
|
||||
}
|
||||
|
||||
func (q *sqlQuerier) GetUserChatPersonalModelOverride(ctx context.Context, arg GetUserChatPersonalModelOverrideParams) (string, error) {
|
||||
row := q.db.QueryRowContext(ctx, getUserChatPersonalModelOverride, arg.UserID, arg.Key)
|
||||
var personal_model_override string
|
||||
err := row.Scan(&personal_model_override)
|
||||
return personal_model_override, err
|
||||
}
|
||||
|
||||
const getUserCount = `-- name: GetUserCount :one
|
||||
SELECT
|
||||
COUNT(*)
|
||||
@@ -25864,6 +25920,41 @@ func (q *sqlQuerier) ListUserChatCompactionThresholds(ctx context.Context, userI
|
||||
return items, nil
|
||||
}
|
||||
|
||||
const listUserChatPersonalModelOverrides = `-- name: ListUserChatPersonalModelOverrides :many
|
||||
SELECT key, value FROM user_configs
|
||||
WHERE user_id = $1
|
||||
AND key LIKE 'chat\_personal\_model\_override:%'
|
||||
ORDER BY key
|
||||
`
|
||||
|
||||
type ListUserChatPersonalModelOverridesRow struct {
|
||||
Key string `db:"key" json:"key"`
|
||||
Value string `db:"value" json:"value"`
|
||||
}
|
||||
|
||||
func (q *sqlQuerier) ListUserChatPersonalModelOverrides(ctx context.Context, userID uuid.UUID) ([]ListUserChatPersonalModelOverridesRow, error) {
|
||||
rows, err := q.db.QueryContext(ctx, listUserChatPersonalModelOverrides, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var items []ListUserChatPersonalModelOverridesRow
|
||||
for rows.Next() {
|
||||
var i ListUserChatPersonalModelOverridesRow
|
||||
if err := rows.Scan(&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
|
||||
@@ -26469,6 +26560,24 @@ func (q *sqlQuerier) UpsertUserChatDebugLoggingEnabled(ctx context.Context, arg
|
||||
return err
|
||||
}
|
||||
|
||||
const upsertUserChatPersonalModelOverride = `-- name: UpsertUserChatPersonalModelOverride :exec
|
||||
INSERT INTO user_configs (user_id, key, value)
|
||||
VALUES ($1::uuid, $2::text, $3::text)
|
||||
ON CONFLICT ON CONSTRAINT user_configs_pkey
|
||||
DO UPDATE SET value = $3::text
|
||||
`
|
||||
|
||||
type UpsertUserChatPersonalModelOverrideParams struct {
|
||||
UserID uuid.UUID `db:"user_id" json:"user_id"`
|
||||
Key string `db:"key" json:"key"`
|
||||
Value string `db:"value" json:"value"`
|
||||
}
|
||||
|
||||
func (q *sqlQuerier) UpsertUserChatPersonalModelOverride(ctx context.Context, arg UpsertUserChatPersonalModelOverrideParams) error {
|
||||
_, err := q.db.ExecContext(ctx, upsertUserChatPersonalModelOverride, arg.UserID, arg.Key, arg.Value)
|
||||
return err
|
||||
}
|
||||
|
||||
const validateUserIDs = `-- name: ValidateUserIDs :one
|
||||
WITH input AS (
|
||||
SELECT
|
||||
|
||||
@@ -259,6 +259,30 @@ SET value = CASE
|
||||
END
|
||||
WHERE site_configs.key = 'agents_chat_debug_logging_allow_users';
|
||||
|
||||
-- GetChatPersonalModelOverridesEnabled returns whether users may configure
|
||||
-- personal chat model overrides. It defaults to false when unset.
|
||||
-- name: GetChatPersonalModelOverridesEnabled :one
|
||||
SELECT
|
||||
COALESCE((SELECT value = 'true' FROM site_configs WHERE key = 'agents_chat_personal_model_overrides_enabled'), false) :: boolean AS enabled;
|
||||
|
||||
-- UpsertChatPersonalModelOverridesEnabled updates whether users may configure
|
||||
-- personal chat model overrides.
|
||||
-- name: UpsertChatPersonalModelOverridesEnabled :exec
|
||||
INSERT INTO site_configs (key, value)
|
||||
VALUES (
|
||||
'agents_chat_personal_model_overrides_enabled',
|
||||
CASE
|
||||
WHEN sqlc.arg(enabled)::bool THEN 'true'
|
||||
ELSE 'false'
|
||||
END
|
||||
)
|
||||
ON CONFLICT (key) DO UPDATE
|
||||
SET value = CASE
|
||||
WHEN sqlc.arg(enabled)::bool THEN 'true'
|
||||
ELSE 'false'
|
||||
END
|
||||
WHERE site_configs.key = 'agents_chat_personal_model_overrides_enabled';
|
||||
|
||||
-- GetChatTemplateAllowlist returns the JSON-encoded template allowlist.
|
||||
-- Returns an empty string when no allowlist has been configured (all templates allowed).
|
||||
-- name: GetChatTemplateAllowlist :one
|
||||
|
||||
@@ -240,6 +240,23 @@ END
|
||||
WHERE user_configs.user_id = @user_id
|
||||
AND user_configs.key = 'chat_debug_logging_enabled';
|
||||
|
||||
-- name: ListUserChatPersonalModelOverrides :many
|
||||
SELECT key, value FROM user_configs
|
||||
WHERE user_id = @user_id
|
||||
AND key LIKE 'chat\_personal\_model\_override:%'
|
||||
ORDER BY key;
|
||||
|
||||
-- name: GetUserChatPersonalModelOverride :one
|
||||
SELECT value AS personal_model_override FROM user_configs
|
||||
WHERE user_id = @user_id
|
||||
AND key = @key;
|
||||
|
||||
-- name: UpsertUserChatPersonalModelOverride :exec
|
||||
INSERT INTO user_configs (user_id, key, value)
|
||||
VALUES (@user_id::uuid, @key::text, @value::text)
|
||||
ON CONFLICT ON CONSTRAINT user_configs_pkey
|
||||
DO UPDATE SET value = @value::text;
|
||||
|
||||
-- name: GetUserTaskNotificationAlertDismissed :one
|
||||
SELECT
|
||||
value::boolean as task_notification_alert_dismissed
|
||||
|
||||
+619
-75
@@ -12,6 +12,7 @@ import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
@@ -601,6 +602,306 @@ func (api *API) upsertChatModelOverrideConfig(
|
||||
return siteConfig.label, siteConfig.upsert(ctx, formatChatModelOverride(modelConfigID))
|
||||
}
|
||||
|
||||
var chatPersonalModelOverrideContexts = []codersdk.ChatPersonalModelOverrideContext{
|
||||
codersdk.ChatPersonalModelOverrideContextRoot,
|
||||
codersdk.ChatPersonalModelOverrideContextGeneral,
|
||||
codersdk.ChatPersonalModelOverrideContextExplore,
|
||||
}
|
||||
|
||||
func parseChatPersonalModelOverrideContext(raw string) (codersdk.ChatPersonalModelOverrideContext, bool) {
|
||||
c := codersdk.ChatPersonalModelOverrideContext(raw)
|
||||
return c, slices.Contains(chatPersonalModelOverrideContexts, c)
|
||||
}
|
||||
|
||||
func chatPersonalModelOverrideContextsJoined() string {
|
||||
values := make([]string, 0, len(chatPersonalModelOverrideContexts))
|
||||
for _, overrideContext := range chatPersonalModelOverrideContexts {
|
||||
values = append(values, string(overrideContext))
|
||||
}
|
||||
return strings.Join(values, ", ")
|
||||
}
|
||||
|
||||
func defaultChatPersonalModelOverrideMode(
|
||||
overrideContext codersdk.ChatPersonalModelOverrideContext,
|
||||
) codersdk.ChatPersonalModelOverrideMode {
|
||||
if overrideContext == codersdk.ChatPersonalModelOverrideContextRoot {
|
||||
return codersdk.ChatPersonalModelOverrideModeChatDefault
|
||||
}
|
||||
return codersdk.ChatPersonalModelOverrideModeDeploymentDefault
|
||||
}
|
||||
|
||||
func parseChatPersonalModelOverrideValue(
|
||||
raw string,
|
||||
overrideContext codersdk.ChatPersonalModelOverrideContext,
|
||||
) chatd.ParsedChatPersonalModelOverride {
|
||||
defaultMode := defaultChatPersonalModelOverrideMode(overrideContext)
|
||||
parsed := chatd.ParseChatPersonalModelOverride(raw, defaultMode)
|
||||
if overrideContext == codersdk.ChatPersonalModelOverrideContextRoot &&
|
||||
parsed.Mode == codersdk.ChatPersonalModelOverrideModeDeploymentDefault {
|
||||
return chatd.ParsedChatPersonalModelOverride{
|
||||
Mode: defaultMode,
|
||||
Malformed: true,
|
||||
}
|
||||
}
|
||||
return parsed
|
||||
}
|
||||
|
||||
func formatChatPersonalModelOverrideValue(
|
||||
mode codersdk.ChatPersonalModelOverrideMode,
|
||||
modelConfigID string,
|
||||
) string {
|
||||
if mode == codersdk.ChatPersonalModelOverrideModeModel {
|
||||
return string(mode) + ":" + strings.TrimSpace(modelConfigID)
|
||||
}
|
||||
return string(mode)
|
||||
}
|
||||
|
||||
func chatPersonalModelOverrideResponse(
|
||||
overrideContext codersdk.ChatPersonalModelOverrideContext,
|
||||
raw string,
|
||||
isSet bool,
|
||||
) codersdk.ChatPersonalModelOverride {
|
||||
parsed := parseChatPersonalModelOverrideValue(raw, overrideContext)
|
||||
modelConfigID := ""
|
||||
if parsed.Mode == codersdk.ChatPersonalModelOverrideModeModel {
|
||||
modelConfigID = parsed.ModelConfigID.String()
|
||||
}
|
||||
return codersdk.ChatPersonalModelOverride{
|
||||
Context: overrideContext,
|
||||
Mode: parsed.Mode,
|
||||
ModelConfigID: modelConfigID,
|
||||
IsSet: isSet,
|
||||
IsMalformed: parsed.Malformed,
|
||||
}
|
||||
}
|
||||
|
||||
func (api *API) chatPersonalModelOverrideDeploymentDefaultResponse(
|
||||
ctx context.Context,
|
||||
overrideContext codersdk.ChatModelOverrideContext,
|
||||
) (codersdk.ChatModelOverrideResponse, error) {
|
||||
// The deployment defaults are global chat configuration, not user-owned
|
||||
// resources. Users may read these values here because the personal settings
|
||||
// UI must explain what deployment_default resolves to.
|
||||
//nolint:gocritic // System context is required to read deployment config.
|
||||
modelConfigID, isMalformed, _, err := api.readChatModelOverrideConfig(
|
||||
dbauthz.AsSystemRestricted(ctx),
|
||||
overrideContext,
|
||||
)
|
||||
if err != nil {
|
||||
return codersdk.ChatModelOverrideResponse{}, err
|
||||
}
|
||||
return codersdk.ChatModelOverrideResponse{
|
||||
Context: overrideContext,
|
||||
ModelConfigID: formatChatModelOverride(modelConfigID),
|
||||
IsMalformed: isMalformed,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (api *API) chatPersonalModelOverrideDeploymentDefaults(
|
||||
ctx context.Context,
|
||||
) (codersdk.ChatPersonalModelOverrideDeploymentDefaults, error) {
|
||||
general, err := api.chatPersonalModelOverrideDeploymentDefaultResponse(
|
||||
ctx,
|
||||
codersdk.ChatModelOverrideContextGeneral,
|
||||
)
|
||||
if err != nil {
|
||||
return codersdk.ChatPersonalModelOverrideDeploymentDefaults{}, err
|
||||
}
|
||||
explore, err := api.chatPersonalModelOverrideDeploymentDefaultResponse(
|
||||
ctx,
|
||||
codersdk.ChatModelOverrideContextExplore,
|
||||
)
|
||||
if err != nil {
|
||||
return codersdk.ChatPersonalModelOverrideDeploymentDefaults{}, err
|
||||
}
|
||||
return codersdk.ChatPersonalModelOverrideDeploymentDefaults{
|
||||
General: general,
|
||||
Explore: explore,
|
||||
}, nil
|
||||
}
|
||||
|
||||
type userChatModelAvailability struct {
|
||||
configuredProviders []chatprovider.ConfiguredProvider
|
||||
configuredModels []chatprovider.ConfiguredModel
|
||||
enabledModels []database.ChatModelConfig
|
||||
providerStatus map[string]chatprovider.ProviderAvailability
|
||||
enabledProviderNames map[string]struct{}
|
||||
}
|
||||
|
||||
// chatModelConfigUnavailableReason reports why a model config cannot be used.
|
||||
// The empty value means the model config is available. Callers must check the
|
||||
// error returned by userCanUseChatModelConfig before interpreting this value.
|
||||
type chatModelConfigUnavailableReason string
|
||||
|
||||
const (
|
||||
chatModelConfigAvailable chatModelConfigUnavailableReason = ""
|
||||
chatModelConfigUnavailableModelNotFoundOrDisabled chatModelConfigUnavailableReason = "model_not_found_or_disabled"
|
||||
chatModelConfigUnavailableProviderDisabled chatModelConfigUnavailableReason = "provider_disabled"
|
||||
chatModelConfigUnavailableCredentialsMissing chatModelConfigUnavailableReason = "credentials_missing"
|
||||
)
|
||||
|
||||
// getUserChatProviderAvailability returns chat provider availability for a
|
||||
// user. Deployment-level enabled providers and models are read with
|
||||
// dbauthz.AsSystemRestricted(ctx) because they are global chat configuration,
|
||||
// not user-owned resources. Callers must pass an authenticated user context so
|
||||
// user-scoped model checks and provider-key lookups run under the caller's
|
||||
// authorization. The returned struct contains configured providers and models
|
||||
// for catalog listing, enabled model rows for ID validation, resolved provider
|
||||
// status, and normalized enabled-provider membership.
|
||||
func (api *API) getUserChatProviderAvailability(
|
||||
ctx context.Context,
|
||||
userID uuid.UUID,
|
||||
) (userChatModelAvailability, error) {
|
||||
//nolint:gocritic // System context is required to read enabled chat config.
|
||||
systemCtx := dbauthz.AsSystemRestricted(ctx)
|
||||
enabledProviders, err := api.Database.GetEnabledChatProviders(systemCtx)
|
||||
if err != nil {
|
||||
return userChatModelAvailability{}, err
|
||||
}
|
||||
enabledModels, err := api.Database.GetEnabledChatModelConfigs(systemCtx)
|
||||
if err != nil {
|
||||
return userChatModelAvailability{}, err
|
||||
}
|
||||
|
||||
availability := userChatModelAvailability{
|
||||
configuredProviders: make([]chatprovider.ConfiguredProvider, 0, len(enabledProviders)),
|
||||
configuredModels: make([]chatprovider.ConfiguredModel, 0, len(enabledModels)),
|
||||
enabledModels: enabledModels,
|
||||
enabledProviderNames: make(map[string]struct{}, len(enabledProviders)),
|
||||
}
|
||||
for _, provider := range enabledProviders {
|
||||
availability.configuredProviders = append(
|
||||
availability.configuredProviders,
|
||||
chatprovider.ConfiguredProvider{
|
||||
ProviderID: provider.ID,
|
||||
Provider: provider.Provider,
|
||||
APIKey: provider.APIKey,
|
||||
BaseURL: provider.BaseUrl,
|
||||
CentralAPIKeyEnabled: provider.CentralApiKeyEnabled,
|
||||
AllowUserAPIKey: provider.AllowUserApiKey,
|
||||
AllowCentralAPIKeyFallback: provider.AllowCentralApiKeyFallback,
|
||||
},
|
||||
)
|
||||
normalizedProvider := chatprovider.NormalizeProvider(provider.Provider)
|
||||
if normalizedProvider != "" {
|
||||
availability.enabledProviderNames[normalizedProvider] = struct{}{}
|
||||
}
|
||||
}
|
||||
for _, model := range enabledModels {
|
||||
availability.configuredModels = append(availability.configuredModels, chatprovider.ConfiguredModel{
|
||||
Provider: model.Provider,
|
||||
Model: model.Model,
|
||||
DisplayName: model.DisplayName,
|
||||
})
|
||||
}
|
||||
|
||||
userKeyRows, err := api.Database.GetUserChatProviderKeys(ctx, userID)
|
||||
if err != nil {
|
||||
return userChatModelAvailability{}, err
|
||||
}
|
||||
userKeys := make([]chatprovider.UserProviderKey, 0, len(userKeyRows))
|
||||
for _, userKey := range userKeyRows {
|
||||
userKeys = append(userKeys, chatprovider.UserProviderKey{
|
||||
ChatProviderID: userKey.ChatProviderID,
|
||||
APIKey: userKey.APIKey,
|
||||
})
|
||||
}
|
||||
|
||||
_, availability.providerStatus = chatprovider.ResolveUserProviderKeys(
|
||||
ChatProviderAPIKeysFromDeploymentValues(api.DeploymentValues),
|
||||
availability.configuredProviders,
|
||||
userKeys,
|
||||
)
|
||||
return availability, nil
|
||||
}
|
||||
|
||||
// userCanUseChatModelConfig returns chatModelConfigAvailable when the user can
|
||||
// use the model config. If err is non-nil, callers must ignore the returned
|
||||
// reason because it may be the zero-value availability sentinel.
|
||||
func (api *API) userCanUseChatModelConfig(
|
||||
ctx context.Context,
|
||||
userID uuid.UUID,
|
||||
modelConfigID uuid.UUID,
|
||||
) (chatModelConfigUnavailableReason, error) {
|
||||
if modelConfigID == uuid.Nil {
|
||||
return chatModelConfigUnavailableModelNotFoundOrDisabled, nil
|
||||
}
|
||||
//nolint:gocritic // Non-admin users need deployment config validation.
|
||||
model, err := api.Database.GetChatModelConfigByID(
|
||||
dbauthz.AsSystemRestricted(ctx),
|
||||
modelConfigID,
|
||||
)
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) || httpapi.Is404Error(err) {
|
||||
return chatModelConfigUnavailableModelNotFoundOrDisabled, nil
|
||||
}
|
||||
return chatModelConfigAvailable, err
|
||||
}
|
||||
if !model.Enabled {
|
||||
return chatModelConfigUnavailableModelNotFoundOrDisabled, nil
|
||||
}
|
||||
|
||||
availability, err := api.getUserChatProviderAvailability(ctx, userID)
|
||||
if err != nil {
|
||||
return chatModelConfigAvailable, err
|
||||
}
|
||||
provider, _, err := chatprovider.ResolveModelWithProviderHint(model.Model, model.Provider)
|
||||
if err != nil {
|
||||
return chatModelConfigUnavailableProviderDisabled, nil
|
||||
}
|
||||
if _, ok := availability.enabledProviderNames[provider]; !ok {
|
||||
return chatModelConfigUnavailableProviderDisabled, nil
|
||||
}
|
||||
providerStatus, ok := availability.providerStatus[provider]
|
||||
if !ok {
|
||||
return chatModelConfigUnavailableProviderDisabled, nil
|
||||
}
|
||||
if !providerStatus.Available {
|
||||
return chatModelConfigUnavailableCredentialsMissing, nil
|
||||
}
|
||||
return chatModelConfigAvailable, nil
|
||||
}
|
||||
|
||||
func (api *API) validateUserChatModelConfigAvailable(
|
||||
ctx context.Context,
|
||||
userID uuid.UUID,
|
||||
modelConfigID uuid.UUID,
|
||||
) (int, *codersdk.Response) {
|
||||
reason, err := api.userCanUseChatModelConfig(ctx, userID, modelConfigID)
|
||||
if err != nil {
|
||||
return http.StatusInternalServerError, &codersdk.Response{
|
||||
Message: "Internal error validating model config override.",
|
||||
Detail: err.Error(),
|
||||
}
|
||||
}
|
||||
switch reason {
|
||||
case chatModelConfigAvailable:
|
||||
return 0, nil
|
||||
case chatModelConfigUnavailableModelNotFoundOrDisabled:
|
||||
return http.StatusBadRequest, &codersdk.Response{
|
||||
Message: "Invalid model_config_id: model config not found or disabled.",
|
||||
}
|
||||
case chatModelConfigUnavailableCredentialsMissing:
|
||||
return http.StatusBadRequest, &codersdk.Response{
|
||||
Message: "Invalid model_config_id: provider credentials unavailable for this model.",
|
||||
}
|
||||
case chatModelConfigUnavailableProviderDisabled:
|
||||
return http.StatusBadRequest, &codersdk.Response{
|
||||
Message: "Invalid model_config_id: provider is not enabled for this model.",
|
||||
}
|
||||
default:
|
||||
api.Logger.Warn(ctx,
|
||||
"unknown chat model config availability reason",
|
||||
slog.F("user_id", userID),
|
||||
slog.F("model_config_id", modelConfigID),
|
||||
slog.F("reason", reason),
|
||||
)
|
||||
return http.StatusBadRequest, &codersdk.Response{
|
||||
Message: "Invalid model_config_id.",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// EXPERIMENTAL: this endpoint is experimental and is subject to change.
|
||||
func (api *API) postChats(rw http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
@@ -685,7 +986,7 @@ func (api *API) postChats(rw http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
modelConfigID, modelConfigStatus, modelConfigError := api.resolveCreateChatModelConfigID(ctx, req)
|
||||
modelConfigID, modelConfigStatus, modelConfigError := api.resolveCreateChatModelConfigID(ctx, apiKey.UserID, req)
|
||||
if modelConfigError != nil {
|
||||
httpapi.Write(ctx, rw, modelConfigStatus, *modelConfigError)
|
||||
return
|
||||
@@ -886,9 +1187,6 @@ func (api *API) postChats(rw http.ResponseWriter, r *http.Request) {
|
||||
func (api *API) listChatModels(rw http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
apiKey := httpmw.APIKey(r)
|
||||
//nolint:gocritic // System context required to read enabled chat models.
|
||||
systemCtx := dbauthz.AsSystemRestricted(ctx)
|
||||
|
||||
if api.chatDaemon == nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Chat processor is unavailable.",
|
||||
@@ -897,9 +1195,7 @@ func (api *API) listChatModels(rw http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
enabledProviders, err := api.Database.GetEnabledChatProviders(
|
||||
systemCtx,
|
||||
)
|
||||
availability, err := api.getUserChatProviderAvailability(ctx, apiKey.UserID)
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Failed to load chat model configuration.",
|
||||
@@ -907,81 +1203,19 @@ func (api *API) listChatModels(rw http.ResponseWriter, r *http.Request) {
|
||||
})
|
||||
return
|
||||
}
|
||||
enabledModels, err := api.Database.GetEnabledChatModelConfigs(
|
||||
systemCtx,
|
||||
)
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Failed to load chat model configuration.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
configuredProviders := make(
|
||||
[]chatprovider.ConfiguredProvider, 0, len(enabledProviders),
|
||||
)
|
||||
enabledProviderNames := make(map[string]struct{}, len(enabledProviders))
|
||||
for _, provider := range enabledProviders {
|
||||
configuredProviders = append(
|
||||
configuredProviders, chatprovider.ConfiguredProvider{
|
||||
ProviderID: provider.ID,
|
||||
Provider: provider.Provider,
|
||||
APIKey: provider.APIKey,
|
||||
BaseURL: provider.BaseUrl,
|
||||
CentralAPIKeyEnabled: provider.CentralApiKeyEnabled,
|
||||
AllowUserAPIKey: provider.AllowUserApiKey,
|
||||
AllowCentralAPIKeyFallback: provider.AllowCentralApiKeyFallback,
|
||||
},
|
||||
)
|
||||
normalizedProvider := chatprovider.NormalizeProvider(provider.Provider)
|
||||
if normalizedProvider == "" {
|
||||
continue
|
||||
}
|
||||
enabledProviderNames[normalizedProvider] = struct{}{}
|
||||
}
|
||||
configuredModels := make(
|
||||
[]chatprovider.ConfiguredModel, 0, len(enabledModels),
|
||||
)
|
||||
for _, model := range enabledModels {
|
||||
configuredModels = append(configuredModels, chatprovider.ConfiguredModel{
|
||||
Provider: model.Provider,
|
||||
Model: model.Model,
|
||||
DisplayName: model.DisplayName,
|
||||
})
|
||||
}
|
||||
|
||||
userKeyRows, err := api.Database.GetUserChatProviderKeys(ctx, apiKey.UserID)
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Failed to load user chat provider keys.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
userKeys := make([]chatprovider.UserProviderKey, 0, len(userKeyRows))
|
||||
for _, userKey := range userKeyRows {
|
||||
userKeys = append(userKeys, chatprovider.UserProviderKey{
|
||||
ChatProviderID: userKey.ChatProviderID,
|
||||
APIKey: userKey.APIKey,
|
||||
})
|
||||
}
|
||||
|
||||
_, providerAvailability := chatprovider.ResolveUserProviderKeys(
|
||||
ChatProviderAPIKeysFromDeploymentValues(api.DeploymentValues),
|
||||
configuredProviders,
|
||||
userKeys,
|
||||
)
|
||||
catalog := chatprovider.NewModelCatalog()
|
||||
var response codersdk.ChatModelsResponse
|
||||
if configured, ok := catalog.ListConfiguredModels(
|
||||
configuredProviders, configuredModels, providerAvailability, enabledProviderNames,
|
||||
availability.configuredProviders,
|
||||
availability.configuredModels,
|
||||
availability.providerStatus,
|
||||
availability.enabledProviderNames,
|
||||
); ok {
|
||||
response = configured
|
||||
} else {
|
||||
response = catalog.ListConfiguredProviderAvailability(
|
||||
providerAvailability,
|
||||
enabledProviderNames,
|
||||
availability.providerStatus,
|
||||
availability.enabledProviderNames,
|
||||
)
|
||||
}
|
||||
|
||||
@@ -3758,6 +3992,7 @@ func (api *API) validateCreateChatWorkspaceSelection(
|
||||
|
||||
func (api *API) resolveCreateChatModelConfigID(
|
||||
ctx context.Context,
|
||||
userID uuid.UUID,
|
||||
req codersdk.CreateChatRequest,
|
||||
) (uuid.UUID, int, *codersdk.Response) {
|
||||
if req.ModelConfigID != nil {
|
||||
@@ -3769,6 +4004,82 @@ func (api *API) resolveCreateChatModelConfigID(
|
||||
return *req.ModelConfigID, 0, nil
|
||||
}
|
||||
|
||||
personalOverridesEnabled, err := api.Database.GetChatPersonalModelOverridesEnabled(ctx)
|
||||
if err != nil {
|
||||
return uuid.Nil, http.StatusInternalServerError, &codersdk.Response{
|
||||
Message: "Failed to resolve chat model config.",
|
||||
Detail: err.Error(),
|
||||
}
|
||||
}
|
||||
if !personalOverridesEnabled {
|
||||
return api.defaultCreateChatModelConfigID(ctx)
|
||||
}
|
||||
|
||||
raw, err := api.Database.GetUserChatPersonalModelOverride(ctx, database.GetUserChatPersonalModelOverrideParams{
|
||||
UserID: userID,
|
||||
Key: chatd.ChatPersonalModelOverrideKey(codersdk.ChatPersonalModelOverrideContextRoot),
|
||||
})
|
||||
if err != nil && !errors.Is(err, sql.ErrNoRows) {
|
||||
return uuid.Nil, http.StatusInternalServerError, &codersdk.Response{
|
||||
Message: "Failed to resolve chat model config.",
|
||||
Detail: err.Error(),
|
||||
}
|
||||
}
|
||||
if err == nil {
|
||||
parsed := parseChatPersonalModelOverrideValue(
|
||||
raw,
|
||||
codersdk.ChatPersonalModelOverrideContextRoot,
|
||||
)
|
||||
if parsed.Malformed {
|
||||
api.Logger.Debug(
|
||||
ctx,
|
||||
"unsupported personal root model override mode, using default model",
|
||||
slog.F("user_id", userID),
|
||||
slog.F("raw_value", raw),
|
||||
)
|
||||
}
|
||||
switch parsed.Mode {
|
||||
case codersdk.ChatPersonalModelOverrideModeChatDefault:
|
||||
// For root context, chat_default and the defensive default
|
||||
// case both fall through to the deployment default model below.
|
||||
case codersdk.ChatPersonalModelOverrideModeModel:
|
||||
reason, err := api.userCanUseChatModelConfig(
|
||||
ctx,
|
||||
userID,
|
||||
parsed.ModelConfigID,
|
||||
)
|
||||
if err != nil {
|
||||
return uuid.Nil, http.StatusInternalServerError, &codersdk.Response{
|
||||
Message: "Failed to resolve chat model config.",
|
||||
Detail: err.Error(),
|
||||
}
|
||||
}
|
||||
if reason == chatModelConfigAvailable {
|
||||
return parsed.ModelConfigID, 0, nil
|
||||
}
|
||||
api.Logger.Debug(
|
||||
ctx,
|
||||
"personal root model override is unavailable, using default model",
|
||||
slog.F("user_id", userID),
|
||||
slog.F("model_config_id", parsed.ModelConfigID),
|
||||
slog.F("reason", reason),
|
||||
)
|
||||
default:
|
||||
api.Logger.Warn(
|
||||
ctx,
|
||||
"unsupported personal root model override mode, using default model",
|
||||
slog.F("user_id", userID),
|
||||
slog.F("mode", parsed.Mode),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
return api.defaultCreateChatModelConfigID(ctx)
|
||||
}
|
||||
|
||||
func (api *API) defaultCreateChatModelConfigID(
|
||||
ctx context.Context,
|
||||
) (uuid.UUID, int, *codersdk.Response) {
|
||||
defaultModelConfig, err := api.Database.GetDefaultChatModelConfig(ctx)
|
||||
if err != nil {
|
||||
if xerrors.Is(err, sql.ErrNoRows) {
|
||||
@@ -4064,6 +4375,239 @@ func (api *API) putChatModelOverride(rw http.ResponseWriter, r *http.Request) {
|
||||
rw.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
func readChatPersonalModelOverrideContext(
|
||||
rw http.ResponseWriter,
|
||||
r *http.Request,
|
||||
) (codersdk.ChatPersonalModelOverrideContext, bool) {
|
||||
ctx := r.Context()
|
||||
rawContext := chi.URLParam(r, "context")
|
||||
overrideContext, ok := parseChatPersonalModelOverrideContext(rawContext)
|
||||
if ok {
|
||||
return overrideContext, true
|
||||
}
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "Invalid chat personal model override context.",
|
||||
Detail: fmt.Sprintf(
|
||||
"Expected one of %s. Got %q.",
|
||||
chatPersonalModelOverrideContextsJoined(),
|
||||
rawContext,
|
||||
),
|
||||
})
|
||||
return "", false
|
||||
}
|
||||
|
||||
// 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) getChatPersonalModelOverridesAdminSettings(rw http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
if !api.Authorize(r, policy.ActionRead, rbac.ResourceDeploymentConfig) {
|
||||
httpapi.ResourceNotFound(rw)
|
||||
return
|
||||
}
|
||||
|
||||
enabled, err := api.Database.GetChatPersonalModelOverridesEnabled(ctx)
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Internal error fetching personal model override setting.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
httpapi.Write(ctx, rw, http.StatusOK, codersdk.ChatPersonalModelOverridesAdminSettings{
|
||||
AllowUsers: enabled,
|
||||
})
|
||||
}
|
||||
|
||||
// EXPERIMENTAL: this endpoint is experimental and is subject to change.
|
||||
func (api *API) putChatPersonalModelOverridesAdminSettings(rw http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
if !api.Authorize(r, policy.ActionUpdate, rbac.ResourceDeploymentConfig) {
|
||||
httpapi.Forbidden(rw)
|
||||
return
|
||||
}
|
||||
|
||||
var req codersdk.UpdateChatPersonalModelOverridesAdminSettingsRequest
|
||||
if !httpapi.Read(ctx, rw, r, &req) {
|
||||
return
|
||||
}
|
||||
if err := api.Database.UpsertChatPersonalModelOverridesEnabled(ctx, req.AllowUsers); err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Internal error updating personal model override setting.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
rw.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
// 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) getUserChatPersonalModelOverrides(rw http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
apiKey := httpmw.APIKey(r)
|
||||
|
||||
enabled, err := api.Database.GetChatPersonalModelOverridesEnabled(ctx)
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Internal error fetching personal model override setting.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
rows, err := api.Database.ListUserChatPersonalModelOverrides(ctx, apiKey.UserID)
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Internal error fetching user personal model overrides.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
values := make(map[codersdk.ChatPersonalModelOverrideContext]string, len(rows))
|
||||
for _, row := range rows {
|
||||
rawContext, ok := strings.CutPrefix(row.Key, chatd.ChatPersonalModelOverrideKeyPrefix)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
overrideContext, ok := parseChatPersonalModelOverrideContext(rawContext)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
values[overrideContext] = row.Value
|
||||
}
|
||||
|
||||
deploymentDefaults, err := api.chatPersonalModelOverrideDeploymentDefaults(ctx)
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Internal error fetching deployment model defaults.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
response := codersdk.UserChatPersonalModelOverridesResponse{
|
||||
Enabled: enabled,
|
||||
DeploymentDefaults: deploymentDefaults,
|
||||
}
|
||||
for _, overrideContext := range chatPersonalModelOverrideContexts {
|
||||
raw, isSet := values[overrideContext]
|
||||
override := chatPersonalModelOverrideResponse(overrideContext, raw, isSet)
|
||||
switch overrideContext {
|
||||
case codersdk.ChatPersonalModelOverrideContextRoot:
|
||||
response.Root = override
|
||||
case codersdk.ChatPersonalModelOverrideContextGeneral:
|
||||
response.General = override
|
||||
case codersdk.ChatPersonalModelOverrideContextExplore:
|
||||
response.Explore = override
|
||||
}
|
||||
}
|
||||
httpapi.Write(ctx, rw, http.StatusOK, response)
|
||||
}
|
||||
|
||||
// EXPERIMENTAL: this endpoint is experimental and is subject to change.
|
||||
func (api *API) putUserChatPersonalModelOverride(rw http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
apiKey := httpmw.APIKey(r)
|
||||
|
||||
enabled, err := api.Database.GetChatPersonalModelOverridesEnabled(ctx)
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Internal error fetching personal model override setting.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
if !enabled {
|
||||
httpapi.Write(ctx, rw, http.StatusForbidden, codersdk.Response{
|
||||
Message: "An administrator has not enabled user personal model overrides.",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
overrideContext, ok := readChatPersonalModelOverrideContext(rw, r)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
|
||||
var req codersdk.UpdateUserChatPersonalModelOverrideRequest
|
||||
if !httpapi.Read(ctx, rw, r, &req) {
|
||||
return
|
||||
}
|
||||
|
||||
modelConfigID := ""
|
||||
rawModelConfigID := strings.TrimSpace(req.ModelConfigID)
|
||||
switch req.Mode {
|
||||
case codersdk.ChatPersonalModelOverrideModeChatDefault:
|
||||
if rawModelConfigID != "" {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "model_config_id must be empty unless mode is model.",
|
||||
})
|
||||
return
|
||||
}
|
||||
case codersdk.ChatPersonalModelOverrideModeDeploymentDefault:
|
||||
if overrideContext == codersdk.ChatPersonalModelOverrideContextRoot {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "deployment_default is not supported for root personal model overrides.",
|
||||
})
|
||||
return
|
||||
}
|
||||
if rawModelConfigID != "" {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "model_config_id must be empty unless mode is model.",
|
||||
})
|
||||
return
|
||||
}
|
||||
case codersdk.ChatPersonalModelOverrideModeModel:
|
||||
if rawModelConfigID == "" {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "model_config_id is required when mode is model.",
|
||||
})
|
||||
return
|
||||
}
|
||||
parsedModelConfigID, err := uuid.Parse(rawModelConfigID)
|
||||
if err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "Invalid model_config_id.",
|
||||
Detail: fmt.Sprintf("Value %q is not a valid UUID.", req.ModelConfigID),
|
||||
})
|
||||
return
|
||||
}
|
||||
if parsedModelConfigID == uuid.Nil {
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "Invalid model_config_id.",
|
||||
})
|
||||
return
|
||||
}
|
||||
status, resp := api.validateUserChatModelConfigAvailable(ctx, apiKey.UserID, parsedModelConfigID)
|
||||
if resp != nil {
|
||||
httpapi.Write(ctx, rw, status, *resp)
|
||||
return
|
||||
}
|
||||
modelConfigID = parsedModelConfigID.String()
|
||||
default:
|
||||
httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{
|
||||
Message: "Invalid personal model override mode.",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
if err := api.Database.UpsertUserChatPersonalModelOverride(ctx, database.UpsertUserChatPersonalModelOverrideParams{
|
||||
UserID: apiKey.UserID,
|
||||
Key: chatd.ChatPersonalModelOverrideKey(overrideContext),
|
||||
Value: formatChatPersonalModelOverrideValue(req.Mode, modelConfigID),
|
||||
}); err != nil {
|
||||
httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{
|
||||
Message: "Internal error updating user personal model override.",
|
||||
Detail: err.Error(),
|
||||
})
|
||||
return
|
||||
}
|
||||
rw.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
// 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.
|
||||
|
||||
@@ -9555,6 +9555,39 @@ func createDisabledChatModelConfig(
|
||||
return updated
|
||||
}
|
||||
|
||||
func enableUserChatProviderKey(
|
||||
t testing.TB,
|
||||
adminClient *codersdk.ExperimentalClient,
|
||||
userClient *codersdk.ExperimentalClient,
|
||||
providerName string,
|
||||
) codersdk.ChatProviderConfig {
|
||||
t.Helper()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
providers, err := adminClient.ListChatProviders(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
var provider codersdk.ChatProviderConfig
|
||||
for _, candidate := range providers {
|
||||
if candidate.Provider == providerName && candidate.Source == codersdk.ChatProviderConfigSourceDatabase {
|
||||
provider = candidate
|
||||
break
|
||||
}
|
||||
}
|
||||
require.NotEqual(t, uuid.Nil, provider.ID)
|
||||
|
||||
updated, err := adminClient.UpdateChatProvider(ctx, provider.ID, codersdk.UpdateChatProviderConfigRequest{
|
||||
AllowUserAPIKey: ptr.Ref(true),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = userClient.UpsertUserChatProviderKey(ctx, updated.ID, codersdk.CreateUserChatProviderKeyRequest{
|
||||
APIKey: "test-user-api-key-" + uuid.NewString(),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
return updated
|
||||
}
|
||||
|
||||
//nolint:tparallel,paralleltest // Subtests share a single coderdtest instance.
|
||||
func TestChatSystemPrompt(t *testing.T) {
|
||||
t.Parallel()
|
||||
@@ -10313,6 +10346,537 @@ func TestChatModelOverrides(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
//nolint:tparallel,paralleltest // Subtests share coderdtest instances.
|
||||
func TestChatPersonalModelOverridesAdminSettings(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
adminClient := newChatClient(t)
|
||||
firstUser := coderdtest.CreateFirstUser(t, adminClient.Client)
|
||||
memberClientRaw, _ := coderdtest.CreateAnotherUser(t, adminClient.Client, firstUser.OrganizationID)
|
||||
memberClient := codersdk.NewExperimentalClient(memberClientRaw)
|
||||
|
||||
resp, err := adminClient.GetChatPersonalModelOverridesAdminSettings(ctx)
|
||||
require.NoError(t, err)
|
||||
require.False(t, resp.AllowUsers)
|
||||
|
||||
err = adminClient.UpdateChatPersonalModelOverridesAdminSettings(ctx, codersdk.UpdateChatPersonalModelOverridesAdminSettingsRequest{
|
||||
AllowUsers: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
resp, err = adminClient.GetChatPersonalModelOverridesAdminSettings(ctx)
|
||||
require.NoError(t, err)
|
||||
require.True(t, resp.AllowUsers)
|
||||
|
||||
err = adminClient.UpdateChatPersonalModelOverridesAdminSettings(ctx, codersdk.UpdateChatPersonalModelOverridesAdminSettingsRequest{
|
||||
AllowUsers: false,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
resp, err = adminClient.GetChatPersonalModelOverridesAdminSettings(ctx)
|
||||
require.NoError(t, err)
|
||||
require.False(t, resp.AllowUsers)
|
||||
|
||||
err = memberClient.UpdateChatPersonalModelOverridesAdminSettings(ctx, codersdk.UpdateChatPersonalModelOverridesAdminSettingsRequest{
|
||||
AllowUsers: true,
|
||||
})
|
||||
requireSDKError(t, err, http.StatusForbidden)
|
||||
|
||||
_, err = memberClient.GetChatPersonalModelOverridesAdminSettings(ctx)
|
||||
requireSDKError(t, err, http.StatusNotFound)
|
||||
}
|
||||
|
||||
//nolint:tparallel,paralleltest // Subtests share coderdtest instances.
|
||||
func TestUserChatPersonalModelOverrides(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
adminClient, db := newChatClientWithDatabase(t)
|
||||
firstUser := coderdtest.CreateFirstUser(t, adminClient.Client)
|
||||
memberClientRaw, member := coderdtest.CreateAnotherUser(t, adminClient.Client, firstUser.OrganizationID)
|
||||
memberClient := codersdk.NewExperimentalClient(memberClientRaw)
|
||||
noKeyClientRaw, noKeyUser := coderdtest.CreateAnotherUser(t, adminClient.Client, firstUser.OrganizationID)
|
||||
noKeyClient := codersdk.NewExperimentalClient(noKeyClientRaw)
|
||||
|
||||
defaultModelConfig := createChatModelConfig(t, adminClient)
|
||||
provider := enableUserChatProviderKey(t, adminClient, memberClient, "openai")
|
||||
modelConfig := createAdditionalChatModelConfig(
|
||||
t,
|
||||
adminClient,
|
||||
"openai",
|
||||
"gpt-4o-personal-"+uuid.NewString(),
|
||||
)
|
||||
err := adminClient.UpdateChatModelOverride(ctx, codersdk.ChatModelOverrideContextGeneral, codersdk.UpdateChatModelOverrideRequest{
|
||||
ModelConfigID: modelConfig.ID.String(),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
err = adminClient.UpdateChatModelOverride(ctx, codersdk.ChatModelOverrideContextExplore, codersdk.UpdateChatModelOverrideRequest{
|
||||
ModelConfigID: defaultModelConfig.ID.String(),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
disabledModelConfig := createDisabledChatModelConfig(
|
||||
t,
|
||||
adminClient,
|
||||
"openai",
|
||||
"gpt-4o-personal-disabled-"+uuid.NewString(),
|
||||
)
|
||||
disabledProvider, err := adminClient.CreateChatProvider(ctx, codersdk.CreateChatProviderConfigRequest{
|
||||
Provider: "anthropic",
|
||||
Enabled: ptr.Ref(false),
|
||||
CentralAPIKeyEnabled: ptr.Ref(false),
|
||||
AllowUserAPIKey: ptr.Ref(true),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
disabledProviderModelConfig := createAdditionalChatModelConfig(
|
||||
t,
|
||||
adminClient,
|
||||
"anthropic",
|
||||
"claude-personal-disabled-provider-"+uuid.NewString(),
|
||||
)
|
||||
require.NotEqual(t, uuid.Nil, provider.ID)
|
||||
require.NotEqual(t, uuid.Nil, disabledProvider.ID)
|
||||
|
||||
personalOverride := func(
|
||||
resp codersdk.UserChatPersonalModelOverridesResponse,
|
||||
overrideContext codersdk.ChatPersonalModelOverrideContext,
|
||||
) codersdk.ChatPersonalModelOverride {
|
||||
t.Helper()
|
||||
switch overrideContext {
|
||||
case codersdk.ChatPersonalModelOverrideContextRoot:
|
||||
return resp.Root
|
||||
case codersdk.ChatPersonalModelOverrideContextGeneral:
|
||||
return resp.General
|
||||
case codersdk.ChatPersonalModelOverrideContextExplore:
|
||||
return resp.Explore
|
||||
default:
|
||||
t.Fatalf("unexpected personal model override context %q", overrideContext)
|
||||
return codersdk.ChatPersonalModelOverride{}
|
||||
}
|
||||
}
|
||||
assertOverride := func(
|
||||
resp codersdk.UserChatPersonalModelOverridesResponse,
|
||||
overrideContext codersdk.ChatPersonalModelOverrideContext,
|
||||
mode codersdk.ChatPersonalModelOverrideMode,
|
||||
modelConfigID string,
|
||||
isSet bool,
|
||||
isMalformed bool,
|
||||
) {
|
||||
t.Helper()
|
||||
override := personalOverride(resp, overrideContext)
|
||||
require.Equal(t, overrideContext, override.Context)
|
||||
require.Equal(t, mode, override.Mode)
|
||||
require.Equal(t, modelConfigID, override.ModelConfigID)
|
||||
require.Equal(t, isSet, override.IsSet)
|
||||
require.Equal(t, isMalformed, override.IsMalformed)
|
||||
}
|
||||
assertDeploymentDefault := func(
|
||||
resp codersdk.UserChatPersonalModelOverridesResponse,
|
||||
overrideContext codersdk.ChatModelOverrideContext,
|
||||
modelConfigID string,
|
||||
isMalformed bool,
|
||||
) {
|
||||
t.Helper()
|
||||
var override codersdk.ChatModelOverrideResponse
|
||||
switch overrideContext {
|
||||
case codersdk.ChatModelOverrideContextGeneral:
|
||||
override = resp.DeploymentDefaults.General
|
||||
case codersdk.ChatModelOverrideContextExplore:
|
||||
override = resp.DeploymentDefaults.Explore
|
||||
default:
|
||||
t.Fatalf("unexpected deployment model override context %q", overrideContext)
|
||||
}
|
||||
require.Equal(t, overrideContext, override.Context)
|
||||
require.Equal(t, modelConfigID, override.ModelConfigID)
|
||||
require.Equal(t, isMalformed, override.IsMalformed)
|
||||
}
|
||||
upsertRaw := func(
|
||||
overrideContext codersdk.ChatPersonalModelOverrideContext,
|
||||
value string,
|
||||
) {
|
||||
t.Helper()
|
||||
err := db.UpsertUserChatPersonalModelOverride(dbauthz.AsSystemRestricted(ctx), database.UpsertUserChatPersonalModelOverrideParams{
|
||||
UserID: member.ID,
|
||||
Key: chatd.ChatPersonalModelOverrideKey(overrideContext),
|
||||
Value: value,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
}
|
||||
getRawFor := func(userID uuid.UUID, overrideContext codersdk.ChatPersonalModelOverrideContext) string {
|
||||
t.Helper()
|
||||
raw, err := db.GetUserChatPersonalModelOverride(dbauthz.AsSystemRestricted(ctx), database.GetUserChatPersonalModelOverrideParams{
|
||||
UserID: userID,
|
||||
Key: chatd.ChatPersonalModelOverrideKey(overrideContext),
|
||||
})
|
||||
if stderrors.Is(err, sql.ErrNoRows) {
|
||||
return ""
|
||||
}
|
||||
require.NoError(t, err)
|
||||
return raw
|
||||
}
|
||||
getRaw := func(overrideContext codersdk.ChatPersonalModelOverrideContext) string {
|
||||
t.Helper()
|
||||
return getRawFor(member.ID, overrideContext)
|
||||
}
|
||||
|
||||
t.Run("GETDisabledReturnsMissingDefaults", func(t *testing.T) {
|
||||
resp, err := memberClient.GetUserChatPersonalModelOverrides(ctx)
|
||||
require.NoError(t, err)
|
||||
require.False(t, resp.Enabled)
|
||||
assertOverride(resp, codersdk.ChatPersonalModelOverrideContextRoot, codersdk.ChatPersonalModelOverrideModeChatDefault, "", false, false)
|
||||
assertOverride(resp, codersdk.ChatPersonalModelOverrideContextGeneral, codersdk.ChatPersonalModelOverrideModeDeploymentDefault, "", false, false)
|
||||
assertOverride(resp, codersdk.ChatPersonalModelOverrideContextExplore, codersdk.ChatPersonalModelOverrideModeDeploymentDefault, "", false, false)
|
||||
})
|
||||
|
||||
upsertRaw(codersdk.ChatPersonalModelOverrideContextRoot, string(codersdk.ChatPersonalModelOverrideModeChatDefault))
|
||||
upsertRaw(codersdk.ChatPersonalModelOverrideContextGeneral, string(codersdk.ChatPersonalModelOverrideModeDeploymentDefault))
|
||||
upsertRaw(codersdk.ChatPersonalModelOverrideContextExplore, "model:"+modelConfig.ID.String())
|
||||
|
||||
t.Run("GETDisabledReturnsSavedValues", func(t *testing.T) {
|
||||
resp, err := memberClient.GetUserChatPersonalModelOverrides(ctx)
|
||||
require.NoError(t, err)
|
||||
require.False(t, resp.Enabled)
|
||||
assertOverride(resp, codersdk.ChatPersonalModelOverrideContextRoot, codersdk.ChatPersonalModelOverrideModeChatDefault, "", true, false)
|
||||
assertOverride(resp, codersdk.ChatPersonalModelOverrideContextGeneral, codersdk.ChatPersonalModelOverrideModeDeploymentDefault, "", true, false)
|
||||
assertOverride(resp, codersdk.ChatPersonalModelOverrideContextExplore, codersdk.ChatPersonalModelOverrideModeModel, modelConfig.ID.String(), true, false)
|
||||
})
|
||||
|
||||
t.Run("GETIncludesDeploymentDefaults", func(t *testing.T) {
|
||||
resp, err := memberClient.GetUserChatPersonalModelOverrides(ctx)
|
||||
require.NoError(t, err)
|
||||
assertDeploymentDefault(resp, codersdk.ChatModelOverrideContextGeneral, modelConfig.ID.String(), false)
|
||||
assertDeploymentDefault(resp, codersdk.ChatModelOverrideContextExplore, defaultModelConfig.ID.String(), false)
|
||||
})
|
||||
|
||||
t.Run("PUTDisabledReturns403AndPreservesRows", func(t *testing.T) {
|
||||
err := memberClient.UpdateUserChatPersonalModelOverride(ctx, codersdk.ChatPersonalModelOverrideContextRoot, codersdk.UpdateUserChatPersonalModelOverrideRequest{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeModel,
|
||||
ModelConfigID: modelConfig.ID.String(),
|
||||
})
|
||||
requireSDKError(t, err, http.StatusForbidden)
|
||||
require.Equal(t, string(codersdk.ChatPersonalModelOverrideModeChatDefault), getRaw(codersdk.ChatPersonalModelOverrideContextRoot))
|
||||
})
|
||||
|
||||
err = adminClient.UpdateChatPersonalModelOverridesAdminSettings(ctx, codersdk.UpdateChatPersonalModelOverridesAdminSettingsRequest{
|
||||
AllowUsers: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
contexts := []codersdk.ChatPersonalModelOverrideContext{
|
||||
codersdk.ChatPersonalModelOverrideContextRoot,
|
||||
codersdk.ChatPersonalModelOverrideContextGeneral,
|
||||
codersdk.ChatPersonalModelOverrideContextExplore,
|
||||
}
|
||||
|
||||
t.Run("PUTRejectsUnknownMode", func(t *testing.T) {
|
||||
rawBefore := getRaw(codersdk.ChatPersonalModelOverrideContextGeneral)
|
||||
err := memberClient.UpdateUserChatPersonalModelOverride(ctx, codersdk.ChatPersonalModelOverrideContextGeneral, codersdk.UpdateUserChatPersonalModelOverrideRequest{
|
||||
Mode: codersdk.ChatPersonalModelOverrideMode("banana"),
|
||||
})
|
||||
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
|
||||
require.Contains(t, sdkErr.Message, "Invalid personal model override mode.")
|
||||
require.Equal(t, rawBefore, getRaw(codersdk.ChatPersonalModelOverrideContextGeneral))
|
||||
})
|
||||
|
||||
t.Run("PUTChatDefaultRoundTrips", func(t *testing.T) {
|
||||
for _, overrideContext := range contexts {
|
||||
err := memberClient.UpdateUserChatPersonalModelOverride(ctx, overrideContext, codersdk.UpdateUserChatPersonalModelOverrideRequest{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeChatDefault,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
resp, err := memberClient.GetUserChatPersonalModelOverrides(ctx)
|
||||
require.NoError(t, err)
|
||||
require.True(t, resp.Enabled)
|
||||
for _, overrideContext := range contexts {
|
||||
assertOverride(resp, overrideContext, codersdk.ChatPersonalModelOverrideModeChatDefault, "", true, false)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("PUTChatDefaultRejectsNonEmptyModelConfigID", func(t *testing.T) {
|
||||
rawBefore := getRaw(codersdk.ChatPersonalModelOverrideContextRoot)
|
||||
err := memberClient.UpdateUserChatPersonalModelOverride(ctx, codersdk.ChatPersonalModelOverrideContextRoot, codersdk.UpdateUserChatPersonalModelOverrideRequest{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeChatDefault,
|
||||
ModelConfigID: modelConfig.ID.String(),
|
||||
})
|
||||
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
|
||||
require.Contains(t, sdkErr.Message, "model_config_id must be empty")
|
||||
require.Equal(t, rawBefore, getRaw(codersdk.ChatPersonalModelOverrideContextRoot))
|
||||
})
|
||||
|
||||
t.Run("PUTDeploymentDefaultRoundTripsForAgentContexts", func(t *testing.T) {
|
||||
for _, overrideContext := range []codersdk.ChatPersonalModelOverrideContext{
|
||||
codersdk.ChatPersonalModelOverrideContextGeneral,
|
||||
codersdk.ChatPersonalModelOverrideContextExplore,
|
||||
} {
|
||||
err := memberClient.UpdateUserChatPersonalModelOverride(ctx, overrideContext, codersdk.UpdateUserChatPersonalModelOverrideRequest{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
resp, err := memberClient.GetUserChatPersonalModelOverrides(ctx)
|
||||
require.NoError(t, err)
|
||||
assertOverride(resp, codersdk.ChatPersonalModelOverrideContextGeneral, codersdk.ChatPersonalModelOverrideModeDeploymentDefault, "", true, false)
|
||||
assertOverride(resp, codersdk.ChatPersonalModelOverrideContextExplore, codersdk.ChatPersonalModelOverrideModeDeploymentDefault, "", true, false)
|
||||
})
|
||||
|
||||
t.Run("PUTDeploymentDefaultRejectsNonEmptyModelConfigID", func(t *testing.T) {
|
||||
rawBefore := getRaw(codersdk.ChatPersonalModelOverrideContextGeneral)
|
||||
err := memberClient.UpdateUserChatPersonalModelOverride(ctx, codersdk.ChatPersonalModelOverrideContextGeneral, codersdk.UpdateUserChatPersonalModelOverrideRequest{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
ModelConfigID: modelConfig.ID.String(),
|
||||
})
|
||||
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
|
||||
require.Contains(t, sdkErr.Message, "model_config_id must be empty")
|
||||
require.Equal(t, rawBefore, getRaw(codersdk.ChatPersonalModelOverrideContextGeneral))
|
||||
})
|
||||
|
||||
t.Run("PUTDeploymentDefaultRejectsRoot", func(t *testing.T) {
|
||||
err := memberClient.UpdateUserChatPersonalModelOverride(ctx, codersdk.ChatPersonalModelOverrideContextRoot, codersdk.UpdateUserChatPersonalModelOverrideRequest{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
})
|
||||
requireSDKError(t, err, http.StatusBadRequest)
|
||||
})
|
||||
|
||||
t.Run("PUTModelRoundTrips", func(t *testing.T) {
|
||||
for _, overrideContext := range contexts {
|
||||
err := memberClient.UpdateUserChatPersonalModelOverride(ctx, overrideContext, codersdk.UpdateUserChatPersonalModelOverrideRequest{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeModel,
|
||||
ModelConfigID: modelConfig.ID.String(),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
resp, err := memberClient.GetUserChatPersonalModelOverrides(ctx)
|
||||
require.NoError(t, err)
|
||||
for _, overrideContext := range contexts {
|
||||
assertOverride(resp, overrideContext, codersdk.ChatPersonalModelOverrideModeModel, modelConfig.ID.String(), true, false)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("PUTModelRejectsInvalidModels", func(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
client *codersdk.ExperimentalClient
|
||||
userID uuid.UUID
|
||||
modelConfigID string
|
||||
wantMessageSubstring string
|
||||
}{
|
||||
{
|
||||
name: "Nil",
|
||||
client: memberClient,
|
||||
userID: member.ID,
|
||||
modelConfigID: uuid.Nil.String(),
|
||||
wantMessageSubstring: "Invalid model_config_id",
|
||||
},
|
||||
{
|
||||
name: "Empty",
|
||||
client: memberClient,
|
||||
userID: member.ID,
|
||||
modelConfigID: "",
|
||||
wantMessageSubstring: "model_config_id is required",
|
||||
},
|
||||
{
|
||||
name: "Malformed",
|
||||
client: memberClient,
|
||||
userID: member.ID,
|
||||
modelConfigID: "not-a-uuid",
|
||||
wantMessageSubstring: "Invalid model_config_id",
|
||||
},
|
||||
{
|
||||
name: "Unknown",
|
||||
client: memberClient,
|
||||
userID: member.ID,
|
||||
modelConfigID: uuid.NewString(),
|
||||
wantMessageSubstring: "Invalid model_config_id: model config " +
|
||||
"not found or disabled.",
|
||||
},
|
||||
{
|
||||
name: "Disabled",
|
||||
client: memberClient,
|
||||
userID: member.ID,
|
||||
modelConfigID: disabledModelConfig.ID.String(),
|
||||
wantMessageSubstring: "Invalid model_config_id: model config " +
|
||||
"not found or disabled.",
|
||||
},
|
||||
{
|
||||
name: "ProviderDisabled",
|
||||
client: memberClient,
|
||||
userID: member.ID,
|
||||
modelConfigID: disabledProviderModelConfig.ID.String(),
|
||||
wantMessageSubstring: "provider is not enabled",
|
||||
},
|
||||
{
|
||||
name: "CredentialUnavailable",
|
||||
client: noKeyClient,
|
||||
userID: noKeyUser.ID,
|
||||
modelConfigID: modelConfig.ID.String(),
|
||||
wantMessageSubstring: "Invalid model_config_id: provider " +
|
||||
"credentials unavailable for this model.",
|
||||
},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
rawBefore := getRawFor(tc.userID, codersdk.ChatPersonalModelOverrideContextGeneral)
|
||||
err := tc.client.UpdateUserChatPersonalModelOverride(ctx, codersdk.ChatPersonalModelOverrideContextGeneral, codersdk.UpdateUserChatPersonalModelOverrideRequest{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeModel,
|
||||
ModelConfigID: tc.modelConfigID,
|
||||
})
|
||||
sdkErr := requireSDKError(t, err, http.StatusBadRequest)
|
||||
require.Contains(t, sdkErr.Message, tc.wantMessageSubstring)
|
||||
rawAfter := getRawFor(tc.userID, codersdk.ChatPersonalModelOverrideContextGeneral)
|
||||
require.Equal(t, rawBefore, rawAfter)
|
||||
})
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("GETMalformedStoredValueFallsBackToContextDefault", func(t *testing.T) {
|
||||
upsertRaw(codersdk.ChatPersonalModelOverrideContextRoot, "model:not-a-uuid")
|
||||
|
||||
resp, err := memberClient.GetUserChatPersonalModelOverrides(ctx)
|
||||
require.NoError(t, err)
|
||||
assertOverride(resp, codersdk.ChatPersonalModelOverrideContextRoot, codersdk.ChatPersonalModelOverrideModeChatDefault, "", true, true)
|
||||
})
|
||||
|
||||
t.Run("GETRootDeploymentDefaultIsMalformed", func(t *testing.T) {
|
||||
upsertRaw(
|
||||
codersdk.ChatPersonalModelOverrideContextRoot,
|
||||
string(codersdk.ChatPersonalModelOverrideModeDeploymentDefault),
|
||||
)
|
||||
|
||||
resp, err := memberClient.GetUserChatPersonalModelOverrides(ctx)
|
||||
require.NoError(t, err)
|
||||
assertOverride(resp, codersdk.ChatPersonalModelOverrideContextRoot, codersdk.ChatPersonalModelOverrideModeChatDefault, "", true, true)
|
||||
})
|
||||
}
|
||||
|
||||
//nolint:tparallel,paralleltest // Subtests share coderdtest instances.
|
||||
func TestCreateChatPersonalModelOverrideRoot(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
adminClient, db := newChatClientWithDatabase(t)
|
||||
firstUser := coderdtest.CreateFirstUser(t, adminClient.Client)
|
||||
defaultModel := createChatModelConfig(t, adminClient)
|
||||
_ = enableUserChatProviderKey(t, adminClient, adminClient, defaultModel.Provider)
|
||||
overrideModel := createAdditionalChatModelConfig(
|
||||
t,
|
||||
adminClient,
|
||||
defaultModel.Provider,
|
||||
"gpt-4o-root-personal-"+uuid.NewString(),
|
||||
)
|
||||
disabledModel := createDisabledChatModelConfig(
|
||||
t,
|
||||
adminClient,
|
||||
defaultModel.Provider,
|
||||
"gpt-4o-root-personal-disabled-"+uuid.NewString(),
|
||||
)
|
||||
memberClientRaw, member := coderdtest.CreateAnotherUser(
|
||||
t,
|
||||
adminClient.Client,
|
||||
firstUser.OrganizationID,
|
||||
rbac.ScopedRoleAgentsAccess(firstUser.OrganizationID),
|
||||
)
|
||||
memberClient := codersdk.NewExperimentalClient(memberClientRaw)
|
||||
|
||||
createChat := func(
|
||||
client *codersdk.ExperimentalClient,
|
||||
text string,
|
||||
modelConfigID *uuid.UUID,
|
||||
) codersdk.Chat {
|
||||
t.Helper()
|
||||
chat, err := client.CreateChat(ctx, codersdk.CreateChatRequest{
|
||||
OrganizationID: firstUser.OrganizationID,
|
||||
Content: []codersdk.ChatInputPart{{
|
||||
Type: codersdk.ChatInputPartTypeText,
|
||||
Text: text,
|
||||
}},
|
||||
ModelConfigID: modelConfigID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
storedChat, err := db.GetChatByID(dbauthz.AsSystemRestricted(ctx), chat.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, chat.LastModelConfigID, storedChat.LastModelConfigID)
|
||||
return chat
|
||||
}
|
||||
upsertRootRaw := func(userID uuid.UUID, value string) {
|
||||
t.Helper()
|
||||
err := db.UpsertUserChatPersonalModelOverride(dbauthz.AsSystemRestricted(ctx), database.UpsertUserChatPersonalModelOverrideParams{
|
||||
UserID: userID,
|
||||
Key: chatd.ChatPersonalModelOverrideKey(codersdk.ChatPersonalModelOverrideContextRoot),
|
||||
Value: value,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
err := adminClient.UpdateChatPersonalModelOverridesAdminSettings(ctx, codersdk.UpdateChatPersonalModelOverridesAdminSettingsRequest{
|
||||
AllowUsers: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
err = adminClient.UpdateUserChatPersonalModelOverride(ctx, codersdk.ChatPersonalModelOverrideContextRoot, codersdk.UpdateUserChatPersonalModelOverrideRequest{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeModel,
|
||||
ModelConfigID: overrideModel.ID.String(),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
t.Run("ExplicitModelConfigWins", func(t *testing.T) {
|
||||
chat := createChat(adminClient, "explicit model config wins", ptr.Ref(defaultModel.ID))
|
||||
require.Equal(t, defaultModel.ID, chat.LastModelConfigID)
|
||||
})
|
||||
|
||||
t.Run("FlagOffIgnoresSavedRootModel", func(t *testing.T) {
|
||||
err := adminClient.UpdateChatPersonalModelOverridesAdminSettings(ctx, codersdk.UpdateChatPersonalModelOverridesAdminSettingsRequest{
|
||||
AllowUsers: false,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
chat := createChat(adminClient, "flag off uses default", nil)
|
||||
require.Equal(t, defaultModel.ID, chat.LastModelConfigID)
|
||||
})
|
||||
|
||||
t.Run("ChatDefaultUsesDefaultModel", func(t *testing.T) {
|
||||
err := adminClient.UpdateChatPersonalModelOverridesAdminSettings(ctx, codersdk.UpdateChatPersonalModelOverridesAdminSettingsRequest{
|
||||
AllowUsers: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
err = adminClient.UpdateUserChatPersonalModelOverride(ctx, codersdk.ChatPersonalModelOverrideContextRoot, codersdk.UpdateUserChatPersonalModelOverrideRequest{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeChatDefault,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
chat := createChat(adminClient, "chat default uses default", nil)
|
||||
require.Equal(t, defaultModel.ID, chat.LastModelConfigID)
|
||||
})
|
||||
|
||||
t.Run("MalformedRootFallsBackToDefault", func(t *testing.T) {
|
||||
upsertRootRaw(firstUser.UserID, "garbage")
|
||||
chat := createChat(adminClient, "malformed root falls back", nil)
|
||||
require.Equal(t, defaultModel.ID, chat.LastModelConfigID)
|
||||
})
|
||||
|
||||
t.Run("RootModelOverrideUsesSavedModel", func(t *testing.T) {
|
||||
err := adminClient.UpdateUserChatPersonalModelOverride(ctx, codersdk.ChatPersonalModelOverrideContextRoot, codersdk.UpdateUserChatPersonalModelOverrideRequest{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeModel,
|
||||
ModelConfigID: overrideModel.ID.String(),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
chat := createChat(adminClient, "root model override uses saved model", nil)
|
||||
require.Equal(t, overrideModel.ID, chat.LastModelConfigID)
|
||||
})
|
||||
|
||||
t.Run("UnavailableRootModelFallsBackToDefault", func(t *testing.T) {
|
||||
upsertRootRaw(firstUser.UserID, "model:"+disabledModel.ID.String())
|
||||
chat := createChat(adminClient, "disabled root model falls back", nil)
|
||||
require.Equal(t, defaultModel.ID, chat.LastModelConfigID)
|
||||
|
||||
upsertRootRaw(member.ID, "model:"+overrideModel.ID.String())
|
||||
chat = createChat(memberClient, "missing user key falls back", nil)
|
||||
require.Equal(t, defaultModel.ID, chat.LastModelConfigID)
|
||||
})
|
||||
}
|
||||
|
||||
func TestChatDesktopEnabled(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -0,0 +1,75 @@
|
||||
package chatd
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/google/uuid"
|
||||
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
|
||||
// ChatPersonalModelOverrideKeyPrefix is the user config key prefix for
|
||||
// chat personal model overrides. Values under this prefix should be parsed
|
||||
// with ParseChatPersonalModelOverride so malformed values use one fallback.
|
||||
const ChatPersonalModelOverrideKeyPrefix = "chat_personal_model_override:"
|
||||
|
||||
// ChatPersonalModelOverrideKey returns the user config key for a chat
|
||||
// personal model override context. Values stored at the returned key should
|
||||
// use ParseChatPersonalModelOverride so malformed values fall back safely.
|
||||
func ChatPersonalModelOverrideKey(
|
||||
overrideContext codersdk.ChatPersonalModelOverrideContext,
|
||||
) string {
|
||||
return ChatPersonalModelOverrideKeyPrefix + string(overrideContext)
|
||||
}
|
||||
|
||||
// ParsedChatPersonalModelOverride is a parsed personal model override value.
|
||||
// When Malformed is true, Mode is the provided default and ModelConfigID is
|
||||
// uuid.Nil.
|
||||
type ParsedChatPersonalModelOverride struct {
|
||||
Mode codersdk.ChatPersonalModelOverrideMode
|
||||
ModelConfigID uuid.UUID
|
||||
Malformed bool
|
||||
}
|
||||
|
||||
// ParseChatPersonalModelOverride parses a stored personal model override.
|
||||
// Empty values return defaultMode without marking the value malformed.
|
||||
// Malformed values return defaultMode, uuid.Nil, and Malformed true.
|
||||
func ParseChatPersonalModelOverride(
|
||||
raw string,
|
||||
defaultMode codersdk.ChatPersonalModelOverrideMode,
|
||||
) ParsedChatPersonalModelOverride {
|
||||
trimmed := strings.TrimSpace(raw)
|
||||
if trimmed == "" {
|
||||
return ParsedChatPersonalModelOverride{Mode: defaultMode}
|
||||
}
|
||||
|
||||
switch trimmed {
|
||||
case string(codersdk.ChatPersonalModelOverrideModeChatDefault):
|
||||
return ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeChatDefault,
|
||||
}
|
||||
case string(codersdk.ChatPersonalModelOverrideModeDeploymentDefault):
|
||||
return ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
}
|
||||
}
|
||||
|
||||
mode, rawModelConfigID, ok := strings.Cut(trimmed, ":")
|
||||
if !ok || mode != string(codersdk.ChatPersonalModelOverrideModeModel) {
|
||||
return ParsedChatPersonalModelOverride{
|
||||
Mode: defaultMode,
|
||||
Malformed: true,
|
||||
}
|
||||
}
|
||||
modelConfigID, err := uuid.Parse(rawModelConfigID)
|
||||
if err != nil {
|
||||
return ParsedChatPersonalModelOverride{
|
||||
Mode: defaultMode,
|
||||
Malformed: true,
|
||||
}
|
||||
}
|
||||
return ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeModel,
|
||||
ModelConfigID: modelConfigID,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,103 @@
|
||||
package chatd_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/x/chatd"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
|
||||
func TestChatPersonalModelOverrideKey(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
require.Equal(
|
||||
t,
|
||||
"chat_personal_model_override:root",
|
||||
chatd.ChatPersonalModelOverrideKey(codersdk.ChatPersonalModelOverrideContextRoot),
|
||||
)
|
||||
}
|
||||
|
||||
func TestParseChatPersonalModelOverride(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
modelConfigID := uuid.MustParse("11111111-1111-1111-1111-111111111111")
|
||||
tests := []struct {
|
||||
name string
|
||||
raw string
|
||||
defaultMode codersdk.ChatPersonalModelOverrideMode
|
||||
want chatd.ParsedChatPersonalModelOverride
|
||||
}{
|
||||
{
|
||||
name: "EmptyUsesDefault",
|
||||
raw: "",
|
||||
defaultMode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
want: chatd.ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "ChatDefault",
|
||||
raw: string(codersdk.ChatPersonalModelOverrideModeChatDefault),
|
||||
defaultMode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
want: chatd.ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeChatDefault,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "DeploymentDefault",
|
||||
raw: string(codersdk.ChatPersonalModelOverrideModeDeploymentDefault),
|
||||
defaultMode: codersdk.ChatPersonalModelOverrideModeChatDefault,
|
||||
want: chatd.ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "Model",
|
||||
raw: "model:" + modelConfigID.String(),
|
||||
defaultMode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
want: chatd.ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeModel,
|
||||
ModelConfigID: modelConfigID,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "InvalidModelUUID",
|
||||
raw: "model:not-a-uuid",
|
||||
defaultMode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
want: chatd.ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
Malformed: true,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "UnknownValue",
|
||||
raw: "unknown",
|
||||
defaultMode: codersdk.ChatPersonalModelOverrideModeChatDefault,
|
||||
want: chatd.ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeChatDefault,
|
||||
Malformed: true,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "OuterWhitespace",
|
||||
raw: " \tmodel:" + modelConfigID.String() + "\n",
|
||||
defaultMode: codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
want: chatd.ParsedChatPersonalModelOverride{
|
||||
Mode: codersdk.ChatPersonalModelOverrideModeModel,
|
||||
ModelConfigID: modelConfigID,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := chatd.ParseChatPersonalModelOverride(tt.raw, tt.defaultMode)
|
||||
require.Equal(t, tt.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
+175
-5
@@ -140,6 +140,22 @@ func readSubagentModelOverride(
|
||||
}
|
||||
}
|
||||
|
||||
func personalModelOverrideContextForSubagent(
|
||||
overrideContext codersdk.ChatModelOverrideContext,
|
||||
) (codersdk.ChatPersonalModelOverrideContext, error) {
|
||||
switch overrideContext {
|
||||
case codersdk.ChatModelOverrideContextGeneral:
|
||||
return codersdk.ChatPersonalModelOverrideContextGeneral, nil
|
||||
case codersdk.ChatModelOverrideContextExplore:
|
||||
return codersdk.ChatPersonalModelOverrideContextExplore, nil
|
||||
default:
|
||||
return "", xerrors.Errorf(
|
||||
"unknown subagent model override context %q",
|
||||
overrideContext,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
func validateModelConfigAndResolveProvider(
|
||||
modelConfig database.ChatModelConfig,
|
||||
) (database.ChatModelConfig, string, error) {
|
||||
@@ -173,6 +189,15 @@ func enabledProviderContainsName(
|
||||
return false
|
||||
}
|
||||
|
||||
func userCanUseProviderKeys(
|
||||
providerKeys chatprovider.ProviderAPIKeys,
|
||||
providerName string,
|
||||
) bool {
|
||||
return providerKeys.APIKey(providerName) != "" ||
|
||||
(chatprovider.ProviderAllowsAmbientCredentials(providerName) &&
|
||||
providerKeys.HasProvider(providerName))
|
||||
}
|
||||
|
||||
type modelOverrideFailureMode int
|
||||
|
||||
const (
|
||||
@@ -274,9 +299,7 @@ func (p *Server) resolveConfiguredModelOverride(
|
||||
err,
|
||||
)
|
||||
}
|
||||
if providerKeys.APIKey(providerName) == "" &&
|
||||
!(chatprovider.ProviderAllowsAmbientCredentials(providerName) &&
|
||||
providerKeys.HasProvider(providerName)) {
|
||||
if !userCanUseProviderKeys(providerKeys, providerName) {
|
||||
if failureMode == modelOverrideFailureModeHard {
|
||||
return database.ChatModelConfig{}, true, xerrors.Errorf(
|
||||
"%s model override credentials are unavailable for provider %q",
|
||||
@@ -296,13 +319,160 @@ func (p *Server) resolveConfiguredModelOverride(
|
||||
return modelConfig, true, nil
|
||||
}
|
||||
|
||||
func (p *Server) resolvePersonalSubagentModelConfigID(
|
||||
ctx context.Context,
|
||||
ownerID uuid.UUID,
|
||||
overrideContext codersdk.ChatModelOverrideContext,
|
||||
) (uuid.UUID, bool, error) {
|
||||
personalContext, err := personalModelOverrideContextForSubagent(overrideContext)
|
||||
if err != nil {
|
||||
return uuid.Nil, false, err
|
||||
}
|
||||
raw, err := p.db.GetUserChatPersonalModelOverride(
|
||||
ctx,
|
||||
database.GetUserChatPersonalModelOverrideParams{
|
||||
UserID: ownerID,
|
||||
Key: ChatPersonalModelOverrideKey(personalContext),
|
||||
},
|
||||
)
|
||||
if err != nil {
|
||||
if !xerrors.Is(err, sql.ErrNoRows) {
|
||||
return uuid.Nil, false, xerrors.Errorf(
|
||||
"get %s personal model override: %w",
|
||||
subagentModelOverrideLogLabel(overrideContext),
|
||||
err,
|
||||
)
|
||||
}
|
||||
raw = ""
|
||||
}
|
||||
|
||||
parsed := ParseChatPersonalModelOverride(
|
||||
raw,
|
||||
codersdk.ChatPersonalModelOverrideModeDeploymentDefault,
|
||||
)
|
||||
if parsed.Malformed {
|
||||
p.logger.Debug(ctx,
|
||||
"personal model override is malformed, using deployment default",
|
||||
slog.F("override_context", overrideContext),
|
||||
slog.F("owner_id", ownerID),
|
||||
slog.F("raw_model_config_id", strings.TrimSpace(raw)),
|
||||
)
|
||||
}
|
||||
switch parsed.Mode {
|
||||
case codersdk.ChatPersonalModelOverrideModeChatDefault:
|
||||
return uuid.Nil, true, nil
|
||||
case codersdk.ChatPersonalModelOverrideModeDeploymentDefault:
|
||||
case codersdk.ChatPersonalModelOverrideModeModel:
|
||||
modelConfig, ok, err := p.resolvePersonalModelOverride(
|
||||
ctx,
|
||||
overrideContext,
|
||||
ownerID,
|
||||
parsed.ModelConfigID,
|
||||
)
|
||||
if err != nil {
|
||||
return uuid.Nil, false, err
|
||||
}
|
||||
if ok {
|
||||
return modelConfig.ID, true, nil
|
||||
}
|
||||
default:
|
||||
p.logger.Warn(ctx,
|
||||
"unsupported personal model override mode, using deployment default",
|
||||
slog.F("override_context", overrideContext),
|
||||
slog.F("owner_id", ownerID),
|
||||
slog.F("mode", parsed.Mode),
|
||||
)
|
||||
}
|
||||
|
||||
return uuid.Nil, false, nil
|
||||
}
|
||||
|
||||
func (p *Server) resolvePersonalModelOverride(
|
||||
ctx context.Context,
|
||||
overrideContext codersdk.ChatModelOverrideContext,
|
||||
ownerID uuid.UUID,
|
||||
modelConfigID uuid.UUID,
|
||||
) (database.ChatModelConfig, bool, error) {
|
||||
modelConfig, providerName, err := p.resolveModelConfigAndNormalizedProvider(
|
||||
ctx,
|
||||
modelConfigID,
|
||||
)
|
||||
if err != nil {
|
||||
switch {
|
||||
case xerrors.Is(err, sql.ErrNoRows):
|
||||
p.logger.Debug(ctx,
|
||||
"personal model override is unavailable, using deployment default",
|
||||
slog.F("override_context", overrideContext),
|
||||
slog.F("owner_id", ownerID),
|
||||
slog.F("model_config_id", modelConfigID),
|
||||
)
|
||||
case errors.Is(err, errInvalidModelOverrideMetadata):
|
||||
p.logger.Debug(ctx,
|
||||
"personal model override metadata is invalid, using deployment default",
|
||||
slog.F("override_context", overrideContext),
|
||||
slog.F("owner_id", ownerID),
|
||||
slog.F("model_config_id", modelConfigID),
|
||||
slog.Error(err),
|
||||
)
|
||||
default:
|
||||
p.logger.Warn(ctx,
|
||||
"failed to resolve personal model override, using deployment default",
|
||||
slog.F("override_context", overrideContext),
|
||||
slog.F("owner_id", ownerID),
|
||||
slog.F("model_config_id", modelConfigID),
|
||||
slog.Error(err),
|
||||
)
|
||||
}
|
||||
return database.ChatModelConfig{}, false, nil
|
||||
}
|
||||
providerKeys, err := p.resolveUserProviderAPIKeys(ctx, ownerID)
|
||||
if err != nil {
|
||||
return database.ChatModelConfig{}, false, xerrors.Errorf(
|
||||
"resolve provider API keys: %w",
|
||||
err,
|
||||
)
|
||||
}
|
||||
if !userCanUseProviderKeys(providerKeys, providerName) {
|
||||
p.logger.Debug(ctx,
|
||||
"personal model override credentials are unavailable, using deployment default",
|
||||
slog.F("override_context", overrideContext),
|
||||
slog.F("owner_id", ownerID),
|
||||
slog.F("model_config_id", modelConfigID),
|
||||
slog.F("provider", providerName),
|
||||
)
|
||||
return database.ChatModelConfig{}, false, nil
|
||||
}
|
||||
return modelConfig, true, nil
|
||||
}
|
||||
|
||||
func (p *Server) resolveSubagentModelConfigID(
|
||||
ctx context.Context,
|
||||
ownerID uuid.UUID,
|
||||
overrideContext codersdk.ChatModelOverrideContext,
|
||||
) (uuid.UUID, error) {
|
||||
//nolint:gocritic // Chatd needs its scoped deployment-config read access here.
|
||||
//nolint:gocritic // Chatd needs its scoped config and user-data access here.
|
||||
chatdCtx := dbauthz.AsChatd(ctx)
|
||||
personalOverridesEnabled, err := p.db.GetChatPersonalModelOverridesEnabled(chatdCtx)
|
||||
if err != nil {
|
||||
return uuid.Nil, xerrors.Errorf(
|
||||
"get chat personal model overrides enabled: %w",
|
||||
err,
|
||||
)
|
||||
}
|
||||
if personalOverridesEnabled {
|
||||
modelConfigID, resolved, err := p.resolvePersonalSubagentModelConfigID(
|
||||
chatdCtx,
|
||||
ownerID,
|
||||
overrideContext,
|
||||
)
|
||||
if err != nil {
|
||||
return uuid.Nil, err
|
||||
}
|
||||
if resolved {
|
||||
return modelConfigID, nil
|
||||
}
|
||||
}
|
||||
|
||||
raw, err := readSubagentModelOverride(chatdCtx, p.db, overrideContext)
|
||||
if err != nil {
|
||||
return uuid.Nil, xerrors.Errorf(
|
||||
@@ -312,7 +482,7 @@ func (p *Server) resolveSubagentModelConfigID(
|
||||
)
|
||||
}
|
||||
modelConfig, ok, err := p.resolveConfiguredModelOverride(
|
||||
ctx,
|
||||
chatdCtx,
|
||||
string(overrideContext),
|
||||
raw,
|
||||
ownerID,
|
||||
|
||||
@@ -409,6 +409,46 @@ func chatdTestContext(t *testing.T) context.Context {
|
||||
return dbauthz.AsChatd(testutil.Context(t, testutil.WaitLong))
|
||||
}
|
||||
|
||||
func systemRestrictedTestContext(t *testing.T) context.Context {
|
||||
t.Helper()
|
||||
return dbauthz.AsSystemRestricted(testutil.Context(t, testutil.WaitLong))
|
||||
}
|
||||
|
||||
func enableInternalChatPersonalModelOverrides(
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
) {
|
||||
t.Helper()
|
||||
require.NoError(
|
||||
t,
|
||||
db.UpsertChatPersonalModelOverridesEnabled(
|
||||
systemRestrictedTestContext(t),
|
||||
true,
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
func upsertInternalUserChatPersonalModelOverride(
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
userID uuid.UUID,
|
||||
overrideContext codersdk.ChatPersonalModelOverrideContext,
|
||||
raw string,
|
||||
) {
|
||||
t.Helper()
|
||||
require.NoError(
|
||||
t,
|
||||
db.UpsertUserChatPersonalModelOverride(
|
||||
systemRestrictedTestContext(t),
|
||||
database.UpsertUserChatPersonalModelOverrideParams{
|
||||
UserID: userID,
|
||||
Key: ChatPersonalModelOverrideKey(overrideContext),
|
||||
Value: raw,
|
||||
},
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
func TestCreateChildSubagentChatInheritsWorkspaceBinding(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -668,6 +708,204 @@ func TestSpawnAgent_GeneralUsesConfiguredModelOverride(t *testing.T) {
|
||||
require.False(t, childChat.PlanMode.Valid)
|
||||
}
|
||||
|
||||
func TestSpawnAgent_GeneralHonorsPersonalModelOverrides(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
enablePersonalOverride bool
|
||||
personalRaw func(database.ChatModelConfig) string
|
||||
personalModel func(context.Context, *testing.T, database.Store, uuid.UUID) database.ChatModelConfig
|
||||
wantModelID func(
|
||||
database.ChatModelConfig,
|
||||
database.ChatModelConfig,
|
||||
database.ChatModelConfig,
|
||||
) uuid.UUID
|
||||
}{
|
||||
{
|
||||
name: "UnsetUsesDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
wantModelID: func(_, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "DeploymentDefaultUsesDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalRaw: func(database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeDeploymentDefault)
|
||||
},
|
||||
wantModelID: func(_, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "ChatDefaultBypassesDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalRaw: func(database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeChatDefault)
|
||||
},
|
||||
wantModelID: func(parentModel, _, _ database.ChatModelConfig) uuid.UUID {
|
||||
return parentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "ModelUsesPersonalOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalRaw: func(personalModel database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeModel) + ":" +
|
||||
personalModel.ID.String()
|
||||
},
|
||||
wantModelID: func(_, _, personalModel database.ChatModelConfig) uuid.UUID {
|
||||
return personalModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "AdminFlagOffIgnoresPersonalOverride",
|
||||
personalRaw: func(database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeChatDefault)
|
||||
},
|
||||
wantModelID: func(_, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "DisabledPersonalModelFallsBackToDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalModel: func(
|
||||
ctx context.Context,
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
userID uuid.UUID,
|
||||
) database.ChatModelConfig {
|
||||
return insertInternalChatModelConfig(
|
||||
t,
|
||||
db,
|
||||
"general-personal-disabled-"+uuid.NewString(),
|
||||
false,
|
||||
)
|
||||
},
|
||||
personalRaw: func(personalModel database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeModel) + ":" +
|
||||
personalModel.ID.String()
|
||||
},
|
||||
wantModelID: func(_, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "MissingCredentialsFallsBackToDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalModel: func(
|
||||
ctx context.Context,
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
userID uuid.UUID,
|
||||
) database.ChatModelConfig {
|
||||
insertInternalChatProvider(
|
||||
t,
|
||||
db,
|
||||
userID,
|
||||
"openai-compat",
|
||||
"",
|
||||
false,
|
||||
true,
|
||||
false,
|
||||
)
|
||||
return insertInternalChatModelConfigForProvider(
|
||||
t,
|
||||
db,
|
||||
"openai-compat",
|
||||
"gpt-4o-mini",
|
||||
true,
|
||||
)
|
||||
},
|
||||
personalRaw: func(personalModel database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeModel) + ":" +
|
||||
personalModel.ID.String()
|
||||
},
|
||||
wantModelID: func(_, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "MalformedValueUsesDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalRaw: func(database.ChatModelConfig) string {
|
||||
return "model:not-a-uuid"
|
||||
},
|
||||
wantModelID: func(_, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{})
|
||||
|
||||
ctx := chatdTestContext(t)
|
||||
user, org, parentModel := seedInternalChatDeps(t, db)
|
||||
deploymentModel := insertInternalChatModelConfig(
|
||||
t,
|
||||
db,
|
||||
"general-deployment-"+uuid.NewString(),
|
||||
true,
|
||||
)
|
||||
require.NoError(t, db.UpsertChatGeneralModelOverride(ctx, deploymentModel.ID.String()))
|
||||
personalModel := insertInternalChatModelConfig(
|
||||
t,
|
||||
db,
|
||||
"general-personal-"+uuid.NewString(),
|
||||
true,
|
||||
)
|
||||
if tt.personalModel != nil {
|
||||
personalModel = tt.personalModel(ctx, t, db, user.ID)
|
||||
}
|
||||
if tt.enablePersonalOverride {
|
||||
enableInternalChatPersonalModelOverrides(t, db)
|
||||
}
|
||||
if tt.personalRaw != nil {
|
||||
upsertInternalUserChatPersonalModelOverride(
|
||||
t,
|
||||
db,
|
||||
user.ID,
|
||||
codersdk.ChatPersonalModelOverrideContextGeneral,
|
||||
tt.personalRaw(personalModel),
|
||||
)
|
||||
}
|
||||
parentChat := createInternalParentChat(
|
||||
ctx,
|
||||
t,
|
||||
server,
|
||||
db,
|
||||
org.ID,
|
||||
user.ID,
|
||||
parentModel.ID,
|
||||
"parent-general-personal-override",
|
||||
)
|
||||
|
||||
resp := runSpawnAgentTool(ctx, t, server, parentChat, spawnAgentArgs{
|
||||
Type: subagentTypeGeneral,
|
||||
Prompt: "delegate general work",
|
||||
})
|
||||
childID := requireSpawnAgentChildChatID(t, resp)
|
||||
|
||||
childChat, err := db.GetChatByID(ctx, childID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(
|
||||
t,
|
||||
tt.wantModelID(parentModel, deploymentModel, personalModel),
|
||||
childChat.LastModelConfigID,
|
||||
)
|
||||
require.False(t, childChat.PlanMode.Valid)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSpawnAgent_GeneralOverrideLogsAndFallsBackWhenCredentialsUnavailable(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -947,6 +1185,218 @@ func TestSpawnAgent_ExploreFallsBackToCurrentTurnModel(t *testing.T) {
|
||||
require.Equal(t, parentModel.ID, parentChat.LastModelConfigID)
|
||||
}
|
||||
|
||||
func TestSpawnAgent_ExploreHonorsPersonalModelOverrides(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
enablePersonalOverride bool
|
||||
personalRaw func(database.ChatModelConfig) string
|
||||
personalModel func(context.Context, *testing.T, database.Store, uuid.UUID) database.ChatModelConfig
|
||||
wantModelID func(
|
||||
database.ChatModelConfig,
|
||||
database.ChatModelConfig,
|
||||
database.ChatModelConfig,
|
||||
database.ChatModelConfig,
|
||||
) uuid.UUID
|
||||
}{
|
||||
{
|
||||
name: "UnsetUsesDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
wantModelID: func(_, _, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "DeploymentDefaultUsesDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalRaw: func(database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeDeploymentDefault)
|
||||
},
|
||||
wantModelID: func(_, _, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "ChatDefaultBypassesDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalRaw: func(database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeChatDefault)
|
||||
},
|
||||
wantModelID: func(_, currentTurnModel, _, _ database.ChatModelConfig) uuid.UUID {
|
||||
return currentTurnModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "ModelUsesPersonalOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalRaw: func(personalModel database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeModel) + ":" +
|
||||
personalModel.ID.String()
|
||||
},
|
||||
wantModelID: func(_, _, _, personalModel database.ChatModelConfig) uuid.UUID {
|
||||
return personalModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "AdminFlagOffIgnoresPersonalOverride",
|
||||
personalRaw: func(database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeChatDefault)
|
||||
},
|
||||
wantModelID: func(_, _, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "DisabledPersonalModelFallsBackToDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalModel: func(
|
||||
ctx context.Context,
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
userID uuid.UUID,
|
||||
) database.ChatModelConfig {
|
||||
return insertInternalChatModelConfig(
|
||||
t,
|
||||
db,
|
||||
"explore-personal-disabled-"+uuid.NewString(),
|
||||
false,
|
||||
)
|
||||
},
|
||||
personalRaw: func(personalModel database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeModel) + ":" +
|
||||
personalModel.ID.String()
|
||||
},
|
||||
wantModelID: func(_, _, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "MissingCredentialsFallsBackToDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalModel: func(
|
||||
ctx context.Context,
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
userID uuid.UUID,
|
||||
) database.ChatModelConfig {
|
||||
insertInternalChatProvider(
|
||||
t,
|
||||
db,
|
||||
userID,
|
||||
"openai-compat",
|
||||
"",
|
||||
false,
|
||||
true,
|
||||
false,
|
||||
)
|
||||
return insertInternalChatModelConfigForProvider(
|
||||
t,
|
||||
db,
|
||||
"openai-compat",
|
||||
"gpt-4o-mini",
|
||||
true,
|
||||
)
|
||||
},
|
||||
personalRaw: func(personalModel database.ChatModelConfig) string {
|
||||
return string(codersdk.ChatPersonalModelOverrideModeModel) + ":" +
|
||||
personalModel.ID.String()
|
||||
},
|
||||
wantModelID: func(_, _, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "MalformedValueUsesDeploymentOverride",
|
||||
enablePersonalOverride: true,
|
||||
personalRaw: func(database.ChatModelConfig) string {
|
||||
return "not-a-mode"
|
||||
},
|
||||
wantModelID: func(_, _, deploymentModel, _ database.ChatModelConfig) uuid.UUID {
|
||||
return deploymentModel.ID
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{})
|
||||
|
||||
ctx := chatdTestContext(t)
|
||||
user, org, parentModel := seedInternalChatDeps(t, db)
|
||||
currentTurnModel := insertInternalChatModelConfig(
|
||||
t,
|
||||
db,
|
||||
"explore-current-turn-"+uuid.NewString(),
|
||||
true,
|
||||
)
|
||||
deploymentModel := insertInternalChatModelConfig(
|
||||
t,
|
||||
db,
|
||||
"explore-deployment-"+uuid.NewString(),
|
||||
true,
|
||||
)
|
||||
require.NoError(t, db.UpsertChatExploreModelOverride(ctx, deploymentModel.ID.String()))
|
||||
personalModel := insertInternalChatModelConfig(
|
||||
t,
|
||||
db,
|
||||
"explore-personal-"+uuid.NewString(),
|
||||
true,
|
||||
)
|
||||
if tt.personalModel != nil {
|
||||
personalModel = tt.personalModel(ctx, t, db, user.ID)
|
||||
}
|
||||
if tt.enablePersonalOverride {
|
||||
enableInternalChatPersonalModelOverrides(t, db)
|
||||
}
|
||||
if tt.personalRaw != nil {
|
||||
upsertInternalUserChatPersonalModelOverride(
|
||||
t,
|
||||
db,
|
||||
user.ID,
|
||||
codersdk.ChatPersonalModelOverrideContextExplore,
|
||||
tt.personalRaw(personalModel),
|
||||
)
|
||||
}
|
||||
parentChat := createInternalParentChat(
|
||||
ctx,
|
||||
t,
|
||||
server,
|
||||
db,
|
||||
org.ID,
|
||||
user.ID,
|
||||
parentModel.ID,
|
||||
"parent-explore-personal-override",
|
||||
)
|
||||
|
||||
resp := runSubagentTool(
|
||||
ctx,
|
||||
t,
|
||||
server,
|
||||
parentChat,
|
||||
currentTurnModel.ID,
|
||||
spawnAgentToolName,
|
||||
spawnAgentArgs{Type: subagentTypeExplore, Prompt: "inspect the codebase"},
|
||||
)
|
||||
childID := requireSpawnAgentChildChatID(t, resp)
|
||||
|
||||
childChat, err := db.GetChatByID(ctx, childID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(
|
||||
t,
|
||||
tt.wantModelID(parentModel, currentTurnModel, deploymentModel, personalModel),
|
||||
childChat.LastModelConfigID,
|
||||
)
|
||||
require.True(t, childChat.Mode.Valid)
|
||||
require.Equal(t, database.ChatModeExplore, childChat.Mode.ChatMode)
|
||||
require.False(t, childChat.PlanMode.Valid)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateChat_ExploreRootStartsWithoutMCPSnapshot(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -607,6 +607,71 @@ type UpdateChatModelOverrideRequest struct {
|
||||
ModelConfigID string `json:"model_config_id"`
|
||||
}
|
||||
|
||||
// ChatPersonalModelOverrideContext identifies which chat context the user
|
||||
// personal model override applies to.
|
||||
type ChatPersonalModelOverrideContext string
|
||||
|
||||
const (
|
||||
ChatPersonalModelOverrideContextRoot ChatPersonalModelOverrideContext = "root"
|
||||
ChatPersonalModelOverrideContextGeneral ChatPersonalModelOverrideContext = "general"
|
||||
ChatPersonalModelOverrideContextExplore ChatPersonalModelOverrideContext = "explore"
|
||||
)
|
||||
|
||||
// ChatPersonalModelOverrideMode identifies how a user personal model override
|
||||
// should resolve the effective model.
|
||||
type ChatPersonalModelOverrideMode string
|
||||
|
||||
const (
|
||||
ChatPersonalModelOverrideModeDeploymentDefault ChatPersonalModelOverrideMode = "deployment_default"
|
||||
ChatPersonalModelOverrideModeChatDefault ChatPersonalModelOverrideMode = "chat_default"
|
||||
ChatPersonalModelOverrideModeModel ChatPersonalModelOverrideMode = "model"
|
||||
)
|
||||
|
||||
// ChatPersonalModelOverride is a resolved user personal model override.
|
||||
type ChatPersonalModelOverride struct {
|
||||
Context ChatPersonalModelOverrideContext `json:"context"`
|
||||
Mode ChatPersonalModelOverrideMode `json:"mode"`
|
||||
ModelConfigID string `json:"model_config_id"`
|
||||
IsSet bool `json:"is_set"`
|
||||
IsMalformed bool `json:"is_malformed"`
|
||||
}
|
||||
|
||||
// ChatPersonalModelOverrideDeploymentDefaults describes the deployment-level
|
||||
// defaults used when a personal override selects deployment_default.
|
||||
type ChatPersonalModelOverrideDeploymentDefaults struct {
|
||||
General ChatModelOverrideResponse `json:"general"`
|
||||
Explore ChatModelOverrideResponse `json:"explore"`
|
||||
}
|
||||
|
||||
// UserChatPersonalModelOverridesResponse is the response body for user
|
||||
// personal model override settings.
|
||||
type UserChatPersonalModelOverridesResponse struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
Root ChatPersonalModelOverride `json:"root"`
|
||||
General ChatPersonalModelOverride `json:"general"`
|
||||
Explore ChatPersonalModelOverride `json:"explore"`
|
||||
DeploymentDefaults ChatPersonalModelOverrideDeploymentDefaults `json:"deployment_defaults"`
|
||||
}
|
||||
|
||||
// UpdateUserChatPersonalModelOverrideRequest is the request body for updating
|
||||
// a user personal model override.
|
||||
type UpdateUserChatPersonalModelOverrideRequest struct {
|
||||
Mode ChatPersonalModelOverrideMode `json:"mode"`
|
||||
ModelConfigID string `json:"model_config_id"`
|
||||
}
|
||||
|
||||
// ChatPersonalModelOverridesAdminSettings describes whether users may manage
|
||||
// personal model override settings.
|
||||
type ChatPersonalModelOverridesAdminSettings struct {
|
||||
AllowUsers bool `json:"allow_users"`
|
||||
}
|
||||
|
||||
// UpdateChatPersonalModelOverridesAdminSettingsRequest is the request body for
|
||||
// updating personal model override admin settings.
|
||||
type UpdateChatPersonalModelOverridesAdminSettingsRequest struct {
|
||||
AllowUsers bool `json:"allow_users"`
|
||||
}
|
||||
|
||||
// UserChatCustomPrompt is the request and response body for the
|
||||
// user chat custom prompt configuration endpoint.
|
||||
type UserChatCustomPrompt struct {
|
||||
@@ -2150,6 +2215,68 @@ func (c *ExperimentalClient) UpdateChatModelOverride(ctx context.Context, overri
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetChatPersonalModelOverridesAdminSettings returns the deployment-wide
|
||||
// personal model override admin settings.
|
||||
func (c *ExperimentalClient) GetChatPersonalModelOverridesAdminSettings(ctx context.Context) (ChatPersonalModelOverridesAdminSettings, error) {
|
||||
res, err := c.Request(ctx, http.MethodGet, "/api/experimental/chats/config/personal-model-overrides", nil)
|
||||
if err != nil {
|
||||
return ChatPersonalModelOverridesAdminSettings{}, err
|
||||
}
|
||||
defer res.Body.Close()
|
||||
if res.StatusCode != http.StatusOK {
|
||||
return ChatPersonalModelOverridesAdminSettings{}, ReadBodyAsError(res)
|
||||
}
|
||||
var resp ChatPersonalModelOverridesAdminSettings
|
||||
return resp, json.NewDecoder(res.Body).Decode(&resp)
|
||||
}
|
||||
|
||||
// UpdateChatPersonalModelOverridesAdminSettings updates the deployment-wide
|
||||
// personal model override admin settings.
|
||||
func (c *ExperimentalClient) UpdateChatPersonalModelOverridesAdminSettings(ctx context.Context, req UpdateChatPersonalModelOverridesAdminSettingsRequest) error {
|
||||
res, err := c.Request(ctx, http.MethodPut, "/api/experimental/chats/config/personal-model-overrides", req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer res.Body.Close()
|
||||
if res.StatusCode != http.StatusNoContent {
|
||||
return ReadBodyAsError(res)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetUserChatPersonalModelOverrides fetches the user's personal model
|
||||
// override settings.
|
||||
func (c *ExperimentalClient) GetUserChatPersonalModelOverrides(ctx context.Context) (UserChatPersonalModelOverridesResponse, error) {
|
||||
res, err := c.Request(ctx, http.MethodGet, "/api/experimental/chats/config/user-personal-model-overrides", nil)
|
||||
if err != nil {
|
||||
return UserChatPersonalModelOverridesResponse{}, err
|
||||
}
|
||||
defer res.Body.Close()
|
||||
if res.StatusCode != http.StatusOK {
|
||||
return UserChatPersonalModelOverridesResponse{}, ReadBodyAsError(res)
|
||||
}
|
||||
var resp UserChatPersonalModelOverridesResponse
|
||||
return resp, json.NewDecoder(res.Body).Decode(&resp)
|
||||
}
|
||||
|
||||
// UpdateUserChatPersonalModelOverride updates the user's personal model
|
||||
// override for the requested context.
|
||||
func (c *ExperimentalClient) UpdateUserChatPersonalModelOverride(ctx context.Context, override ChatPersonalModelOverrideContext, req UpdateUserChatPersonalModelOverrideRequest) error {
|
||||
path := fmt.Sprintf(
|
||||
"/api/experimental/chats/config/user-personal-model-overrides/%s",
|
||||
url.PathEscape(string(override)),
|
||||
)
|
||||
res, err := c.Request(ctx, http.MethodPut, path, req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer res.Body.Close()
|
||||
if res.StatusCode != http.StatusNoContent {
|
||||
return ReadBodyAsError(res)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetUserChatCustomPrompt fetches the user's custom chat prompt.
|
||||
func (c *ExperimentalClient) GetUserChatCustomPrompt(ctx context.Context) (UserChatCustomPrompt, error) {
|
||||
res, err := c.Request(ctx, http.MethodGet, "/api/experimental/chats/config/user-prompt", nil)
|
||||
|
||||
Generated
+81
@@ -2189,6 +2189,55 @@ export interface ChatModelsResponse {
|
||||
readonly providers: readonly ChatModelProvider[];
|
||||
}
|
||||
|
||||
// From codersdk/chats.go
|
||||
/**
|
||||
* ChatPersonalModelOverride is a resolved user personal model override.
|
||||
*/
|
||||
export interface ChatPersonalModelOverride {
|
||||
readonly context: ChatPersonalModelOverrideContext;
|
||||
readonly mode: ChatPersonalModelOverrideMode;
|
||||
readonly model_config_id: string;
|
||||
readonly is_set: boolean;
|
||||
readonly is_malformed: boolean;
|
||||
}
|
||||
|
||||
// From codersdk/chats.go
|
||||
export type ChatPersonalModelOverrideContext = "explore" | "general" | "root";
|
||||
|
||||
export const ChatPersonalModelOverrideContexts: ChatPersonalModelOverrideContext[] =
|
||||
["explore", "general", "root"];
|
||||
|
||||
// From codersdk/chats.go
|
||||
/**
|
||||
* ChatPersonalModelOverrideDeploymentDefaults describes the deployment-level
|
||||
* defaults used when a personal override selects deployment_default.
|
||||
*/
|
||||
export interface ChatPersonalModelOverrideDeploymentDefaults {
|
||||
readonly general: ChatModelOverrideResponse;
|
||||
readonly explore: ChatModelOverrideResponse;
|
||||
}
|
||||
|
||||
// From codersdk/chats.go
|
||||
export type ChatPersonalModelOverrideMode =
|
||||
| "chat_default"
|
||||
| "deployment_default"
|
||||
| "model";
|
||||
|
||||
export const ChatPersonalModelOverrideModes: ChatPersonalModelOverrideMode[] = [
|
||||
"chat_default",
|
||||
"deployment_default",
|
||||
"model",
|
||||
];
|
||||
|
||||
// From codersdk/chats.go
|
||||
/**
|
||||
* ChatPersonalModelOverridesAdminSettings describes whether users may manage
|
||||
* personal model override settings.
|
||||
*/
|
||||
export interface ChatPersonalModelOverridesAdminSettings {
|
||||
readonly allow_users: boolean;
|
||||
}
|
||||
|
||||
// From codersdk/chats.go
|
||||
export type ChatPlanMode = "plan";
|
||||
|
||||
@@ -7876,6 +7925,15 @@ export interface UpdateChatModelOverrideRequest {
|
||||
readonly model_config_id: string;
|
||||
}
|
||||
|
||||
// From codersdk/chats.go
|
||||
/**
|
||||
* UpdateChatPersonalModelOverridesAdminSettingsRequest is the request body for
|
||||
* updating personal model override admin settings.
|
||||
*/
|
||||
export interface UpdateChatPersonalModelOverridesAdminSettingsRequest {
|
||||
readonly allow_users: boolean;
|
||||
}
|
||||
|
||||
// From codersdk/chats.go
|
||||
/**
|
||||
* UpdateChatPlanModeInstructionsRequest is the request body for
|
||||
@@ -8181,6 +8239,16 @@ export interface UpdateUserChatDebugLoggingRequest {
|
||||
readonly debug_logging_enabled: boolean;
|
||||
}
|
||||
|
||||
// From codersdk/chats.go
|
||||
/**
|
||||
* UpdateUserChatPersonalModelOverrideRequest is the request body for updating
|
||||
* a user personal model override.
|
||||
*/
|
||||
export interface UpdateUserChatPersonalModelOverrideRequest {
|
||||
readonly mode: ChatPersonalModelOverrideMode;
|
||||
readonly model_config_id: string;
|
||||
}
|
||||
|
||||
// From codersdk/notifications.go
|
||||
export interface UpdateUserNotificationPreferences {
|
||||
readonly template_disabled_map: Record<string, boolean>;
|
||||
@@ -8492,6 +8560,19 @@ export interface UserChatDebugLoggingSettings {
|
||||
readonly forced_by_deployment: boolean;
|
||||
}
|
||||
|
||||
// From codersdk/chats.go
|
||||
/**
|
||||
* UserChatPersonalModelOverridesResponse is the response body for user
|
||||
* personal model override settings.
|
||||
*/
|
||||
export interface UserChatPersonalModelOverridesResponse {
|
||||
readonly enabled: boolean;
|
||||
readonly root: ChatPersonalModelOverride;
|
||||
readonly general: ChatPersonalModelOverride;
|
||||
readonly explore: ChatPersonalModelOverride;
|
||||
readonly deployment_defaults: ChatPersonalModelOverrideDeploymentDefaults;
|
||||
}
|
||||
|
||||
// From codersdk/chats.go
|
||||
/**
|
||||
* UserChatProviderConfig is a summary of a provider that allows
|
||||
|
||||
Reference in New Issue
Block a user