diff --git a/coderd/chatd/chatd.go b/coderd/chatd/chatd.go index ea3a2f87be..b1f1c2de00 100644 --- a/coderd/chatd/chatd.go +++ b/coderd/chatd/chatd.go @@ -2229,9 +2229,11 @@ func (p *Server) runChat( if instruction := p.resolveInstructions(ctx, chat, getWorkspaceConn); instruction != "" { prompt = chatprompt.InsertSystem(prompt, instruction) } + if userPrompt := p.resolveUserPrompt(ctx, chat.OwnerID); userPrompt != "" { + prompt = chatprompt.InsertSystem(prompt, userPrompt) + } - // Use the model config's context_limit as a fallback when the LLM - // provider doesn't include context_limit in its response metadata + // Use the model config's context_limit as a fallback when the LLM // provider doesn't include context_limit in its response metadata // (which is the common case). modelConfigContextLimit := modelConfig.ContextLimit @@ -2510,6 +2512,9 @@ func (p *Server) runChat( if instruction := p.resolveInstructions(reloadCtx, chat, getWorkspaceConn); instruction != "" { reloadedPrompt = chatprompt.InsertSystem(reloadedPrompt, instruction) } + if userPrompt := p.resolveUserPrompt(reloadCtx, chat.OwnerID); userPrompt != "" { + reloadedPrompt = chatprompt.InsertSystem(reloadedPrompt, userPrompt) + } return reloadedPrompt, nil }, @@ -2884,6 +2889,22 @@ func (p *Server) resolveInstructions( return instruction } +// resolveUserPrompt fetches the user's custom chat prompt from the +// database and wraps it in tags. Returns empty +// string if no prompt is set. +func (p *Server) resolveUserPrompt(ctx context.Context, userID uuid.UUID) string { + raw, err := p.db.GetUserChatCustomPrompt(ctx, userID) + if err != nil { + // sql.ErrNoRows is the normal "not set" case. + return "" + } + trimmed := strings.TrimSpace(raw) + if trimmed == "" { + return "" + } + return "\n" + trimmed + "\n" +} + func (p *Server) recoverStaleChats(ctx context.Context) { staleAfter := time.Now().Add(-p.inFlightChatStaleAfter) staleChats, err := p.db.GetStaleChats(ctx, staleAfter) diff --git a/coderd/chats.go b/coderd/chats.go index d3e4a557c1..24fd6f3d97 100644 --- a/coderd/chats.go +++ b/coderd/chats.go @@ -2301,6 +2301,72 @@ func (api *API) putChatSystemPrompt(rw http.ResponseWriter, r *http.Request) { 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) getUserChatCustomPrompt(rw http.ResponseWriter, r *http.Request) { + var ( + ctx = r.Context() + apiKey = httpmw.APIKey(r) + ) + + customPrompt, err := api.Database.GetUserChatCustomPrompt(ctx, apiKey.UserID) + if err != nil { + if !errors.Is(err, sql.ErrNoRows) { + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Error reading user chat custom prompt.", + Detail: err.Error(), + }) + return + } + + customPrompt = "" + } + + httpapi.Write(ctx, rw, http.StatusOK, codersdk.UserChatCustomPromptResponse{ + CustomPrompt: customPrompt, + }) +} + +// EXPERIMENTAL: this endpoint is experimental and is subject to change. +func (api *API) putUserChatCustomPrompt(rw http.ResponseWriter, r *http.Request) { + var ( + ctx = r.Context() + apiKey = httpmw.APIKey(r) + ) + + var params codersdk.UpdateUserChatCustomPromptRequest + if !httpapi.Read(ctx, rw, r, ¶ms) { + return + } + + trimmedPrompt := strings.TrimSpace(params.CustomPrompt) + // Apply the same 128 KiB limit as the deployment system prompt. + if len(trimmedPrompt) > maxSystemPromptLenBytes { + httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ + Message: "Custom prompt exceeds maximum length.", + Detail: fmt.Sprintf("Maximum length is %d bytes, got %d.", maxSystemPromptLenBytes, len(trimmedPrompt)), + }) + return + } + + updatedConfig, err := api.Database.UpdateUserChatCustomPrompt(ctx, database.UpdateUserChatCustomPromptParams{ + UserID: apiKey.UserID, + ChatCustomPrompt: trimmedPrompt, + }) + if err != nil { + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Error updating user chat custom prompt.", + Detail: err.Error(), + }) + return + } + + httpapi.Write(ctx, rw, http.StatusOK, codersdk.UserChatCustomPromptResponse{ + CustomPrompt: updatedConfig.Value, + }) +} + func (api *API) resolvedChatSystemPrompt(ctx context.Context) string { custom, err := api.Database.GetChatSystemPrompt(ctx) if err != nil { diff --git a/coderd/coderd.go b/coderd/coderd.go index 9a084ef59f..9d2a25c360 100644 --- a/coderd/coderd.go +++ b/coderd/coderd.go @@ -1130,6 +1130,8 @@ func New(options *Options) *API { r.Route("/config", func(r chi.Router) { r.Get("/system-prompt", api.getChatSystemPrompt) r.Put("/system-prompt", api.putChatSystemPrompt) + r.Get("/user-prompt", api.getUserChatCustomPrompt) + r.Put("/user-prompt", api.putUserChatCustomPrompt) }) // TODO(cian): place under /api/experimental/chats/config r.Route("/providers", func(r chi.Router) { @@ -1463,6 +1465,7 @@ func New(options *Options) *API { r.Put("/appearance", api.putUserAppearanceSettings) r.Get("/preferences", api.userPreferenceSettings) r.Put("/preferences", api.putUserPreferenceSettings) + r.Route("/password", func(r chi.Router) { r.Use(httpmw.RateLimit(options.LoginRateLimit, time.Minute)) r.Put("/", api.putUserPassword) diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index a803260676..a6f0a47e3e 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -3800,6 +3800,17 @@ func (q *querier) GetUserByID(ctx context.Context, id uuid.UUID) (database.User, return fetch(q.log, q.auth, q.db.GetUserByID)(ctx, id) } +func (q *querier) GetUserChatCustomPrompt(ctx context.Context, userID uuid.UUID) (string, error) { + u, err := q.db.GetUserByID(ctx, userID) + if err != nil { + return "", err + } + if err := q.authorizeContext(ctx, policy.ActionReadPersonal, u); err != nil { + return "", err + } + return q.db.GetUserChatCustomPrompt(ctx, userID) +} + func (q *querier) GetUserCount(ctx context.Context, includeSystem bool) (int64, error) { if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceSystem); err != nil { return 0, err @@ -6029,6 +6040,17 @@ func (q *querier) UpdateUsageEventsPostPublish(ctx context.Context, arg database return q.db.UpdateUsageEventsPostPublish(ctx, arg) } +func (q *querier) UpdateUserChatCustomPrompt(ctx context.Context, arg database.UpdateUserChatCustomPromptParams) (database.UserConfig, error) { + u, err := q.db.GetUserByID(ctx, arg.UserID) + if err != nil { + return database.UserConfig{}, err + } + if err := q.authorizeContext(ctx, policy.ActionUpdatePersonal, u); err != nil { + return database.UserConfig{}, err + } + return q.db.UpdateUserChatCustomPrompt(ctx, arg) +} + func (q *querier) UpdateUserDeletedByID(ctx context.Context, id uuid.UUID) error { return deleteQ(q.log, q.auth, q.db.GetUserByID, q.db.UpdateUserDeletedByID)(ctx, id) } diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index 5eb590af7b..bc0324002a 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -1922,6 +1922,20 @@ func (s *MethodTestSuite) TestUser() { dbm.EXPECT().GetUserTaskNotificationAlertDismissed(gomock.Any(), u.ID).Return(false, nil).AnyTimes() check.Args(u.ID).Asserts(u, policy.ActionReadPersonal).Returns(false) })) + s.Run("GetUserChatCustomPrompt", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + u := testutil.Fake(s.T(), faker, database.User{}) + dbm.EXPECT().GetUserByID(gomock.Any(), u.ID).Return(u, nil).AnyTimes() + dbm.EXPECT().GetUserChatCustomPrompt(gomock.Any(), u.ID).Return("my custom prompt", nil).AnyTimes() + check.Args(u.ID).Asserts(u, policy.ActionReadPersonal).Returns("my custom prompt") + })) + 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"} + arg := database.UpdateUserChatCustomPromptParams{UserID: u.ID, ChatCustomPrompt: uc.Value} + dbm.EXPECT().GetUserByID(gomock.Any(), u.ID).Return(u, nil).AnyTimes() + dbm.EXPECT().UpdateUserChatCustomPrompt(gomock.Any(), arg).Return(uc, nil).AnyTimes() + check.Args(arg).Asserts(u, policy.ActionUpdatePersonal).Returns(uc) + })) s.Run("UpdateUserTaskNotificationAlertDismissed", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { user := testutil.Fake(s.T(), faker, database.User{}) userConfig := database.UserConfig{UserID: user.ID, Key: "task_notification_alert_dismissed", Value: "false"} diff --git a/coderd/database/dbmetrics/querymetrics.go b/coderd/database/dbmetrics/querymetrics.go index 7739254bef..8436e6db9b 100644 --- a/coderd/database/dbmetrics/querymetrics.go +++ b/coderd/database/dbmetrics/querymetrics.go @@ -2295,6 +2295,14 @@ func (m queryMetricsStore) GetUserByID(ctx context.Context, id uuid.UUID) (datab return r0, r1 } +func (m queryMetricsStore) GetUserChatCustomPrompt(ctx context.Context, userID uuid.UUID) (string, error) { + start := time.Now() + r0, r1 := m.s.GetUserChatCustomPrompt(ctx, userID) + m.queryLatencies.WithLabelValues("GetUserChatCustomPrompt").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetUserChatCustomPrompt").Inc() + return r0, r1 +} + func (m queryMetricsStore) GetUserCount(ctx context.Context, includeSystem bool) (int64, error) { start := time.Now() r0, r1 := m.s.GetUserCount(ctx, includeSystem) @@ -4158,6 +4166,14 @@ func (m queryMetricsStore) UpdateUsageEventsPostPublish(ctx context.Context, arg return r0 } +func (m queryMetricsStore) UpdateUserChatCustomPrompt(ctx context.Context, arg database.UpdateUserChatCustomPromptParams) (database.UserConfig, error) { + start := time.Now() + r0, r1 := m.s.UpdateUserChatCustomPrompt(ctx, arg) + m.queryLatencies.WithLabelValues("UpdateUserChatCustomPrompt").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "UpdateUserChatCustomPrompt").Inc() + return r0, r1 +} + func (m queryMetricsStore) UpdateUserDeletedByID(ctx context.Context, id uuid.UUID) error { start := time.Now() r0 := m.s.UpdateUserDeletedByID(ctx, id) diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index 33bc2653d6..f914011018 100644 --- a/coderd/database/dbmock/dbmock.go +++ b/coderd/database/dbmock/dbmock.go @@ -4282,6 +4282,21 @@ func (mr *MockStoreMockRecorder) GetUserByID(ctx, id any) *gomock.Call { return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetUserByID", reflect.TypeOf((*MockStore)(nil).GetUserByID), ctx, id) } +// GetUserChatCustomPrompt mocks base method. +func (m *MockStore) GetUserChatCustomPrompt(ctx context.Context, userID uuid.UUID) (string, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetUserChatCustomPrompt", ctx, userID) + ret0, _ := ret[0].(string) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetUserChatCustomPrompt indicates an expected call of GetUserChatCustomPrompt. +func (mr *MockStoreMockRecorder) GetUserChatCustomPrompt(ctx, userID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetUserChatCustomPrompt", reflect.TypeOf((*MockStore)(nil).GetUserChatCustomPrompt), ctx, userID) +} + // GetUserCount mocks base method. func (m *MockStore) GetUserCount(ctx context.Context, includeSystem bool) (int64, error) { m.ctrl.T.Helper() @@ -7801,6 +7816,21 @@ func (mr *MockStoreMockRecorder) UpdateUsageEventsPostPublish(ctx, arg any) *gom return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateUsageEventsPostPublish", reflect.TypeOf((*MockStore)(nil).UpdateUsageEventsPostPublish), ctx, arg) } +// UpdateUserChatCustomPrompt mocks base method. +func (m *MockStore) UpdateUserChatCustomPrompt(ctx context.Context, arg database.UpdateUserChatCustomPromptParams) (database.UserConfig, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "UpdateUserChatCustomPrompt", ctx, arg) + ret0, _ := ret[0].(database.UserConfig) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// UpdateUserChatCustomPrompt indicates an expected call of UpdateUserChatCustomPrompt. +func (mr *MockStoreMockRecorder) UpdateUserChatCustomPrompt(ctx, arg any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateUserChatCustomPrompt", reflect.TypeOf((*MockStore)(nil).UpdateUserChatCustomPrompt), ctx, arg) +} + // UpdateUserDeletedByID mocks base method. func (m *MockStore) UpdateUserDeletedByID(ctx context.Context, id uuid.UUID) error { m.ctrl.T.Helper() diff --git a/coderd/database/querier.go b/coderd/database/querier.go index e29227e459..4c992997d9 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -490,6 +490,7 @@ type sqlcQuerier interface { GetUserActivityInsights(ctx context.Context, arg GetUserActivityInsightsParams) ([]GetUserActivityInsightsRow, error) GetUserByEmailOrUsername(ctx context.Context, arg GetUserByEmailOrUsernameParams) (User, error) GetUserByID(ctx context.Context, id uuid.UUID) (User, error) + GetUserChatCustomPrompt(ctx context.Context, userID uuid.UUID) (string, error) GetUserCount(ctx context.Context, includeSystem bool) (int64, error) // GetUserLatencyInsights returns the median and 95th percentile connection // latency that users have experienced. The result can be filtered on @@ -789,6 +790,7 @@ type sqlcQuerier interface { UpdateTemplateVersionFlagsByJobID(ctx context.Context, arg UpdateTemplateVersionFlagsByJobIDParams) error UpdateTemplateWorkspacesLastUsedAt(ctx context.Context, arg UpdateTemplateWorkspacesLastUsedAtParams) error UpdateUsageEventsPostPublish(ctx context.Context, arg UpdateUsageEventsPostPublishParams) error + UpdateUserChatCustomPrompt(ctx context.Context, arg UpdateUserChatCustomPromptParams) (UserConfig, error) UpdateUserDeletedByID(ctx context.Context, id uuid.UUID) error UpdateUserGithubComUserID(ctx context.Context, arg UpdateUserGithubComUserIDParams) error UpdateUserHashedOneTimePasscode(ctx context.Context, arg UpdateUserHashedOneTimePasscodeParams) error diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index 0f6c2db116..0552b7958b 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -18744,6 +18744,23 @@ func (q *sqlQuerier) GetUserByID(ctx context.Context, id uuid.UUID) (User, error return i, err } +const getUserChatCustomPrompt = `-- name: GetUserChatCustomPrompt :one +SELECT + value as chat_custom_prompt +FROM + user_configs +WHERE + user_id = $1 + AND key = 'chat_custom_prompt' +` + +func (q *sqlQuerier) GetUserChatCustomPrompt(ctx context.Context, userID uuid.UUID) (string, error) { + row := q.db.QueryRowContext(ctx, getUserChatCustomPrompt, userID) + var chat_custom_prompt string + err := row.Scan(&chat_custom_prompt) + return chat_custom_prompt, err +} + const getUserCount = `-- name: GetUserCount :one SELECT COUNT(*) @@ -19191,6 +19208,33 @@ func (q *sqlQuerier) UpdateInactiveUsersToDormant(ctx context.Context, arg Updat return items, nil } +const updateUserChatCustomPrompt = `-- name: UpdateUserChatCustomPrompt :one +INSERT INTO + user_configs (user_id, key, value) +VALUES + ($1, 'chat_custom_prompt', $2) +ON CONFLICT + ON CONSTRAINT user_configs_pkey +DO UPDATE +SET + value = $2 +WHERE user_configs.user_id = $1 + AND user_configs.key = 'chat_custom_prompt' +RETURNING user_id, key, value +` + +type UpdateUserChatCustomPromptParams struct { + UserID uuid.UUID `db:"user_id" json:"user_id"` + ChatCustomPrompt string `db:"chat_custom_prompt" json:"chat_custom_prompt"` +} + +func (q *sqlQuerier) UpdateUserChatCustomPrompt(ctx context.Context, arg UpdateUserChatCustomPromptParams) (UserConfig, error) { + row := q.db.QueryRowContext(ctx, updateUserChatCustomPrompt, arg.UserID, arg.ChatCustomPrompt) + var i UserConfig + err := row.Scan(&i.UserID, &i.Key, &i.Value) + return i, err +} + const updateUserDeletedByID = `-- name: UpdateUserDeletedByID :exec UPDATE users diff --git a/coderd/database/queries/users.sql b/coderd/database/queries/users.sql index 2e8649507a..26fc53084b 100644 --- a/coderd/database/queries/users.sql +++ b/coderd/database/queries/users.sql @@ -168,6 +168,29 @@ WHERE user_configs.user_id = @user_id AND user_configs.key = 'terminal_font' RETURNING *; +-- name: GetUserChatCustomPrompt :one +SELECT + value as chat_custom_prompt +FROM + user_configs +WHERE + user_id = @user_id + AND key = 'chat_custom_prompt'; + +-- name: UpdateUserChatCustomPrompt :one +INSERT INTO + user_configs (user_id, key, value) +VALUES + (@user_id, 'chat_custom_prompt', @chat_custom_prompt) +ON CONFLICT + ON CONSTRAINT user_configs_pkey +DO UPDATE +SET + value = @chat_custom_prompt +WHERE user_configs.user_id = @user_id + AND user_configs.key = 'chat_custom_prompt' +RETURNING *; + -- name: GetUserTaskNotificationAlertDismissed :one SELECT value::boolean as task_notification_alert_dismissed diff --git a/codersdk/chats.go b/codersdk/chats.go index 9100278090..f6f53e7342 100644 --- a/codersdk/chats.go +++ b/codersdk/chats.go @@ -211,6 +211,18 @@ type UpdateChatSystemPromptRequest struct { SystemPrompt string `json:"system_prompt"` } +// UserChatCustomPromptResponse is the response for getting a user's +// custom chat prompt. +type UserChatCustomPromptResponse struct { + CustomPrompt string `json:"custom_prompt"` +} + +// UpdateUserChatCustomPromptRequest is the request to update a user's +// custom chat prompt. +type UpdateUserChatCustomPromptRequest struct { + CustomPrompt string `json:"custom_prompt"` +} + // ChatProviderConfigSource describes how a provider entry is sourced. type ChatProviderConfigSource string @@ -725,6 +737,34 @@ func (c *Client) UpdateChatSystemPrompt(ctx context.Context, req UpdateChatSyste return nil } +// GetUserChatCustomPrompt fetches the user's custom chat prompt. +func (c *Client) GetUserChatCustomPrompt(ctx context.Context) (UserChatCustomPromptResponse, error) { + res, err := c.Request(ctx, http.MethodGet, "/api/experimental/chats/config/user-prompt", nil) + if err != nil { + return UserChatCustomPromptResponse{}, err + } + defer res.Body.Close() + if res.StatusCode != http.StatusOK { + return UserChatCustomPromptResponse{}, ReadBodyAsError(res) + } + var resp UserChatCustomPromptResponse + return resp, json.NewDecoder(res.Body).Decode(&resp) +} + +// UpdateUserChatCustomPrompt updates the user's custom chat prompt. +func (c *Client) UpdateUserChatCustomPrompt(ctx context.Context, req UpdateUserChatCustomPromptRequest) (UserChatCustomPromptResponse, error) { + res, err := c.Request(ctx, http.MethodPut, "/api/experimental/chats/config/user-prompt", req) + if err != nil { + return UserChatCustomPromptResponse{}, err + } + defer res.Body.Close() + if res.StatusCode != http.StatusOK { + return UserChatCustomPromptResponse{}, ReadBodyAsError(res) + } + var resp UserChatCustomPromptResponse + return resp, json.NewDecoder(res.Body).Decode(&resp) +} + // CreateChat creates a new chat. func (c *Client) CreateChat(ctx context.Context, req CreateChatRequest) (Chat, error) { res, err := c.Request(ctx, http.MethodPost, "/api/experimental/chats", req) diff --git a/site/src/api/typesGenerated.ts b/site/src/api/typesGenerated.ts index 0adb9cd195..7f6aadcf40 100644 --- a/site/src/api/typesGenerated.ts +++ b/site/src/api/typesGenerated.ts @@ -6479,6 +6479,15 @@ export interface UpdateUserAppearanceSettingsRequest { readonly terminal_font: TerminalFontName; } +// From codersdk/chats.go +/** + * UpdateUserChatCustomPromptRequest is the request to update a user's + * custom chat prompt. + */ +export interface UpdateUserChatCustomPromptRequest { + readonly custom_prompt: string; +} + // From codersdk/notifications.go export interface UpdateUserNotificationPreferences { readonly template_disabled_map: Record; @@ -6703,6 +6712,15 @@ export interface UserAppearanceSettings { readonly terminal_font: TerminalFontName; } +// From codersdk/chats.go +/** + * UserChatCustomPromptResponse is the response for getting a user's + * custom chat prompt. + */ +export interface UserChatCustomPromptResponse { + readonly custom_prompt: string; +} + // From codersdk/insights.go /** * UserLatency shows the connection latency for a user.