From 44c68c79bfa712d994fb8227ac9a8d3ab7f7968f Mon Sep 17 00:00:00 2001 From: Michael Suchacz <203725896+ibetitsmike@users.noreply.github.com> Date: Fri, 10 Jul 2026 16:07:37 +0000 Subject: [PATCH 1/2] feat(coderd/x/chatd): add synthetic gateway keys --- coderd/database/dbauthz/dbauthz.go | 52 +++ coderd/database/dbauthz/dbauthz_test.go | 44 +++ coderd/database/dbmetrics/querymetrics.go | 32 ++ coderd/database/dbmock/dbmock.go | 60 +++ coderd/database/dump.sql | 25 +- coderd/database/foreign_key_constraint.go | 4 +- .../000546_chat_synthetic_api_keys.down.sql | 27 ++ .../000546_chat_synthetic_api_keys.up.sql | 12 + coderd/database/migrations/migrate_test.go | 59 +++ .../000546_chat_synthetic_api_keys.up.sql | 5 + coderd/database/models.go | 7 + coderd/database/querier.go | 4 + coderd/database/queries.sql.go | 93 +++++ coderd/database/queries/chats.sql | 17 + coderd/database/queries/users.sql | 5 + coderd/database/unique_constraint.go | 2 + coderd/exp_chats.go | 6 - coderd/exp_chats_test.go | 8 +- coderd/rbac/authz.go | 3 +- coderd/x/chatd/ARCHITECTURE.md | 10 + coderd/x/chatd/chatd.go | 82 ++-- coderd/x/chatd/chatd_helpers_test.go | 2 - coderd/x/chatd/chatd_internal_test.go | 4 + coderd/x/chatd/chatd_test.go | 264 ++++++------- coderd/x/chatd/generation.go | 7 +- coderd/x/chatd/generation_preparer.go | 16 +- coderd/x/chatd/helpers_test.go | 6 - coderd/x/chatd/integration_responses_test.go | 4 - coderd/x/chatd/model_routing.go | 17 - coderd/x/chatd/model_routing_internal_test.go | 211 ----------- coderd/x/chatd/quickgen.go | 11 +- coderd/x/chatd/subagent.go | 28 +- .../x/chatd/subagent_context_internal_test.go | 3 - coderd/x/chatd/subagent_internal_test.go | 351 ++++-------------- coderd/x/chatd/subagent_test.go | 8 - coderd/x/chatd/synthetickey.go | 121 ++++++ coderd/x/chatd/synthetickey_internal_test.go | 176 +++++++++ .../x/chatd/title_override_internal_test.go | 161 ++++---- 38 files changed, 1071 insertions(+), 876 deletions(-) create mode 100644 coderd/database/migrations/000546_chat_synthetic_api_keys.down.sql create mode 100644 coderd/database/migrations/000546_chat_synthetic_api_keys.up.sql create mode 100644 coderd/database/migrations/testdata/fixtures/000546_chat_synthetic_api_keys.up.sql create mode 100644 coderd/x/chatd/synthetickey.go create mode 100644 coderd/x/chatd/synthetickey_internal_test.go diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index 5a18c5fb8f3..f04b6502286 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -473,6 +473,27 @@ var ( }.WithCachedASTValue() } + subjectChatdKeyMinter = func(userID uuid.UUID) rbac.Subject { + return rbac.Subject{ + Type: rbac.SubjectTypeChatdKeyMinter, + FriendlyName: "Chatd Key Minter", + ID: userID.String(), + Roles: rbac.Roles([]rbac.Role{ + { + Identifier: rbac.RoleIdentifier{Name: "chatdkeyminter"}, + DisplayName: "Chatd Key Minter", + Site: []rbac.Permission{}, + User: rbac.Permissions(map[string][]policy.Action{ + rbac.ResourceApiKey.Type: {policy.ActionRead, policy.ActionCreate, policy.ActionUpdate, policy.ActionDelete}, + rbac.ResourceUser.Type: {policy.ActionReadPersonal}, + }), + ByOrgID: map[string]rbac.OrgPermissions{}, + }, + }), + Scope: rbac.ScopeAll, + }.WithCachedASTValue() + } + subjectSystemRestricted = rbac.Subject{ Type: rbac.SubjectTypeSystemRestricted, FriendlyName: "System", @@ -874,6 +895,12 @@ func AsAPIKeyRevoker(ctx context.Context, userID uuid.UUID) context.Context { return As(ctx, subjectAPIKeyRevoker(userID)) } +// AsChatdKeyMinter returns a context with an actor that manages the synthetic +// gateway API key owned by the specified user. +func AsChatdKeyMinter(ctx context.Context, userID uuid.UUID) context.Context { + return As(ctx, subjectChatdKeyMinter(userID)) +} + // AsSystemRestricted returns a context with an actor that has permissions // required for various system operations (login, logout, metrics cache). // DO NOT USE THIS UNLESS YOU HAVE ABSOLUTELY NO OTHER CHOICE. Prefer using a @@ -3494,6 +3521,13 @@ func (q *querier) GetChatStreamSyncRows(ctx context.Context, ids []uuid.UUID) ([ return q.db.GetChatStreamSyncRows(ctx, ids) } +func (q *querier) GetChatSyntheticAPIKeyByUserID(ctx context.Context, userID uuid.UUID) (database.ChatSyntheticApiKey, error) { + if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceApiKey.WithOwner(userID.String())); err != nil { + return database.ChatSyntheticApiKey{}, err + } + return q.db.GetChatSyntheticAPIKeyByUserID(ctx, userID) +} + func (q *querier) GetChatSystemPrompt(ctx context.Context) (string, error) { // The system prompt is a deployment-wide setting read during chat // creation by every authenticated user, so no RBAC policy check @@ -5072,6 +5106,10 @@ func (q *querier) GetUserCount(ctx context.Context, includeSystem bool) (int64, return q.db.GetUserCount(ctx, includeSystem) } +func (q *querier) GetUserForChatSyntheticAPIKeyByID(ctx context.Context, id uuid.UUID) (database.User, error) { + return fetchWithAction(q.log, q.auth, policy.ActionReadPersonal, q.db.GetUserForChatSyntheticAPIKeyByID)(ctx, id) +} + func (q *querier) GetUserGroupSpendLimit(ctx context.Context, arg database.GetUserGroupSpendLimitParams) (int64, error) { if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceChat.WithOwner(arg.UserID.String())); err != nil { return 0, err @@ -6011,6 +6049,13 @@ func (q *querier) InsertChatQueuedMessageWithCreator(ctx context.Context, arg da return q.db.InsertChatQueuedMessageWithCreator(ctx, arg) } +func (q *querier) InsertChatSyntheticAPIKey(ctx context.Context, arg database.InsertChatSyntheticAPIKeyParams) (int64, error) { + if err := q.authorizeContext(ctx, policy.ActionCreate, rbac.ResourceApiKey.WithOwner(arg.UserID.String())); err != nil { + return 0, err + } + return q.db.InsertChatSyntheticAPIKey(ctx, arg) +} + func (q *querier) InsertCryptoKey(ctx context.Context, arg database.InsertCryptoKeyParams) (database.CryptoKey, error) { if err := q.authorizeContext(ctx, policy.ActionCreate, rbac.ResourceCryptoKey); err != nil { return database.CryptoKey{}, err @@ -7402,6 +7447,13 @@ func (q *querier) UpdateChatStatus(ctx context.Context, arg database.UpdateChatS return q.db.UpdateChatStatus(ctx, arg) } +func (q *querier) UpdateChatSyntheticAPIKey(ctx context.Context, arg database.UpdateChatSyntheticAPIKeyParams) (int64, error) { + if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceApiKey.WithOwner(arg.UserID.String())); err != nil { + return 0, err + } + return q.db.UpdateChatSyntheticAPIKey(ctx, arg) +} + func (q *querier) UpdateChatTitleByID(ctx context.Context, arg database.UpdateChatTitleByIDParams) (database.Chat, error) { chat, err := q.db.GetChatByID(ctx, arg.ID) if err != nil { diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index d3263972e6b..7b4ad1f8f81 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -339,6 +339,32 @@ func defaultIPAddress() pqtype.Inet { } } +func (s *MethodTestSuite) TestChatSyntheticAPIKey() { + s.Run("GetUserForChatSyntheticAPIKeyByID", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + user := testutil.Fake(s.T(), faker, database.User{}) + dbm.EXPECT().GetUserForChatSyntheticAPIKeyByID(gomock.Any(), user.ID).Return(user, nil).AnyTimes() + check.Args(user.ID).Asserts(user, policy.ActionReadPersonal).Returns(user) + })) + s.Run("GetChatSyntheticAPIKeyByUserID", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + user := testutil.Fake(s.T(), faker, database.User{}) + mapping := testutil.Fake(s.T(), faker, database.ChatSyntheticApiKey{UserID: user.ID}) + dbm.EXPECT().GetChatSyntheticAPIKeyByUserID(gomock.Any(), user.ID).Return(mapping, nil).AnyTimes() + check.Args(user.ID).Asserts(rbac.ResourceApiKey.WithOwner(user.ID.String()), policy.ActionRead).Returns(mapping) + })) + s.Run("InsertChatSyntheticAPIKey", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + user := testutil.Fake(s.T(), faker, database.User{}) + arg := database.InsertChatSyntheticAPIKeyParams{UserID: user.ID, APIKeyID: uuid.NewString()} + dbm.EXPECT().InsertChatSyntheticAPIKey(gomock.Any(), arg).Return(int64(1), nil).AnyTimes() + check.Args(arg).Asserts(rbac.ResourceApiKey.WithOwner(user.ID.String()), policy.ActionCreate).Returns(int64(1)) + })) + s.Run("UpdateChatSyntheticAPIKey", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + user := testutil.Fake(s.T(), faker, database.User{}) + arg := database.UpdateChatSyntheticAPIKeyParams{UserID: user.ID, OldApiKeyID: uuid.NewString(), NewApiKeyID: uuid.NewString()} + dbm.EXPECT().UpdateChatSyntheticAPIKey(gomock.Any(), arg).Return(int64(1), nil).AnyTimes() + check.Args(arg).Asserts(rbac.ResourceApiKey.WithOwner(user.ID.String()), policy.ActionUpdate).Returns(int64(1)) + })) +} + func (s *MethodTestSuite) TestAPIKey() { s.Run("DeleteAPIKeyByID", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { key := testutil.Fake(s.T(), faker, database.APIKey{}) @@ -7481,6 +7507,24 @@ func TestAsAPIKeyRevoker(t *testing.T) { }) } +func TestAsChatdKeyMinter(t *testing.T) { + t.Parallel() + + userID := uuid.New() + ctx := dbauthz.AsChatdKeyMinter(context.Background(), userID) + actor, ok := dbauthz.ActorFromContext(ctx) + require.True(t, ok) + require.Equal(t, rbac.SubjectTypeChatdKeyMinter, actor.Type) + require.Equal(t, userID.String(), actor.ID) + + auth := rbac.NewStrictCachingAuthorizer(prometheus.NewRegistry()) + for _, action := range []policy.Action{policy.ActionCreate, policy.ActionRead, policy.ActionUpdate, policy.ActionDelete} { + require.NoError(t, auth.Authorize(ctx, actor, action, rbac.ResourceApiKey.WithOwner(userID.String()))) + require.Error(t, auth.Authorize(ctx, actor, action, rbac.ResourceApiKey.WithOwner(uuid.NewString()))) + } + require.NoError(t, auth.Authorize(ctx, actor, policy.ActionReadPersonal, rbac.ResourceUserObject(userID))) +} + func TestAsChatd(t *testing.T) { t.Parallel() diff --git a/coderd/database/dbmetrics/querymetrics.go b/coderd/database/dbmetrics/querymetrics.go index 16ba3084ebd..307d6908a15 100644 --- a/coderd/database/dbmetrics/querymetrics.go +++ b/coderd/database/dbmetrics/querymetrics.go @@ -1769,6 +1769,14 @@ func (m queryMetricsStore) GetChatStreamSyncRows(ctx context.Context, ids []uuid return r0, r1 } +func (m queryMetricsStore) GetChatSyntheticAPIKeyByUserID(ctx context.Context, userID uuid.UUID) (database.ChatSyntheticApiKey, error) { + start := time.Now() + r0, r1 := m.s.GetChatSyntheticAPIKeyByUserID(ctx, userID) + m.queryLatencies.WithLabelValues("GetChatSyntheticAPIKeyByUserID").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetChatSyntheticAPIKeyByUserID").Inc() + return r0, r1 +} + func (m queryMetricsStore) GetChatSystemPrompt(ctx context.Context) (string, error) { start := time.Now() r0, r1 := m.s.GetChatSystemPrompt(ctx) @@ -3273,6 +3281,14 @@ func (m queryMetricsStore) GetUserCount(ctx context.Context, includeSystem bool) return r0, r1 } +func (m queryMetricsStore) GetUserForChatSyntheticAPIKeyByID(ctx context.Context, id uuid.UUID) (database.User, error) { + start := time.Now() + r0, r1 := m.s.GetUserForChatSyntheticAPIKeyByID(ctx, id) + m.queryLatencies.WithLabelValues("GetUserForChatSyntheticAPIKeyByID").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetUserForChatSyntheticAPIKeyByID").Inc() + return r0, r1 +} + func (m queryMetricsStore) GetUserGroupSpendLimit(ctx context.Context, userID database.GetUserGroupSpendLimitParams) (int64, error) { start := time.Now() r0, r1 := m.s.GetUserGroupSpendLimit(ctx, userID) @@ -4121,6 +4137,14 @@ func (m queryMetricsStore) InsertChatQueuedMessageWithCreator(ctx context.Contex return r0, r1 } +func (m queryMetricsStore) InsertChatSyntheticAPIKey(ctx context.Context, arg database.InsertChatSyntheticAPIKeyParams) (int64, error) { + start := time.Now() + r0, r1 := m.s.InsertChatSyntheticAPIKey(ctx, arg) + m.queryLatencies.WithLabelValues("InsertChatSyntheticAPIKey").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "InsertChatSyntheticAPIKey").Inc() + return r0, r1 +} + func (m queryMetricsStore) InsertCryptoKey(ctx context.Context, arg database.InsertCryptoKeyParams) (database.CryptoKey, error) { start := time.Now() r0, r1 := m.s.InsertCryptoKey(ctx, arg) @@ -5265,6 +5289,14 @@ func (m queryMetricsStore) UpdateChatStatus(ctx context.Context, arg database.Up return r0, r1 } +func (m queryMetricsStore) UpdateChatSyntheticAPIKey(ctx context.Context, arg database.UpdateChatSyntheticAPIKeyParams) (int64, error) { + start := time.Now() + r0, r1 := m.s.UpdateChatSyntheticAPIKey(ctx, arg) + m.queryLatencies.WithLabelValues("UpdateChatSyntheticAPIKey").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "UpdateChatSyntheticAPIKey").Inc() + return r0, r1 +} + func (m queryMetricsStore) UpdateChatTitleByID(ctx context.Context, arg database.UpdateChatTitleByIDParams) (database.Chat, error) { start := time.Now() r0, r1 := m.s.UpdateChatTitleByID(ctx, arg) diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index 67e42764a88..64422a6e3d3 100644 --- a/coderd/database/dbmock/dbmock.go +++ b/coderd/database/dbmock/dbmock.go @@ -3268,6 +3268,21 @@ func (mr *MockStoreMockRecorder) GetChatStreamSyncRows(ctx, ids any) *gomock.Cal return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatStreamSyncRows", reflect.TypeOf((*MockStore)(nil).GetChatStreamSyncRows), ctx, ids) } +// GetChatSyntheticAPIKeyByUserID mocks base method. +func (m *MockStore) GetChatSyntheticAPIKeyByUserID(ctx context.Context, userID uuid.UUID) (database.ChatSyntheticApiKey, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetChatSyntheticAPIKeyByUserID", ctx, userID) + ret0, _ := ret[0].(database.ChatSyntheticApiKey) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetChatSyntheticAPIKeyByUserID indicates an expected call of GetChatSyntheticAPIKeyByUserID. +func (mr *MockStoreMockRecorder) GetChatSyntheticAPIKeyByUserID(ctx, userID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatSyntheticAPIKeyByUserID", reflect.TypeOf((*MockStore)(nil).GetChatSyntheticAPIKeyByUserID), ctx, userID) +} + // GetChatSystemPrompt mocks base method. func (m *MockStore) GetChatSystemPrompt(ctx context.Context) (string, error) { m.ctrl.T.Helper() @@ -6118,6 +6133,21 @@ func (mr *MockStoreMockRecorder) GetUserCount(ctx, includeSystem any) *gomock.Ca return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetUserCount", reflect.TypeOf((*MockStore)(nil).GetUserCount), ctx, includeSystem) } +// GetUserForChatSyntheticAPIKeyByID mocks base method. +func (m *MockStore) GetUserForChatSyntheticAPIKeyByID(ctx context.Context, id uuid.UUID) (database.User, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetUserForChatSyntheticAPIKeyByID", ctx, id) + ret0, _ := ret[0].(database.User) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetUserForChatSyntheticAPIKeyByID indicates an expected call of GetUserForChatSyntheticAPIKeyByID. +func (mr *MockStoreMockRecorder) GetUserForChatSyntheticAPIKeyByID(ctx, id any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetUserForChatSyntheticAPIKeyByID", reflect.TypeOf((*MockStore)(nil).GetUserForChatSyntheticAPIKeyByID), ctx, id) +} + // GetUserGroupSpendLimit mocks base method. func (m *MockStore) GetUserGroupSpendLimit(ctx context.Context, arg database.GetUserGroupSpendLimitParams) (int64, error) { m.ctrl.T.Helper() @@ -7720,6 +7750,21 @@ func (mr *MockStoreMockRecorder) InsertChatQueuedMessageWithCreator(ctx, arg any return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "InsertChatQueuedMessageWithCreator", reflect.TypeOf((*MockStore)(nil).InsertChatQueuedMessageWithCreator), ctx, arg) } +// InsertChatSyntheticAPIKey mocks base method. +func (m *MockStore) InsertChatSyntheticAPIKey(ctx context.Context, arg database.InsertChatSyntheticAPIKeyParams) (int64, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "InsertChatSyntheticAPIKey", ctx, arg) + ret0, _ := ret[0].(int64) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// InsertChatSyntheticAPIKey indicates an expected call of InsertChatSyntheticAPIKey. +func (mr *MockStoreMockRecorder) InsertChatSyntheticAPIKey(ctx, arg any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "InsertChatSyntheticAPIKey", reflect.TypeOf((*MockStore)(nil).InsertChatSyntheticAPIKey), ctx, arg) +} + // InsertCryptoKey mocks base method. func (m *MockStore) InsertCryptoKey(ctx context.Context, arg database.InsertCryptoKeyParams) (database.CryptoKey, error) { m.ctrl.T.Helper() @@ -9920,6 +9965,21 @@ func (mr *MockStoreMockRecorder) UpdateChatStatus(ctx, arg any) *gomock.Call { return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateChatStatus", reflect.TypeOf((*MockStore)(nil).UpdateChatStatus), ctx, arg) } +// UpdateChatSyntheticAPIKey mocks base method. +func (m *MockStore) UpdateChatSyntheticAPIKey(ctx context.Context, arg database.UpdateChatSyntheticAPIKeyParams) (int64, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "UpdateChatSyntheticAPIKey", ctx, arg) + ret0, _ := ret[0].(int64) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// UpdateChatSyntheticAPIKey indicates an expected call of UpdateChatSyntheticAPIKey. +func (mr *MockStoreMockRecorder) UpdateChatSyntheticAPIKey(ctx, arg any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateChatSyntheticAPIKey", reflect.TypeOf((*MockStore)(nil).UpdateChatSyntheticAPIKey), ctx, arg) +} + // UpdateChatTitleByID mocks base method. func (m *MockStore) UpdateChatTitleByID(ctx context.Context, arg database.UpdateChatTitleByIDParams) (database.Chat, error) { m.ctrl.T.Helper() diff --git a/coderd/database/dump.sql b/coderd/database/dump.sql index a8e815ec25a..e36ad5bb568 100644 --- a/coderd/database/dump.sql +++ b/coderd/database/dump.sql @@ -2033,6 +2033,13 @@ CREATE SEQUENCE chat_queued_messages_id_seq ALTER SEQUENCE chat_queued_messages_id_seq OWNED BY chat_queued_messages.id; +CREATE TABLE chat_synthetic_api_keys ( + user_id uuid NOT NULL, + api_key_id text NOT NULL, + created_at timestamp with time zone DEFAULT now() NOT NULL, + updated_at timestamp with time zone DEFAULT now() NOT NULL +); + CREATE TABLE chat_usage_limit_config ( id bigint NOT NULL, singleton boolean DEFAULT true NOT NULL, @@ -4301,6 +4308,12 @@ ALTER TABLE ONLY chat_model_configs ALTER TABLE ONLY chat_queued_messages ADD CONSTRAINT chat_queued_messages_pkey PRIMARY KEY (id); +ALTER TABLE ONLY chat_synthetic_api_keys + ADD CONSTRAINT chat_synthetic_api_keys_api_key_id_key UNIQUE (api_key_id); + +ALTER TABLE ONLY chat_synthetic_api_keys + ADD CONSTRAINT chat_synthetic_api_keys_pkey PRIMARY KEY (user_id); + ALTER TABLE ONLY chat_usage_limit_config ADD CONSTRAINT chat_usage_limit_config_pkey PRIMARY KEY (id); @@ -5139,9 +5152,6 @@ ALTER TABLE ONLY chat_files ALTER TABLE ONLY chat_heartbeats ADD CONSTRAINT chat_heartbeats_chat_id_fkey FOREIGN KEY (chat_id) REFERENCES chats(id) ON DELETE CASCADE; -ALTER TABLE ONLY chat_messages - ADD CONSTRAINT chat_messages_api_key_id_fkey FOREIGN KEY (api_key_id) REFERENCES api_keys(id) ON DELETE SET NULL; - ALTER TABLE ONLY chat_messages ADD CONSTRAINT chat_messages_chat_id_fkey FOREIGN KEY (chat_id) REFERENCES chats(id) ON DELETE CASCADE; @@ -5157,12 +5167,15 @@ ALTER TABLE ONLY chat_model_configs ALTER TABLE ONLY chat_model_configs ADD CONSTRAINT chat_model_configs_updated_by_fkey FOREIGN KEY (updated_by) REFERENCES users(id); -ALTER TABLE ONLY chat_queued_messages - ADD CONSTRAINT chat_queued_messages_api_key_id_fkey FOREIGN KEY (api_key_id) REFERENCES api_keys(id) ON DELETE SET NULL; - ALTER TABLE ONLY chat_queued_messages ADD CONSTRAINT chat_queued_messages_chat_id_fkey FOREIGN KEY (chat_id) REFERENCES chats(id) ON DELETE CASCADE; +ALTER TABLE ONLY chat_synthetic_api_keys + ADD CONSTRAINT chat_synthetic_api_keys_api_key_id_fkey FOREIGN KEY (api_key_id) REFERENCES api_keys(id) ON DELETE CASCADE; + +ALTER TABLE ONLY chat_synthetic_api_keys + ADD CONSTRAINT chat_synthetic_api_keys_user_id_fkey FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE; + ALTER TABLE ONLY chats ADD CONSTRAINT chats_agent_id_fkey FOREIGN KEY (agent_id) REFERENCES workspace_agents(id) ON DELETE SET NULL; diff --git a/coderd/database/foreign_key_constraint.go b/coderd/database/foreign_key_constraint.go index 4b40dcf679e..94bb3f0ed22 100644 --- a/coderd/database/foreign_key_constraint.go +++ b/coderd/database/foreign_key_constraint.go @@ -24,14 +24,14 @@ const ( ForeignKeyChatFilesOrganizationID ForeignKeyConstraint = "chat_files_organization_id_fkey" // ALTER TABLE ONLY chat_files ADD CONSTRAINT chat_files_organization_id_fkey FOREIGN KEY (organization_id) REFERENCES organizations(id) ON DELETE CASCADE; ForeignKeyChatFilesOwnerID ForeignKeyConstraint = "chat_files_owner_id_fkey" // ALTER TABLE ONLY chat_files ADD CONSTRAINT chat_files_owner_id_fkey FOREIGN KEY (owner_id) REFERENCES users(id) ON DELETE CASCADE; ForeignKeyChatHeartbeatsChatID ForeignKeyConstraint = "chat_heartbeats_chat_id_fkey" // ALTER TABLE ONLY chat_heartbeats ADD CONSTRAINT chat_heartbeats_chat_id_fkey FOREIGN KEY (chat_id) REFERENCES chats(id) ON DELETE CASCADE; - ForeignKeyChatMessagesAPIKeyID ForeignKeyConstraint = "chat_messages_api_key_id_fkey" // ALTER TABLE ONLY chat_messages ADD CONSTRAINT chat_messages_api_key_id_fkey FOREIGN KEY (api_key_id) REFERENCES api_keys(id) ON DELETE SET NULL; ForeignKeyChatMessagesChatID ForeignKeyConstraint = "chat_messages_chat_id_fkey" // ALTER TABLE ONLY chat_messages ADD CONSTRAINT chat_messages_chat_id_fkey FOREIGN KEY (chat_id) REFERENCES chats(id) ON DELETE CASCADE; ForeignKeyChatMessagesModelConfigID ForeignKeyConstraint = "chat_messages_model_config_id_fkey" // ALTER TABLE ONLY chat_messages ADD CONSTRAINT chat_messages_model_config_id_fkey FOREIGN KEY (model_config_id) REFERENCES chat_model_configs(id); ForeignKeyChatModelConfigsAIProviderID ForeignKeyConstraint = "chat_model_configs_ai_provider_id_fkey" // ALTER TABLE ONLY chat_model_configs ADD CONSTRAINT chat_model_configs_ai_provider_id_fkey FOREIGN KEY (ai_provider_id) REFERENCES ai_providers(id); ForeignKeyChatModelConfigsCreatedBy ForeignKeyConstraint = "chat_model_configs_created_by_fkey" // ALTER TABLE ONLY chat_model_configs ADD CONSTRAINT chat_model_configs_created_by_fkey FOREIGN KEY (created_by) REFERENCES users(id); ForeignKeyChatModelConfigsUpdatedBy ForeignKeyConstraint = "chat_model_configs_updated_by_fkey" // ALTER TABLE ONLY chat_model_configs ADD CONSTRAINT chat_model_configs_updated_by_fkey FOREIGN KEY (updated_by) REFERENCES users(id); - ForeignKeyChatQueuedMessagesAPIKeyID ForeignKeyConstraint = "chat_queued_messages_api_key_id_fkey" // ALTER TABLE ONLY chat_queued_messages ADD CONSTRAINT chat_queued_messages_api_key_id_fkey FOREIGN KEY (api_key_id) REFERENCES api_keys(id) ON DELETE SET NULL; ForeignKeyChatQueuedMessagesChatID ForeignKeyConstraint = "chat_queued_messages_chat_id_fkey" // ALTER TABLE ONLY chat_queued_messages ADD CONSTRAINT chat_queued_messages_chat_id_fkey FOREIGN KEY (chat_id) REFERENCES chats(id) ON DELETE CASCADE; + ForeignKeyChatSyntheticAPIKeysAPIKeyID ForeignKeyConstraint = "chat_synthetic_api_keys_api_key_id_fkey" // ALTER TABLE ONLY chat_synthetic_api_keys ADD CONSTRAINT chat_synthetic_api_keys_api_key_id_fkey FOREIGN KEY (api_key_id) REFERENCES api_keys(id) ON DELETE CASCADE; + ForeignKeyChatSyntheticAPIKeysUserID ForeignKeyConstraint = "chat_synthetic_api_keys_user_id_fkey" // ALTER TABLE ONLY chat_synthetic_api_keys ADD CONSTRAINT chat_synthetic_api_keys_user_id_fkey FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE; ForeignKeyChatsAgentID ForeignKeyConstraint = "chats_agent_id_fkey" // ALTER TABLE ONLY chats ADD CONSTRAINT chats_agent_id_fkey FOREIGN KEY (agent_id) REFERENCES workspace_agents(id) ON DELETE SET NULL; ForeignKeyChatsBuildID ForeignKeyConstraint = "chats_build_id_fkey" // ALTER TABLE ONLY chats ADD CONSTRAINT chats_build_id_fkey FOREIGN KEY (build_id) REFERENCES workspace_builds(id) ON DELETE SET NULL; ForeignKeyChatsLastModelConfigID ForeignKeyConstraint = "chats_last_model_config_id_fkey" // ALTER TABLE ONLY chats ADD CONSTRAINT chats_last_model_config_id_fkey FOREIGN KEY (last_model_config_id) REFERENCES chat_model_configs(id); diff --git a/coderd/database/migrations/000546_chat_synthetic_api_keys.down.sql b/coderd/database/migrations/000546_chat_synthetic_api_keys.down.sql new file mode 100644 index 00000000000..88c6ffc34da --- /dev/null +++ b/coderd/database/migrations/000546_chat_synthetic_api_keys.down.sql @@ -0,0 +1,27 @@ +DROP TABLE chat_synthetic_api_keys; + +UPDATE chat_messages +SET api_key_id = NULL +WHERE api_key_id IS NOT NULL + AND NOT EXISTS ( + SELECT 1 + FROM api_keys + WHERE api_keys.id = chat_messages.api_key_id + ); + +UPDATE chat_queued_messages +SET api_key_id = NULL +WHERE api_key_id IS NOT NULL + AND NOT EXISTS ( + SELECT 1 + FROM api_keys + WHERE api_keys.id = chat_queued_messages.api_key_id + ); + +ALTER TABLE chat_messages +ADD CONSTRAINT chat_messages_api_key_id_fkey +FOREIGN KEY (api_key_id) REFERENCES api_keys(id) ON DELETE SET NULL; + +ALTER TABLE chat_queued_messages +ADD CONSTRAINT chat_queued_messages_api_key_id_fkey +FOREIGN KEY (api_key_id) REFERENCES api_keys(id) ON DELETE SET NULL; diff --git a/coderd/database/migrations/000546_chat_synthetic_api_keys.up.sql b/coderd/database/migrations/000546_chat_synthetic_api_keys.up.sql new file mode 100644 index 00000000000..bde888b0dcc --- /dev/null +++ b/coderd/database/migrations/000546_chat_synthetic_api_keys.up.sql @@ -0,0 +1,12 @@ +CREATE TABLE chat_synthetic_api_keys ( + user_id uuid NOT NULL PRIMARY KEY REFERENCES users(id) ON DELETE CASCADE, + api_key_id text NOT NULL UNIQUE REFERENCES api_keys(id) ON DELETE CASCADE, + created_at timestamptz NOT NULL DEFAULT now(), + updated_at timestamptz NOT NULL DEFAULT now() +); + +ALTER TABLE chat_messages +DROP CONSTRAINT chat_messages_api_key_id_fkey; + +ALTER TABLE chat_queued_messages +DROP CONSTRAINT chat_queued_messages_api_key_id_fkey; diff --git a/coderd/database/migrations/migrate_test.go b/coderd/database/migrations/migrate_test.go index ae5b2edc489..2e5be77104e 100644 --- a/coderd/database/migrations/migrate_test.go +++ b/coderd/database/migrations/migrate_test.go @@ -1658,6 +1658,65 @@ func TestMigration000542ChatReasoningEffortBackfill(t *testing.T) { require.Equal(t, sql.NullString{}, got["bedrock:anthropic.invalid-effort"]) } +func TestMigration000546ChatHistoryAPIKeyConstraints(t *testing.T) { + t.Parallel() + + const priorMigrationVersion = 545 + + sqlDB := testSQLDB(t) + next, err := migrations.Stepper(sqlDB) + require.NoError(t, err) + for { + version, more, err := next() + require.NoError(t, err) + if !more || version == priorMigrationVersion { + break + } + } + + ctx := testutil.Context(t, testutil.WaitSuperLong) + constraintNames := []string{ + "chat_messages_api_key_id_fkey", + "chat_queued_messages_api_key_id_fkey", + } + assertConstraintCount := func(t *testing.T, want int) { + t.Helper() + for _, name := range constraintNames { + var got int + err := sqlDB.QueryRowContext(ctx, ` + SELECT COUNT(*) + FROM pg_constraint + WHERE conname = $1 + `, name).Scan(&got) + require.NoError(t, err) + require.Equal(t, want, got, name) + } + } + + upSQL, err := os.ReadFile("000546_chat_synthetic_api_keys.up.sql") + require.NoError(t, err) + _, err = sqlDB.ExecContext(ctx, string(upSQL)) + require.NoError(t, err) + assertConstraintCount(t, 0) + + downSQL, err := os.ReadFile("000546_chat_synthetic_api_keys.down.sql") + require.NoError(t, err) + _, err = sqlDB.ExecContext(ctx, string(downSQL)) + require.NoError(t, err) + assertConstraintCount(t, 1) + + for _, name := range constraintNames { + var count int + err := sqlDB.QueryRowContext(ctx, ` + SELECT COUNT(*) + FROM pg_constraint + WHERE conname = $1 AND confdeltype = 'n' + `, name).Scan(&count) + require.NoError(t, err) + require.Equal(t, 1, count, name) + } +} + func TestMigration000498SoftDeleteStaleWorkspaceAgents(t *testing.T) { t.Parallel() diff --git a/coderd/database/migrations/testdata/fixtures/000546_chat_synthetic_api_keys.up.sql b/coderd/database/migrations/testdata/fixtures/000546_chat_synthetic_api_keys.up.sql new file mode 100644 index 00000000000..e6d0257e8ee --- /dev/null +++ b/coderd/database/migrations/testdata/fixtures/000546_chat_synthetic_api_keys.up.sql @@ -0,0 +1,5 @@ +INSERT INTO chat_synthetic_api_keys (user_id, api_key_id) +VALUES ( + '30095c71-380b-457a-8995-97b8ee6e5307', + 'fixture-api-key' +); diff --git a/coderd/database/models.go b/coderd/database/models.go index 73c11d14680..a0e88e5bb8b 100644 --- a/coderd/database/models.go +++ b/coderd/database/models.go @@ -5158,6 +5158,13 @@ type ChatQueuedMessage struct { ReasoningEffort NullChatReasoningEffort `db:"reasoning_effort" json:"reasoning_effort"` } +type ChatSyntheticApiKey struct { + UserID uuid.UUID `db:"user_id" json:"user_id"` + APIKeyID string `db:"api_key_id" json:"api_key_id"` + CreatedAt time.Time `db:"created_at" json:"created_at"` + UpdatedAt time.Time `db:"updated_at" json:"updated_at"` +} + type ChatTable struct { ID uuid.UUID `db:"id" json:"id"` OwnerID uuid.UUID `db:"owner_id" json:"owner_id"` diff --git a/coderd/database/querier.go b/coderd/database/querier.go index cd0a2bfa7fb..6cd2625d4a1 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -471,6 +471,7 @@ type sqlcQuerier interface { // A value of 0 disables chat purging entirely. GetChatRetentionDays(ctx context.Context) (int32, error) GetChatStreamSyncRows(ctx context.Context, ids []uuid.UUID) ([]GetChatStreamSyncRowsRow, error) + GetChatSyntheticAPIKeyByUserID(ctx context.Context, userID uuid.UUID) (ChatSyntheticApiKey, error) GetChatSystemPrompt(ctx context.Context) (string, error) // GetChatSystemPromptConfig returns both chat system prompt settings in a // single read to avoid torn reads between separate site-config lookups. @@ -857,6 +858,7 @@ type sqlcQuerier interface { GetUserChatSpendInPeriod(ctx context.Context, arg GetUserChatSpendInPeriodParams) (int64, error) GetUserCodeDiffDisplayMode(ctx context.Context, userID uuid.UUID) (string, error) GetUserCount(ctx context.Context, includeSystem bool) (int64, error) + GetUserForChatSyntheticAPIKeyByID(ctx context.Context, id uuid.UUID) (User, error) // Returns the minimum (most restrictive) group limit for a user. // Returns -1 if no group limits match the specified scope. // When organization_id is NULL, groups across all organizations are @@ -1063,6 +1065,7 @@ type sqlcQuerier interface { // sequence) and an explicit created_by reference. Use this when the // queued-message creator differs from the chat owner. InsertChatQueuedMessageWithCreator(ctx context.Context, arg InsertChatQueuedMessageWithCreatorParams) (ChatQueuedMessage, error) + InsertChatSyntheticAPIKey(ctx context.Context, arg InsertChatSyntheticAPIKeyParams) (int64, error) InsertCryptoKey(ctx context.Context, arg InsertCryptoKeyParams) (CryptoKey, error) InsertCustomRole(ctx context.Context, arg InsertCustomRoleParams) (CustomRole, error) InsertDBCryptKey(ctx context.Context, arg InsertDBCryptKeyParams) error @@ -1399,6 +1402,7 @@ type sqlcQuerier interface { // assigned by trigger from the current snapshot_version. UpdateChatRetryState(ctx context.Context, arg UpdateChatRetryStateParams) (Chat, error) UpdateChatStatus(ctx context.Context, arg UpdateChatStatusParams) (Chat, error) + UpdateChatSyntheticAPIKey(ctx context.Context, arg UpdateChatSyntheticAPIKeyParams) (int64, error) UpdateChatTitleByID(ctx context.Context, arg UpdateChatTitleByIDParams) (Chat, error) UpdateChatWorkspaceBinding(ctx context.Context, arg UpdateChatWorkspaceBindingParams) (Chat, error) UpdateCryptoKeyDeletesAt(ctx context.Context, arg UpdateCryptoKeyDeletesAtParams) (CryptoKey, error) diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index a65852cf8de..de9ddbce507 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -8202,6 +8202,24 @@ func (q *sqlQuerier) GetChatStreamSyncRows(ctx context.Context, ids []uuid.UUID) return items, nil } +const getChatSyntheticAPIKeyByUserID = `-- name: GetChatSyntheticAPIKeyByUserID :one +SELECT user_id, api_key_id, created_at, updated_at +FROM chat_synthetic_api_keys +WHERE user_id = $1::uuid +` + +func (q *sqlQuerier) GetChatSyntheticAPIKeyByUserID(ctx context.Context, userID uuid.UUID) (ChatSyntheticApiKey, error) { + row := q.db.QueryRowContext(ctx, getChatSyntheticAPIKeyByUserID, userID) + var i ChatSyntheticApiKey + err := row.Scan( + &i.UserID, + &i.APIKeyID, + &i.CreatedAt, + &i.UpdatedAt, + ) + return i, err +} + const getChatUsageLimitConfig = `-- name: GetChatUsageLimitConfig :one SELECT id, singleton, enabled, default_limit_micros, period, created_at, updated_at FROM chat_usage_limit_config WHERE singleton = TRUE LIMIT 1 ` @@ -9985,6 +10003,25 @@ func (q *sqlQuerier) InsertChatQueuedMessageWithCreator(ctx context.Context, arg return i, err } +const insertChatSyntheticAPIKey = `-- name: InsertChatSyntheticAPIKey :execrows +INSERT INTO chat_synthetic_api_keys (user_id, api_key_id) +VALUES ($1::uuid, $2::text) +ON CONFLICT (user_id) DO NOTHING +` + +type InsertChatSyntheticAPIKeyParams struct { + UserID uuid.UUID `db:"user_id" json:"user_id"` + APIKeyID string `db:"api_key_id" json:"api_key_id"` +} + +func (q *sqlQuerier) InsertChatSyntheticAPIKey(ctx context.Context, arg InsertChatSyntheticAPIKeyParams) (int64, error) { + result, err := q.db.ExecContext(ctx, insertChatSyntheticAPIKey, arg.UserID, arg.APIKeyID) + if err != nil { + return 0, err + } + return result.RowsAffected() +} + const isChatHeartbeatStale = `-- name: IsChatHeartbeatStale :one SELECT NOT EXISTS ( SELECT 1 FROM chat_heartbeats @@ -12170,6 +12207,28 @@ func (q *sqlQuerier) UpdateChatStatus(ctx context.Context, arg UpdateChatStatusP return i, err } +const updateChatSyntheticAPIKey = `-- name: UpdateChatSyntheticAPIKey :execrows +UPDATE chat_synthetic_api_keys +SET api_key_id = $1::text, + updated_at = NOW() +WHERE user_id = $2::uuid + AND api_key_id = $3::text +` + +type UpdateChatSyntheticAPIKeyParams struct { + NewApiKeyID string `db:"new_api_key_id" json:"new_api_key_id"` + UserID uuid.UUID `db:"user_id" json:"user_id"` + OldApiKeyID string `db:"old_api_key_id" json:"old_api_key_id"` +} + +func (q *sqlQuerier) UpdateChatSyntheticAPIKey(ctx context.Context, arg UpdateChatSyntheticAPIKeyParams) (int64, error) { + result, err := q.db.ExecContext(ctx, updateChatSyntheticAPIKey, arg.NewApiKeyID, arg.UserID, arg.OldApiKeyID) + if err != nil { + return 0, err + } + return result.RowsAffected() +} + const updateChatTitleByID = `-- name: UpdateChatTitleByID :one WITH updated_chat AS ( UPDATE @@ -29857,6 +29916,40 @@ func (q *sqlQuerier) GetUserCount(ctx context.Context, includeSystem bool) (int6 return count, err } +const getUserForChatSyntheticAPIKeyByID = `-- name: GetUserForChatSyntheticAPIKeyByID :one +SELECT id, email, username, hashed_password, created_at, updated_at, status, rbac_roles, login_type, avatar_url, deleted, last_seen_at, quiet_hours_schedule, name, github_com_user_id, hashed_one_time_passcode, one_time_passcode_expires_at, is_system, is_service_account, chat_spend_limit_micros +FROM users +WHERE id = $1::uuid +` + +func (q *sqlQuerier) GetUserForChatSyntheticAPIKeyByID(ctx context.Context, id uuid.UUID) (User, error) { + row := q.db.QueryRowContext(ctx, getUserForChatSyntheticAPIKeyByID, id) + var i User + err := row.Scan( + &i.ID, + &i.Email, + &i.Username, + &i.HashedPassword, + &i.CreatedAt, + &i.UpdatedAt, + &i.Status, + &i.RBACRoles, + &i.LoginType, + &i.AvatarURL, + &i.Deleted, + &i.LastSeenAt, + &i.QuietHoursSchedule, + &i.Name, + &i.GithubComUserID, + &i.HashedOneTimePasscode, + &i.OneTimePasscodeExpiresAt, + &i.IsSystem, + &i.IsServiceAccount, + &i.ChatSpendLimitMicros, + ) + return i, err +} + const getUserShellToolDisplayMode = `-- name: GetUserShellToolDisplayMode :one SELECT value AS shell_tool_display_mode diff --git a/coderd/database/queries/chats.sql b/coderd/database/queries/chats.sql index 51bd02e34bd..31b0050708d 100644 --- a/coderd/database/queries/chats.sql +++ b/coderd/database/queries/chats.sql @@ -3008,3 +3008,20 @@ LEFT JOIN to_archive t ON t.id = a.id -- created_at ASC flows through to dbpurge's digest truncation; see -- buildDigestData in dbpurge.go for the tradeoff rationale. ORDER BY (a.root_chat_id IS NULL) DESC, a.owner_id ASC, a.created_at ASC, a.id ASC; + +-- name: GetChatSyntheticAPIKeyByUserID :one +SELECT * +FROM chat_synthetic_api_keys +WHERE user_id = @user_id::uuid; + +-- name: InsertChatSyntheticAPIKey :execrows +INSERT INTO chat_synthetic_api_keys (user_id, api_key_id) +VALUES (@user_id::uuid, @api_key_id::text) +ON CONFLICT (user_id) DO NOTHING; + +-- name: UpdateChatSyntheticAPIKey :execrows +UPDATE chat_synthetic_api_keys +SET api_key_id = @new_api_key_id::text, + updated_at = NOW() +WHERE user_id = @user_id::uuid + AND api_key_id = @old_api_key_id::text; diff --git a/coderd/database/queries/users.sql b/coderd/database/queries/users.sql index 92dc26a4d7d..3c79b405225 100644 --- a/coderd/database/queries/users.sql +++ b/coderd/database/queries/users.sql @@ -689,3 +689,8 @@ SET WHERE id = $1 ; + +-- name: GetUserForChatSyntheticAPIKeyByID :one +SELECT * +FROM users +WHERE id = @id::uuid; diff --git a/coderd/database/unique_constraint.go b/coderd/database/unique_constraint.go index 4b1a4376f2d..aa954b5cf0c 100644 --- a/coderd/database/unique_constraint.go +++ b/coderd/database/unique_constraint.go @@ -32,6 +32,8 @@ const ( UniqueChatMessagesPkey UniqueConstraint = "chat_messages_pkey" // ALTER TABLE ONLY chat_messages ADD CONSTRAINT chat_messages_pkey PRIMARY KEY (id); UniqueChatModelConfigsPkey UniqueConstraint = "chat_model_configs_pkey" // ALTER TABLE ONLY chat_model_configs ADD CONSTRAINT chat_model_configs_pkey PRIMARY KEY (id); UniqueChatQueuedMessagesPkey UniqueConstraint = "chat_queued_messages_pkey" // ALTER TABLE ONLY chat_queued_messages ADD CONSTRAINT chat_queued_messages_pkey PRIMARY KEY (id); + UniqueChatSyntheticAPIKeysAPIKeyIDKey UniqueConstraint = "chat_synthetic_api_keys_api_key_id_key" // ALTER TABLE ONLY chat_synthetic_api_keys ADD CONSTRAINT chat_synthetic_api_keys_api_key_id_key UNIQUE (api_key_id); + UniqueChatSyntheticAPIKeysPkey UniqueConstraint = "chat_synthetic_api_keys_pkey" // ALTER TABLE ONLY chat_synthetic_api_keys ADD CONSTRAINT chat_synthetic_api_keys_pkey PRIMARY KEY (user_id); UniqueChatUsageLimitConfigPkey UniqueConstraint = "chat_usage_limit_config_pkey" // ALTER TABLE ONLY chat_usage_limit_config ADD CONSTRAINT chat_usage_limit_config_pkey PRIMARY KEY (id); UniqueChatUsageLimitConfigSingletonKey UniqueConstraint = "chat_usage_limit_config_singleton_key" // ALTER TABLE ONLY chat_usage_limit_config ADD CONSTRAINT chat_usage_limit_config_singleton_key UNIQUE (singleton); UniqueChatsPkey UniqueConstraint = "chats_pkey" // ALTER TABLE ONLY chats ADD CONSTRAINT chats_pkey PRIMARY KEY (id); diff --git a/coderd/exp_chats.go b/coderd/exp_chats.go index 46323846834..39f6673ab30 100644 --- a/coderd/exp_chats.go +++ b/coderd/exp_chats.go @@ -27,7 +27,6 @@ import ( "cdr.dev/slog/v3" "github.com/coder/coder/v2/agent/agentssh" - "github.com/coder/coder/v2/coderd/aibridge" "github.com/coder/coder/v2/coderd/audit" "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/database/db2sdk" @@ -1420,7 +1419,6 @@ func (api *API) postChats(rw http.ResponseWriter, r *http.Request) { ClientType: clientType, SystemPrompt: req.SystemPrompt, InitialUserContent: contentBlocks, - APIKeyID: apiKey.ID, MCPServerIDs: mcpServerIDs, Labels: labels, DynamicTools: dynamicToolsJSON, @@ -3351,7 +3349,6 @@ func (api *API) postChatMessages(rw http.ResponseWriter, r *http.Request) { Content: contentBlocks, ModelConfigID: modelConfigID, ReasoningEffort: reasoningEffort, - APIKeyID: apiKey.ID, BusyBehavior: busyBehavior, PlanMode: sendPlanMode, MCPServerIDs: req.MCPServerIDs, @@ -3513,7 +3510,6 @@ func (api *API) patchChatMessage(rw http.ResponseWriter, r *http.Request) { CreatedBy: apiKey.UserID, EditedMessageID: messageID, Content: contentBlocks, - APIKeyID: apiKey.ID, ModelConfigID: editModelConfigID, ReasoningEffort: editReasoningEffort, }) @@ -4026,7 +4022,6 @@ func (api *API) regenerateChatTitle(rw http.ResponseWriter, r *http.Request) { return } - ctx = aibridge.WithDelegatedAPIKeyID(ctx, apiKey.ID) updatedChat, err := api.chatDaemon.RegenerateChatTitle(ctx, chat) if err != nil { if errors.Is(err, chatd.ErrNoDefaultChatModelConfig) { @@ -4076,7 +4071,6 @@ func (api *API) proposeChatTitle(rw http.ResponseWriter, r *http.Request) { return } - ctx = aibridge.WithDelegatedAPIKeyID(ctx, apiKey.ID) title, err := api.chatDaemon.ProposeChatTitle(ctx, chat) if err != nil { if errors.Is(err, chatd.ErrNoDefaultChatModelConfig) { diff --git a/coderd/exp_chats_test.go b/coderd/exp_chats_test.go index 209e51d1e71..d48ab69db2c 100644 --- a/coderd/exp_chats_test.go +++ b/coderd/exp_chats_test.go @@ -9540,7 +9540,7 @@ func TestProposeChatTitle(t *testing.T) { }) } -func TestManualTitleEndpointsPassCallerAPIKeyToAIGateway(t *testing.T) { +func TestManualTitleEndpointsPassOwnerSyntheticAPIKeyToAIGateway(t *testing.T) { t.Parallel() for _, tt := range []struct { @@ -9574,7 +9574,6 @@ func TestManualTitleEndpointsPassCallerAPIKeyToAIGateway(t *testing.T) { }) firstUser := coderdtest.CreateFirstUser(t, client.Client) modelConfig := createAdditionalChatModelConfig(t, client, "openai", "gpt-4.1") - wantAPIKeyID := strings.Split(client.SessionToken(), "-")[0] wantTitle := "Fallback title" seenAPIKeyID := make(chan string, 1) stub := &stubTransportFactory{ @@ -9607,7 +9606,6 @@ func TestManualTitleEndpointsPassCallerAPIKeyToAIGateway(t *testing.T) { _ = dbgen.ChatMessage(t, db, database.ChatMessage{ ChatID: chat.ID, CreatedBy: uuid.NullUUID{UUID: firstUser.UserID, Valid: true}, - APIKeyID: sql.NullString{String: wantAPIKeyID, Valid: true}, ModelConfigID: uuid.NullUUID{UUID: modelConfig.ID, Valid: true}, Role: database.ChatMessageRoleUser, Visibility: database.ChatMessageVisibilityBoth, @@ -9615,7 +9613,9 @@ func TestManualTitleEndpointsPassCallerAPIKeyToAIGateway(t *testing.T) { }) require.NoError(t, tt.call(ctx, client, chat.ID)) - require.Equal(t, wantAPIKeyID, testutil.RequireReceive(ctx, t, seenAPIKeyID)) + mapping, err := db.GetChatSyntheticAPIKeyByUserID(dbauthz.AsSystemRestricted(ctx), firstUser.UserID) + require.NoError(t, err) + require.Equal(t, mapping.APIKeyID, testutil.RequireReceive(ctx, t, seenAPIKeyID)) }) } } diff --git a/coderd/rbac/authz.go b/coderd/rbac/authz.go index 0be01384941..58086159d9d 100644 --- a/coderd/rbac/authz.go +++ b/coderd/rbac/authz.go @@ -75,7 +75,8 @@ const ( SubjectTypeSystemReadProvisionerDaemons SubjectType = "system_read_provisioner_daemons" SubjectTypeSystemRestricted SubjectType = "system_restricted" SubjectTypeSystemOAuth SubjectType = "system_oauth" - SubjectTypeAPIKeyRevoker SubjectType = "api_key_revoker" // #nosec G101, not a credential. + SubjectTypeAPIKeyRevoker SubjectType = "api_key_revoker" // #nosec G101, not a credential. + SubjectTypeChatdKeyMinter SubjectType = "chatd_key_minter" // #nosec G101, not a credential. SubjectTypeNotifier SubjectType = "notifier" SubjectTypeSubAgentAPI SubjectType = "sub_agent_api" SubjectTypeFileReader SubjectType = "file_reader" diff --git a/coderd/x/chatd/ARCHITECTURE.md b/coderd/x/chatd/ARCHITECTURE.md index d532d1b8448..119d154412c 100644 --- a/coderd/x/chatd/ARCHITECTURE.md +++ b/coderd/x/chatd/ARCHITECTURE.md @@ -7,6 +7,16 @@ Chatd has 4 main pieces: - **chat worker**: lives inside every coderd replica. It acquires chats, calls the LLM API, executes tools, handles interrupts and tool-result waits, and commits completed outcomes through the core state machine. - **stream loop**: powers `GET /api/experimental/chats/{chat}/stream`, the WebSocket endpoint that the UI uses to consume a live chat. It combines two kinds of data: messages committed to the database and streaming message parts emitted by the chat worker. It receives notifications over pubsub whenever the chat state is updated, fetches messages from the database, and connects to the coderd replica that currently owns the chat to relay the streaming message parts to the client. +# Gateway attribution keys + +Chatd attributes AI Gateway requests with a synthetic API key owned by the chat owner. The mapping is stored in `chat_synthetic_api_keys`, with one active key per user. All chatd AI Gateway attribution resolves the key from `chats.owner_id`. Callers do not provide the key ID. + +Synthetic keys expire after 30 days and renew when less than 24 hours remain. Renewal replaces the mapping but does not delete the previous key, because an in-flight request may still use it. Concurrent mints use conditional insert/update operations, and losing keys are deleted. The generated token is discarded, so the stored key cannot be used as a bearer credential. + +The legacy `api_key_id` columns on messages and queued messages are still stamped with the synthetic key for rolling deployment compatibility, but they no longer have foreign keys to `api_keys`. They are not the source of gateway routing. Stale IDs are harmless because chatd resolves attribution from `chats.owner_id`. + +Deleting a synthetic key removes its mapping without updating chat messages, queued messages, or their version fields. Deleting all user API keys or resetting a password therefore causes chatd to mint a replacement on the next request without mutating history. User suspension and deletion still block delegated gateway authorization. + # Core state machine The core state machine describes how a chat's execution state in the database can change over time. A fundamental component of the state machine is the set of valid **states** it can be in. We will consider 2 kinds of states: **execution states** and **ownership states**. These states let us describe what the runtime components of chatd can do with a chat at a given point in time. diff --git a/coderd/x/chatd/chatd.go b/coderd/x/chatd/chatd.go index fa2012b2ade..c110b016662 100644 --- a/coderd/x/chatd/chatd.go +++ b/coderd/x/chatd/chatd.go @@ -1174,7 +1174,6 @@ type CreateOptions struct { ClientType database.ChatClientType SystemPrompt string InitialUserContent []codersdk.ChatMessagePart - APIKeyID string MCPServerIDs []uuid.UUID Labels database.StringMap DynamicTools json.RawMessage @@ -1200,7 +1199,6 @@ type SendMessageOptions struct { Content []codersdk.ChatMessagePart ModelConfigID uuid.UUID ReasoningEffort *string - APIKeyID string BusyBehavior SendMessageBusyBehavior PlanMode *database.NullChatPlanMode MCPServerIDs *[]uuid.UUID @@ -1220,7 +1218,6 @@ type EditMessageOptions struct { CreatedBy uuid.UUID EditedMessageID int64 Content []codersdk.ChatMessagePart - APIKeyID string // ModelConfigID, when non-zero, overrides the model used for // the replacement user message. When set to uuid.Nil the // original message's model is preserved. @@ -1249,13 +1246,6 @@ type PromoteQueuedResult struct { PromotedMessage database.ChatMessage } -func validateChatUserMessageAPIKeyID(apiKeyID string) error { - if apiKeyID == "" { - return xerrors.New("api_key_id is required for user chat messages") - } - return nil -} - // CreateChat creates a chat with its initial history through // chatstate.CreateChat. The new chat starts in `running` status per // the chat execution state model. Ownership hints wake chat workers. @@ -1272,9 +1262,6 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C if len(opts.InitialUserContent) == 0 { return database.Chat{}, xerrors.New("initial user content is required") } - if err := validateChatUserMessageAPIKeyID(opts.APIKeyID); err != nil { - return database.Chat{}, err - } // Ensure MCPServerIDs is non-nil so pq.Array produces '{}' // instead of SQL NULL, which violates the NOT NULL column // constraint. @@ -1298,6 +1285,11 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C return database.Chat{}, limitErr } + apiKeyID, err := p.ensureSyntheticAPIKeyID(ctx, opts.OwnerID) + if err != nil { + return database.Chat{}, xerrors.Errorf("ensure synthetic API key: %w", err) + } + labelsJSON, err := json.Marshal(opts.Labels) if err != nil { return database.Chat{}, xerrors.Errorf("marshal labels: %w", err) @@ -1339,7 +1331,7 @@ func (p *Server) CreateChat(ctx context.Context, opts CreateOptions) (database.C initialMessages = append(initialMessages, systemMessage(userPromptContent, opts.ModelConfigID)) } initialMessages = append(initialMessages, systemMessage(workspaceAwarenessContent, opts.ModelConfigID)) - initialMessages = append(initialMessages, userMessageWithAPIKeyID(userContent, opts.ModelConfigID, opts.OwnerID, opts.APIKeyID, opts.ReasoningEffort)) + initialMessages = append(initialMessages, userMessageWithAPIKeyID(userContent, opts.ModelConfigID, opts.OwnerID, apiKeyID, opts.ReasoningEffort)) result, err := chatstate.CreateChat(ctx, p.db, p.pubsub, chatstate.CreateChatInput{ OrganizationID: opts.OrganizationID, @@ -1400,9 +1392,6 @@ func (p *Server) SendMessage( if len(opts.Content) == 0 { return SendMessageResult{}, xerrors.New("content is required") } - if err := validateChatUserMessageAPIKeyID(opts.APIKeyID); err != nil { - return SendMessageResult{}, err - } busyBehavior := opts.BusyBehavior if busyBehavior == "" { @@ -1422,6 +1411,15 @@ func (p *Server) SendMessage( requestedPlanMode := opts.PlanMode requestedMCPServerIDs := opts.MCPServerIDs + chat, err := p.db.GetChatByID(ctx, opts.ChatID) + if err != nil { + return SendMessageResult{}, xerrors.Errorf("load chat: %w", err) + } + apiKeyID, err := p.ensureSyntheticAPIKeyID(ctx, chat.OwnerID) + if err != nil { + return SendMessageResult{}, xerrors.Errorf("ensure synthetic API key: %w", err) + } + var result SendMessageResult machine := p.newChatMachine(opts.ChatID) updateErr := machine.Update(ctx, func(tx *chatstate.Tx, store database.Store) error { @@ -1486,7 +1484,7 @@ func (p *Server) SendMessage( // Queue capacity is enforced inside tx.SendMessage; this // wrapper only propagates the typed error. sendResult, err := tx.SendMessage(chatstate.SendMessageInput{ - Message: userMessageWithAPIKeyID(content, modelConfigID, messageCreatedBy, opts.APIKeyID, opts.ReasoningEffort), + Message: userMessageWithAPIKeyID(content, modelConfigID, messageCreatedBy, apiKeyID, opts.ReasoningEffort), BusyBehavior: busyBehaviorToChatState(busyBehavior), }) if err != nil { @@ -1627,14 +1625,19 @@ func (p *Server) EditMessage( if len(opts.Content) == 0 { return EditMessageResult{}, xerrors.New("content is required") } - if err := validateChatUserMessageAPIKeyID(opts.APIKeyID); err != nil { - return EditMessageResult{}, err - } content, err := chatprompt.MarshalParts(opts.Content) if err != nil { return EditMessageResult{}, xerrors.Errorf("marshal message content: %w", err) } + chat, err := p.db.GetChatByID(ctx, opts.ChatID) + if err != nil { + return EditMessageResult{}, xerrors.Errorf("load chat: %w", err) + } + apiKeyID, err := p.ensureSyntheticAPIKeyID(ctx, chat.OwnerID) + if err != nil { + return EditMessageResult{}, xerrors.Errorf("ensure synthetic API key: %w", err) + } var ( result EditMessageResult @@ -1705,7 +1708,7 @@ func (p *Server) EditMessage( Content: content, ModelConfigIDOverride: modelOverride, ReasoningEffortOverride: reasoningEffortOverride, - APIKeyID: sql.NullString{String: opts.APIKeyID, Valid: opts.APIKeyID != ""}, + APIKeyID: sql.NullString{String: apiKeyID, Valid: true}, }) if err != nil { if errors.Is(err, chatstate.ErrEditedMessageNotUser) { @@ -2248,7 +2251,6 @@ func (p *Server) ProposeChatTitle( // generateManualTitleCandidate generates a title candidate from the chat's // visible messages. It returns "" when the chat has no messages to summarize. // Endpoint-specific commit paths decide whether to persist the title. -// The context may carry the caller's delegated API key for manual title routes. func (p *Server) generateManualTitleCandidate( ctx context.Context, store database.Store, @@ -2288,14 +2290,11 @@ func (p *Server) generateManualTitleCandidate( if err != nil { return "", xerrors.Errorf("get pasted-text attachments for manual title: %w", err) } - modelOpts := modelBuildOptionsFromMessages(messages) - // Manual title routes can run over messages that lack API key attribution. - // Fall back to the authenticated caller's delegated key for AI Gateway routing. - if modelOpts.ActiveAPIKeyID == "" { - if apiKeyID, ok := aibridge.DelegatedAPIKeyIDFromContext(ctx); ok { - modelOpts.ActiveAPIKeyID = apiKeyID - } + apiKeyID, err := p.ensureSyntheticAPIKeyID(ctx, chat.OwnerID) + if err != nil { + return "", xerrors.Errorf("ensure synthetic API key: %w", err) } + modelOpts := modelBuildOptions{ActiveAPIKeyID: apiKeyID} model, modelConfig, err := p.resolveManualTitleModel(ctx, store, chat, modelOpts) if err != nil { @@ -3286,29 +3285,6 @@ type runChatResult struct { HistoryTipMessageID int64 } -func activeTurnAPIKeyIDFromMessages(messages []database.ChatMessage) (string, bool) { - for i := len(messages) - 1; i >= 0; i-- { - message := messages[i] - if message.Role != database.ChatMessageRoleUser { - continue - } - if !isUserVisibleChatMessage(message) && - !(message.Visibility == database.ChatMessageVisibilityModel && message.Compressed) { - continue - } - if !message.APIKeyID.Valid || message.APIKeyID.String == "" { - return "", false - } - return message.APIKeyID.String, true - } - return "", false -} - -func isUserVisibleChatMessage(message database.ChatMessage) bool { - return message.Visibility == database.ChatMessageVisibilityBoth || - message.Visibility == database.ChatMessageVisibilityUser -} - func allToolNames(allTools []fantasy.AgentTool) []string { toolNames := make([]string, 0, len(allTools)) for _, tool := range allTools { diff --git a/coderd/x/chatd/chatd_helpers_test.go b/coderd/x/chatd/chatd_helpers_test.go index b51d7c55c6d..a271016ab82 100644 --- a/coderd/x/chatd/chatd_helpers_test.go +++ b/coderd/x/chatd/chatd_helpers_test.go @@ -57,7 +57,6 @@ func filterAnthropicStreamingRequests(requests []chattest.AnthropicRequest) []ch func seedAnthropicChatDependencies(t *testing.T, db database.Store, baseURL string) (database.User, database.Organization, database.ChatModelConfig) { t.Helper() user := dbgen.User(t, db, database.User{}) - _ = testAPIKeyID(t, db, user.ID) org := dbgen.Organization(t, db, database.Organization{}) dbgen.OrganizationMember(t, db, database.OrganizationMember{UserID: user.ID, OrganizationID: org.ID}) provider := dbgen.AIProvider(t, db, database.AIProvider{Type: database.AIProviderTypeAnthropic}, func(params *database.InsertAIProviderParams) { @@ -157,7 +156,6 @@ func createChatThroughServer( chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: orgID, OwnerID: userID, - APIKeyID: testAPIKeyID(t, db, userID), Title: "test chat", InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText(text)}, ModelConfigID: modelID, diff --git a/coderd/x/chatd/chatd_internal_test.go b/coderd/x/chatd/chatd_internal_test.go index 916df0f80fb..dacc9116b9b 100644 --- a/coderd/x/chatd/chatd_internal_test.go +++ b/coderd/x/chatd/chatd_internal_test.go @@ -858,6 +858,8 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts(t *testing.T) { db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes() db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), gomock.Any()).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes() db.EXPECT().GetChatUsageLimitConfig(gomock.Any()).Return(database.ChatUsageLimitConfig{}, sql.ErrNoRows) + db.EXPECT().GetChatSyntheticAPIKeyByUserID(gomock.Any(), ownerID).Return(database.ChatSyntheticApiKey{UserID: ownerID, APIKeyID: activeAPIKeyID}, nil) + db.EXPECT().GetAPIKeyByID(gomock.Any(), activeAPIKeyID).Return(database.APIKey{ID: activeAPIKeyID, UserID: ownerID, ExpiresAt: time.Now().Add(48 * time.Hour)}, nil) db.EXPECT().GetChatMessagesByChatIDAscPaginated( gomock.Any(), database.GetChatMessagesByChatIDAscPaginatedParams{ @@ -1009,6 +1011,8 @@ func TestRegenerateChatTitle_SkipsPersistWhenTitleChangedConcurrently(t *testing db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes() db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), gomock.Any()).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes() db.EXPECT().GetChatUsageLimitConfig(gomock.Any()).Return(database.ChatUsageLimitConfig{}, sql.ErrNoRows) + db.EXPECT().GetChatSyntheticAPIKeyByUserID(gomock.Any(), ownerID).Return(database.ChatSyntheticApiKey{UserID: ownerID, APIKeyID: activeAPIKeyID}, nil) + db.EXPECT().GetAPIKeyByID(gomock.Any(), activeAPIKeyID).Return(database.APIKey{ID: activeAPIKeyID, UserID: ownerID, ExpiresAt: time.Now().Add(48 * time.Hour)}, nil) db.EXPECT().GetChatMessagesByChatIDAscPaginated( gomock.Any(), database.GetChatMessagesByChatIDAscPaginatedParams{ diff --git a/coderd/x/chatd/chatd_test.go b/coderd/x/chatd/chatd_test.go index 53ca0b9078b..d843105cbe1 100644 --- a/coderd/x/chatd/chatd_test.go +++ b/coderd/x/chatd/chatd_test.go @@ -903,7 +903,6 @@ func TestRootExploreChatStaysBuiltinOnlyAtRuntime(t *testing.T) { exploreChat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "root-explore-builtin-only", ModelConfigID: model.ID, ChatMode: database.NullChatMode{ @@ -992,7 +991,6 @@ func TestRootExploreChatExcludesWebSearchProviderToolAtRuntime(t *testing.T) { exploreChat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "root-explore-no-provider-web-search", ModelConfigID: webSearchModel.ID, ChatMode: database.NullChatMode{ @@ -1119,7 +1117,6 @@ func TestExploreChatSendMessageCannotMutateMCPSnapshot(t *testing.T) { rootChat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "runtime-parent", ModelConfigID: model.ID, MCPServerIDs: []uuid.UUID{parentConfig.ID}, @@ -1162,7 +1159,6 @@ func TestExploreChatSendMessageCannotMutateMCPSnapshot(t *testing.T) { _, err = server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: exploreChat.ID, CreatedBy: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("inspect the codebase again")}, MCPServerIDs: &updatedMCPServerIDs, }) @@ -1325,7 +1321,6 @@ func TestPlanModeRootChatAllowsApprovedExternalMCPTools(t *testing.T) { planChat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "plan-mode-root-mcp-visibility", ModelConfigID: model.ID, WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, @@ -1345,7 +1340,6 @@ func TestPlanModeRootChatAllowsApprovedExternalMCPTools(t *testing.T) { askChat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "ask-mode-root-mcp-visibility", ModelConfigID: model.ID, WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, @@ -1573,7 +1567,6 @@ func TestUpdateChatHeartbeatsRequiresOwnership(t *testing.T) { chat, err := replica.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "heartbeat-ownership", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -1610,7 +1603,7 @@ func TestUpdateChatHeartbeatsRequiresOwnership(t *testing.T) { require.Equal(t, chat.ID, ids[0]) } -func TestCreateChatPersistsAPIKeyIDOnInitialUserMessage(t *testing.T) { +func TestCreateChatPersistsSyntheticAPIKeyIDOnInitialUserMessage(t *testing.T) { t.Parallel() db, ps := dbtestutil.NewDB(t) @@ -1618,15 +1611,12 @@ func TestCreateChatPersistsAPIKeyIDOnInitialUserMessage(t *testing.T) { ctx := testutil.Context(t, testutil.WaitLong) user, org, model := seedChatDependencies(t, db) - apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID}) - chat, err := replica.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - Title: "create-chat-api-key-id", + Title: "create-chat-synthetic-api-key-id", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, - APIKeyID: apiKey.ID, }) require.NoError(t, err) @@ -1638,10 +1628,12 @@ func TestCreateChatPersistsAPIKeyIDOnInitialUserMessage(t *testing.T) { require.Len(t, messages, 1) require.Equal(t, database.ChatMessageRoleUser, messages[0].Role) require.True(t, messages[0].APIKeyID.Valid) - require.Equal(t, apiKey.ID, messages[0].APIKeyID.String) + mapping, err := db.GetChatSyntheticAPIKeyByUserID(ctx, user.ID) + require.NoError(t, err) + require.Equal(t, mapping.APIKeyID, messages[0].APIKeyID.String) } -func TestSendMessagePersistsAPIKeyIDOnUserMessage(t *testing.T) { +func TestSendMessagePersistsSyntheticAPIKeyIDOnUserMessage(t *testing.T) { t.Parallel() db, ps := dbtestutil.NewDB(t) @@ -1649,32 +1641,123 @@ func TestSendMessagePersistsAPIKeyIDOnUserMessage(t *testing.T) { ctx := testutil.Context(t, testutil.WaitLong) user, org, model := seedChatDependencies(t, db) - apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID}) chat := dbgen.Chat(t, db, database.Chat{ OrganizationID: org.ID, OwnerID: user.ID, LastModelConfigID: model.ID, - Title: "send-message-api-key-id", + Title: "send-message-synthetic-api-key-id", }) result, err := replica.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, CreatedBy: user.ID, Content: []codersdk.ChatMessagePart{ - codersdk.ChatMessageText("message with api key id"), + codersdk.ChatMessageText("message with synthetic api key id"), }, - APIKeyID: apiKey.ID, }) require.NoError(t, err) require.False(t, result.Queued) require.True(t, result.Message.APIKeyID.Valid) - require.Equal(t, apiKey.ID, result.Message.APIKeyID.String) + mapping, err := db.GetChatSyntheticAPIKeyByUserID(ctx, user.ID) + require.NoError(t, err) + require.Equal(t, mapping.APIKeyID, result.Message.APIKeyID.String) + + stored, err := db.GetChatMessageByID(ctx, result.Message.ID) + require.NoError(t, err) + require.True(t, stored.APIKeyID.Valid) + require.Equal(t, mapping.APIKeyID, stored.APIKeyID.String) +} + +func TestSendMessagePersistsSyntheticAPIKeyIDOnQueuedUserMessage(t *testing.T) { + t.Parallel() + + db, ps := dbtestutil.NewDB(t) + replica := newTestServer(t, db, ps, uuid.New()) + + ctx := testutil.Context(t, testutil.WaitLong) + user, org, model := seedChatDependencies(t, db) + chat, err := replica.CreateChat(ctx, chatd.CreateOptions{ + OrganizationID: org.ID, + OwnerID: user.ID, + Title: "queue-synthetic-api-key-id", + ModelConfigID: model.ID, + InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, + }) + require.NoError(t, err) + + chat, err = db.UpdateChatStatus(ctx, database.UpdateChatStatusParams{ + ID: chat.ID, + Status: database.ChatStatusRunning, + WorkerID: uuid.NullUUID{UUID: uuid.New(), Valid: true}, + StartedAt: sql.NullTime{Time: time.Now(), Valid: true}, + HeartbeatAt: sql.NullTime{Time: time.Now(), Valid: true}, + }) + require.NoError(t, err) + + result, err := replica.SendMessage(ctx, chatd.SendMessageOptions{ + ChatID: chat.ID, + CreatedBy: user.ID, + Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("queued")}, + BusyBehavior: chatd.SendMessageBusyBehaviorQueue, + }) + require.NoError(t, err) + require.True(t, result.Queued) + require.NotNil(t, result.QueuedMessage) + + mapping, err := db.GetChatSyntheticAPIKeyByUserID(ctx, user.ID) + require.NoError(t, err) + require.True(t, result.QueuedMessage.APIKeyID.Valid) + require.Equal(t, mapping.APIKeyID, result.QueuedMessage.APIKeyID.String) + + queued, err := db.GetChatQueuedMessages(ctx, chat.ID) + require.NoError(t, err) + require.Len(t, queued, 1) + require.True(t, queued[0].APIKeyID.Valid) + require.Equal(t, mapping.APIKeyID, queued[0].APIKeyID.String) +} + +func TestEditMessagePersistsSyntheticAPIKeyIDOnReplacement(t *testing.T) { + t.Parallel() + + db, ps := dbtestutil.NewDB(t) + replica := newTestServer(t, db, ps, uuid.New()) + + ctx := testutil.Context(t, testutil.WaitLong) + user, org, model := seedChatDependencies(t, db) + chat, err := replica.CreateChat(ctx, chatd.CreateOptions{ + OrganizationID: org.ID, + OwnerID: user.ID, + Title: "edit-synthetic-api-key-id", + ModelConfigID: model.ID, + InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("original")}, + }) + require.NoError(t, err) + + messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{ + ChatID: chat.ID, + AfterID: 0, + }) + require.NoError(t, err) + require.Len(t, messages, 1) + + result, err := replica.EditMessage(ctx, chatd.EditMessageOptions{ + ChatID: chat.ID, + EditedMessageID: messages[0].ID, + CreatedBy: user.ID, + Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edited")}, + }) + require.NoError(t, err) + + mapping, err := db.GetChatSyntheticAPIKeyByUserID(ctx, user.ID) + require.NoError(t, err) + require.True(t, result.Message.APIKeyID.Valid) + require.Equal(t, mapping.APIKeyID, result.Message.APIKeyID.String) stored, err := db.GetChatMessageByID(ctx, result.Message.ID) require.NoError(t, err) require.True(t, stored.APIKeyID.Valid) - require.Equal(t, apiKey.ID, stored.APIKeyID.String) + require.Equal(t, mapping.APIKeyID, stored.APIKeyID.String) } func TestSendMessageQueueBehaviorQueuesWhenBusy(t *testing.T) { @@ -1689,7 +1772,6 @@ func TestSendMessageQueueBehaviorQueuesWhenBusy(t *testing.T) { chat, err := replica.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "queue-when-busy", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -1708,7 +1790,6 @@ func TestSendMessageQueueBehaviorQueuesWhenBusy(t *testing.T) { result, err := replica.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("queued")}, BusyBehavior: chatd.SendMessageBusyBehaviorQueue, }) @@ -1767,7 +1848,6 @@ func TestPlanTurnPromptContract(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), OrganizationID: org.ID, Title: "plan-turn-prompt-contract", ModelConfigID: model.ID, @@ -1828,7 +1908,6 @@ func TestSendMessageRejectsInvalidQueuedModelConfigID(t *testing.T) { invalidModelConfigID := uuid.New() _, err := replica.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("queued")}, ModelConfigID: invalidModelConfigID, }) @@ -1869,7 +1948,6 @@ func TestCreateChatInsertsWorkspaceAwarenessMessage(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), WorkspaceID: uuid.NullUUID{UUID: workspace.ID, Valid: true}, Title: "test-with-workspace", ModelConfigID: model.ID, @@ -1907,7 +1985,6 @@ func TestCreateChatInsertsWorkspaceAwarenessMessage(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "test-without-workspace", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -1987,7 +2064,6 @@ func TestCreateChatRejectsWhenUsageLimitReached(t *testing.T) { _, err = replica.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "over-limit", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -2085,7 +2161,6 @@ func TestAutoPromoteQueuedMessagesPreservesPerTurnModelOrder(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "auto-promote per-turn model order", ModelConfigID: modelConfigA.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -2096,7 +2171,6 @@ func TestAutoPromoteQueuedMessagesPreservesPerTurnModelOrder(t *testing.T) { queuedB, err := server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("queued b")}, ModelConfigID: modelConfigB.ID, BusyBehavior: chatd.SendMessageBusyBehaviorQueue, @@ -2106,7 +2180,6 @@ func TestAutoPromoteQueuedMessagesPreservesPerTurnModelOrder(t *testing.T) { queuedC, err := server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("queued c")}, ModelConfigID: modelConfigC.ID, BusyBehavior: chatd.SendMessageBusyBehaviorQueue, @@ -2253,7 +2326,6 @@ func TestInterruptAutoPromotionIgnoresLaterUsageLimitIncrease(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "interrupt-autopromote-limit", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -2264,7 +2336,6 @@ func TestInterruptAutoPromotionIgnoresLaterUsageLimitIncrease(t *testing.T) { queuedResult, err := server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("queued")}, BusyBehavior: chatd.SendMessageBusyBehaviorInterrupt, }) @@ -2278,9 +2349,8 @@ func TestInterruptAutoPromotionIgnoresLaterUsageLimitIncrease(t *testing.T) { testutil.TryReceive(ctx, t, secondRequestStarted) laterQueuedResult, err := server.SendMessage(ctx, chatd.SendMessageOptions{ - ChatID: chat.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), - Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("later queued")}, + ChatID: chat.ID, + Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("later queued")}, }) require.NoError(t, err) require.True(t, laterQueuedResult.Queued) @@ -2362,7 +2432,6 @@ func TestEditMessageRejectsWhenUsageLimitReached(t *testing.T) { chat, err := replica.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "edit-limit-reached", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("original")}, @@ -2393,7 +2462,6 @@ func TestEditMessageRejectsWhenUsageLimitReached(t *testing.T) { _, err = replica.EditMessage(ctx, chatd.EditMessageOptions{ ChatID: chat.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), EditedMessageID: editedMessageID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edited")}, }) @@ -2427,7 +2495,6 @@ func TestEditMessageRejectsMissingMessage(t *testing.T) { chat, err := replica.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "missing-edited-message", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -2436,7 +2503,6 @@ func TestEditMessageRejectsMissingMessage(t *testing.T) { _, err = replica.EditMessage(ctx, chatd.EditMessageOptions{ ChatID: chat.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), EditedMessageID: 999999, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edited")}, }) @@ -2456,7 +2522,6 @@ func TestEditMessageRejectsNonUserMessage(t *testing.T) { chat, err := replica.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "non-user-edited-message", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -2478,7 +2543,6 @@ func TestEditMessageRejectsNonUserMessage(t *testing.T) { _, err = replica.EditMessage(ctx, chatd.EditMessageOptions{ ChatID: chat.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), EditedMessageID: assistantMessage.ID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edited")}, }) @@ -2505,7 +2569,6 @@ func TestEditMessageDebugCleanupDeletesPreEditRuns(t *testing.T) { chat, err := replica.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "debug-edit-cleanup", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("first")}, @@ -2555,7 +2618,6 @@ func TestEditMessageDebugCleanupDeletesPreEditRuns(t *testing.T) { _, err = replica.EditMessage(ctx, chatd.EditMessageOptions{ ChatID: chat.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), EditedMessageID: editedMsgID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edited")}, }) @@ -2613,7 +2675,6 @@ func TestEditMessageDebugCleanupPreservesRecentRuns(t *testing.T) { chat, err := replica.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "debug-edit-buffer", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("first")}, @@ -2646,7 +2707,6 @@ func TestEditMessageDebugCleanupPreservesRecentRuns(t *testing.T) { _, err = replica.EditMessage(ctx, chatd.EditMessageOptions{ ChatID: chat.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), EditedMessageID: editedMsgID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edited")}, }) @@ -2966,7 +3026,6 @@ func TestSubscribeSnapshotIncludesStatusEvent(t *testing.T) { chat, err := replica.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "status-snapshot", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -3086,7 +3145,6 @@ func TestPersistToolResultWithBinaryData(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "binary-tool-result", ModelConfigID: model.ID, WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, @@ -3219,7 +3277,6 @@ func TestRequiresActionChatPersistsWaitingStatusLabel(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "requires-action-status-label", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{ @@ -3341,7 +3398,6 @@ func TestActiveServer_InterruptionBehavior(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, Title: "interrupt-partial-tool", @@ -3356,7 +3412,6 @@ func TestActiveServer_InterruptionBehavior(t *testing.T) { queued, err := server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, CreatedBy: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ModelConfigID: model.ID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("queued after interrupt")}, BusyBehavior: chatd.SendMessageBusyBehaviorInterrupt, @@ -3448,7 +3503,6 @@ func TestActiveServer_InterruptionBehavior(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, Title: "interrupt-tool-execution", @@ -3463,7 +3517,6 @@ func TestActiveServer_InterruptionBehavior(t *testing.T) { queued, err := server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, CreatedBy: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ModelConfigID: model.ID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("continue after interrupt")}, BusyBehavior: chatd.SendMessageBusyBehaviorInterrupt, @@ -3564,7 +3617,6 @@ func TestActiveServer_InterruptionBehavior(t *testing.T) { queued, err := server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, CreatedBy: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ModelConfigID: model.ID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("continue after provider interrupt")}, BusyBehavior: chatd.SendMessageBusyBehaviorInterrupt, @@ -3662,7 +3714,6 @@ func TestActiveServer_InterruptionBehavior(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, Title: "anthropic-mixed-interrupt", @@ -3676,7 +3727,6 @@ func TestActiveServer_InterruptionBehavior(t *testing.T) { queued, err := server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, CreatedBy: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ModelConfigID: model.ID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("continue after mixed interrupt")}, BusyBehavior: chatd.SendMessageBusyBehaviorInterrupt, @@ -3763,7 +3813,6 @@ func TestActiveServer_InterruptionBehavior(t *testing.T) { queued, err := server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, CreatedBy: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ModelConfigID: model.ID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("continue after reasoning")}, BusyBehavior: chatd.SendMessageBusyBehaviorInterrupt, @@ -3818,7 +3867,6 @@ func TestActiveServer_DynamicToolsAndStopAfterToolBehavior(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "dynamic-tool-requires-action", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{ @@ -3879,7 +3927,6 @@ func TestActiveServer_DynamicToolsAndStopAfterToolBehavior(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "stop-after-success", ModelConfigID: model.ID, WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, @@ -3929,7 +3976,6 @@ func TestActiveServer_DynamicToolsAndStopAfterToolBehavior(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "stop-after-error", ModelConfigID: model.ID, WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, @@ -4022,7 +4068,6 @@ func TestDynamicToolCallPausesAndResumes(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "dynamic-tool-pause-resume", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{ @@ -4193,7 +4238,6 @@ func TestDynamicToolNamedProposePlanRemainsAvailableOutsidePlanMode(t *testing.T chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "dynamic-propose-plan-collision", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{ @@ -4311,7 +4355,6 @@ func TestDynamicToolCallMixedWithBuiltIn(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "mixed-builtin-dynamic", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{ @@ -4454,7 +4497,6 @@ func TestSubmitToolResultsConcurrency(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "concurrency-tool-results", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{ @@ -4583,7 +4625,6 @@ func TestSubscribeNoDuplicateMessageParts(t *testing.T) { chat, err := replica.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "no-dup-parts", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -5052,7 +5093,6 @@ func TestStoppedWorkspaceWithPersistedAgentBindingDoesNotBlockChat(t *testing.T) chat, err := inactive.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "stopped-workspace-regression", ModelConfigID: model.ID, WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, @@ -5220,7 +5260,6 @@ func TestHeartbeatNoWorkspaceNoBump(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "no-workspace-test", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -5379,7 +5418,6 @@ func TestActiveServer_RoutingPreservesAPIKeyAfterCompaction(t *testing.T) { model = updateChatModelCompressionThreshold(t, db, model, contextLimit, thresholdPercent) provider, err := db.GetAIProviderByID(ctx, model.AIProviderID.UUID) require.NoError(t, err) - apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID}) _, err = db.UpsertUserAIProviderKey(ctx, database.UpsertUserAIProviderKeyParams{ ID: uuid.New(), UserID: user.ID, @@ -5404,12 +5442,13 @@ func TestActiveServer_RoutingPreservesAPIKeyAfterCompaction(t *testing.T) { ModelConfigID: model.ID, WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, - APIKeyID: apiKey.ID, InitialUserContent: []codersdk.ChatMessagePart{ codersdk.ChatMessageText("trigger compaction"), }, }) require.NoError(t, err) + mapping, err := db.GetChatSyntheticAPIKeyByUserID(ctx, user.ID) + require.NoError(t, err) contextContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{{ Type: codersdk.ChatMessagePartTypeContextFile, ContextFileAgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, @@ -5421,7 +5460,7 @@ func TestActiveServer_RoutingPreservesAPIKeyAfterCompaction(t *testing.T) { require.NoError(t, err) _, err = db.InsertChatMessages(ctx, chatd.BuildSingleUserChatMessageInsertParams( chat.ID, - apiKey.ID, + mapping.APIKeyID, contextContent, database.ChatMessageVisibilityBoth, model.ID, @@ -5450,14 +5489,14 @@ func TestActiveServer_RoutingPreservesAPIKeyAfterCompaction(t *testing.T) { compressed := compressedChatSummarizedMessages(t, append(promptMessages, messages...)) require.Len(t, compressed.summaries, 1) require.True(t, compressed.summaries[0].APIKeyID.Valid) - require.Equal(t, apiKey.ID, compressed.summaries[0].APIKeyID.String) + require.Equal(t, mapping.APIKeyID, compressed.summaries[0].APIKeyID.String) requests := factory.RequestsSnapshot() require.NotEmpty(t, requests) for _, req := range requests { require.Equal(t, provider.Name, req.ProviderName) require.Equal(t, aibridge.SourceAgents, req.Source) - require.Equal(t, apiKey.ID, req.APIKeyID) + require.Equal(t, mapping.APIKeyID, req.APIKeyID) require.Equal(t, "sk-user-aibridge", req.Request.Header.Get("X-Api-Key")) require.Equal(t, "delegated", req.Request.Header.Get(aibridge.HeaderCoderToken)) } @@ -5520,7 +5559,6 @@ func TestActiveServer_CompactionRecordsMetric(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, Title: "compaction-metric", @@ -5613,7 +5651,6 @@ func TestActiveServer_Compaction(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, Title: "compaction-continues", @@ -5741,7 +5778,6 @@ func TestActiveServer_Compaction(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, Title: "compaction-next-message-over-limit", @@ -5760,7 +5796,6 @@ func TestActiveServer_Compaction(t *testing.T) { _, err = server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, CreatedBy: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ModelConfigID: model.ID, Content: []codersdk.ChatMessagePart{ codersdk.ChatMessageText("continue after the large compacted turn"), @@ -5943,7 +5978,6 @@ func TestActiveServer_CompactionModelOverride(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, Title: "compaction-override", @@ -6048,7 +6082,6 @@ func TestActiveServer_CompactionModelOverride(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, Title: "compaction-override-limit", @@ -6093,7 +6126,6 @@ func TestActiveServer_BasicAssistantGenerationAndPromptPreparation(t *testing.T) _, err := server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, CreatedBy: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ModelConfigID: model.ID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("continue")}, BusyBehavior: chatd.SendMessageBusyBehaviorQueue, @@ -6132,7 +6164,6 @@ func TestActiveServer_BasicAssistantGenerationAndPromptPreparation(t *testing.T) _, err = server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: planChat.ID, CreatedBy: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ModelConfigID: model.ID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("continue")}, BusyBehavior: chatd.SendMessageBusyBehaviorQueue, @@ -6185,7 +6216,6 @@ func TestActiveServer_ToolExecutionAndPolicy(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, Title: "active-tool-reject", @@ -6266,7 +6296,6 @@ func TestActiveServer_ToolExecutionAndPolicy(t *testing.T) { _, err = server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, CreatedBy: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ModelConfigID: model.ID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("continue")}, BusyBehavior: chatd.SendMessageBusyBehaviorQueue, @@ -6329,7 +6358,6 @@ func TestActiveServer_ToolExecutionAndPolicy(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, Title: "multi-step-tool", @@ -6411,7 +6439,6 @@ func TestActiveServer_ToolExecutionAndPolicy(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, Title: "parallel-timestamps", @@ -6577,7 +6604,6 @@ func TestActiveServer_ToolErrorRecordsMetric(t *testing.T) { chatOpts := chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, Title: "tool-error-metric", @@ -7013,7 +7039,6 @@ func TestActiveServer_ChatTurnDebugRunRecordsMultipleStreamSteps(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, Title: "multi-step-debug", @@ -7101,7 +7126,6 @@ func TestActiveServer_AnthropicSanitizesProviderToolBeforeRequest(t *testing.T) _, err := server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, CreatedBy: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ModelConfigID: model.ID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("continue")}, BusyBehavior: chatd.SendMessageBusyBehaviorQueue, @@ -7209,7 +7233,6 @@ func TestActiveServer_AnthropicProviderToolPreRequestGuard(t *testing.T) { _, err := server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, CreatedBy: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ModelConfigID: model.ID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("continue")}, BusyBehavior: chatd.SendMessageBusyBehaviorQueue, @@ -7246,7 +7269,6 @@ func TestActiveServer_AnthropicProviderToolPreRequestGuard(t *testing.T) { _, err := server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, CreatedBy: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ModelConfigID: model.ID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("continue")}, BusyBehavior: chatd.SendMessageBusyBehaviorQueue, @@ -7406,7 +7428,6 @@ func TestActiveServer_AnthropicWebSearchFollowUpHasNoSyntheticCancellation(t *te _, err := server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, CreatedBy: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("thanks, tell me more")}, }) require.NoError(t, err) @@ -7484,7 +7505,6 @@ func TestActiveServer_AnthropicSanitizesWebSearchBeforeContinuation(t *testing.T chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, Title: "anthropic-web-search-continuation", @@ -7557,7 +7577,6 @@ func TestActiveServer_ExclusiveToolPolicy(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, AgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, Title: "exclusive-local-policy", @@ -7614,7 +7633,6 @@ func TestActiveServer_ExclusiveToolPolicy(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "exclusive-dynamic-policy", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{ @@ -8441,7 +8459,6 @@ func TestPassiveServerDoesNotProcess(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "should-stay-pending", InitialUserContent: []codersdk.ChatMessagePart{{Type: codersdk.ChatMessagePartTypeText, Text: "hello"}}, ModelConfigID: model.ID, @@ -8746,7 +8763,6 @@ func seedChatDependenciesWithProvider( t.Helper() user := dbgen.User(t, db, database.User{}) - _ = testAPIKeyID(t, db, user.ID) org := dbgen.Organization(t, db, database.Organization{}) dbgen.OrganizationMember(t, db, database.OrganizationMember{ UserID: user.ID, @@ -8777,7 +8793,6 @@ func seedChatDependenciesWithProviderPolicy( t.Helper() user := dbgen.User(t, db, database.User{}) - _ = testAPIKeyID(t, db, user.ID) org := dbgen.Organization(t, db, database.Organization{}) dbgen.OrganizationMember(t, db, database.OrganizationMember{ UserID: user.ID, @@ -9048,7 +9063,6 @@ func TestInterruptChatDoesNotSendWebPushNotification(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "interrupt-no-push", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -9171,7 +9185,6 @@ func TestSuccessfulChatSendsWebPushWithNavigationData(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "push-nav-test", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -9259,7 +9272,6 @@ func TestCloseDuringShutdownContextCanceledShouldRetryOnNewReplica(t *testing.T) chat, err := serverA.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "shutdown-retry", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -9360,7 +9372,6 @@ func TestSuccessfulChatSendsWebPushWithSummary(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "summary-push-test", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("do the thing")}, @@ -9419,7 +9430,6 @@ func TestSuccessfulChatPersistsTurnSummaryWithoutWebPush(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "summary-no-webpush-test", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("do the thing")}, @@ -9481,7 +9491,6 @@ func TestSuccessfulChatSendsWebPushFallbackWithoutSummaryForEmptyAssistantText(t chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "empty-summary-push-test", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("do the thing")}, @@ -9542,7 +9551,6 @@ func TestErroredChatClearsLastTurnSummaryAndSendsWebPush(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "error-summary-clear-test", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("do the thing")}, @@ -9776,7 +9784,6 @@ func TestComputerUseSubagentToolsAndModel(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "computer-use-detection", ModelConfigID: model.ID, WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, @@ -9943,7 +9950,6 @@ func TestInterruptChatPersistsPartialResponse(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "interrupt-persist-test", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -10048,7 +10054,6 @@ func TestProcessChat_UserProviderKey_Success(t *testing.T) { chat, err := creator.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "user-provider-key-success", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{ @@ -10080,6 +10085,7 @@ func seedAIGatewayOpenAITestDependencies( user := dbgen.User(t, db, database.User{}) org := dbgen.Organization(t, db, database.Organization{}) + apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID}) dbgen.OrganizationMember(t, db, database.OrganizationMember{ UserID: user.ID, OrganizationID: org.ID, @@ -10094,7 +10100,6 @@ func seedAIGatewayOpenAITestDependencies( IsDefault: true, AIProviderID: uuid.NullUUID{UUID: provider.ID, Valid: true}, }) - apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID}) _, err := db.UpsertUserAIProviderKey(context.Background(), database.UpsertUserAIProviderKeyParams{ ID: uuid.New(), UserID: user.ID, @@ -10122,7 +10127,7 @@ func TestProcessChat_RoutingUsesDelegatedAPIKey(t *testing.T) { }) factory := chattest.NewMockAIBridgeTransport(t, openAIURL) - user, org, provider, model, apiKey := seedAIGatewayOpenAITestDependencies(t, db, openAIURL) + user, org, provider, model, _ := seedAIGatewayOpenAITestDependencies(t, db, openAIURL) creator := newTestServer(t, db, ps, uuid.New()) chat, err := creator.CreateChat(ctx, chatd.CreateOptions{ @@ -10130,12 +10135,13 @@ func TestProcessChat_RoutingUsesDelegatedAPIKey(t *testing.T) { OwnerID: user.ID, Title: "aigateway-routing", ModelConfigID: model.ID, - APIKeyID: apiKey.ID, InitialUserContent: []codersdk.ChatMessagePart{ codersdk.ChatMessageText("say hello"), }, }) require.NoError(t, err) + mapping, err := db.GetChatSyntheticAPIKeyByUserID(ctx, user.ID) + require.NoError(t, err) _, events, cancel, ok := creator.Subscribe(ctx, chat.ID, nil, 0) require.True(t, ok) @@ -10161,7 +10167,7 @@ func TestProcessChat_RoutingUsesDelegatedAPIKey(t *testing.T) { for _, req := range requests { require.Equal(t, provider.Name, req.ProviderName) require.Equal(t, aibridge.SourceAgents, req.Source) - require.Equal(t, apiKey.ID, req.APIKeyID) + require.Equal(t, mapping.APIKeyID, req.APIKeyID) require.Equal(t, "Bearer sk-user-aibridge", req.Request.Header.Get("Authorization")) require.Empty(t, req.Request.Header.Get("X-Api-Key")) require.Equal(t, "delegated", req.Request.Header.Get(aibridge.HeaderCoderToken)) @@ -10184,7 +10190,7 @@ func TestProcessChat_RoutingPreservesAPIKeyAfterWorkspaceContext(t *testing.T) { return chattest.OpenAINonStreamingResponse(`{"title":"AI Gateway Workspace"}`) }) factory := chattest.NewMockAIBridgeTransport(t, openAIURL) - user, org, provider, model, apiKey := seedAIGatewayOpenAITestDependencies(t, db, openAIURL) + user, org, provider, model, _ := seedAIGatewayOpenAITestDependencies(t, db, openAIURL) ws, dbAgent := seedWorkspaceWithAgent(t, db, user.ID) creator := newTestServer(t, db, ps, uuid.New()) @@ -10194,12 +10200,13 @@ func TestProcessChat_RoutingPreservesAPIKeyAfterWorkspaceContext(t *testing.T) { Title: "aigateway-workspace-context", ModelConfigID: model.ID, WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, - APIKeyID: apiKey.ID, InitialUserContent: []codersdk.ChatMessagePart{ codersdk.ChatMessageText("use the workspace context"), }, }) require.NoError(t, err) + mapping, err := db.GetChatSyntheticAPIKeyByUserID(ctx, user.ID) + require.NoError(t, err) const contextText = "# Project instructions\nAlways keep routing metadata." // Workspace context is sourced from the agent's pinned snapshot. Seed it so @@ -10236,7 +10243,7 @@ func TestProcessChat_RoutingPreservesAPIKeyAfterWorkspaceContext(t *testing.T) { for _, req := range requests { require.Equal(t, provider.Name, req.ProviderName) require.Equal(t, aibridge.SourceAgents, req.Source) - require.Equal(t, apiKey.ID, req.APIKeyID) + require.Equal(t, mapping.APIKeyID, req.APIKeyID) require.Equal(t, "Bearer sk-user-aibridge", req.Request.Header.Get("Authorization")) require.Equal(t, "delegated", req.Request.Header.Get(aibridge.HeaderCoderToken)) } @@ -10274,7 +10281,6 @@ func TestProcessChat_UserProviderKey_MissingKeyError(t *testing.T) { chat, err := creator.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "user-provider-key-missing", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{ @@ -10338,7 +10344,6 @@ func TestProcessChatPanicRecovery(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "panic-recovery", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{ @@ -10522,7 +10527,6 @@ func TestMCPServerToolInvocation(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "mcp-tool-test", ModelConfigID: model.ID, WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, @@ -10681,7 +10685,6 @@ func TestPlanModeRootChatApprovedExternalMCPToolInvocation(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "plan-mode-mcp-invocation", ModelConfigID: model.ID, PlanMode: database.NullChatPlanMode{ChatPlanMode: database.ChatPlanModePlan, Valid: true}, @@ -10812,7 +10815,6 @@ func TestPlanModeRootChatApprovedExternalMCPWorkflowCanReachProposePlan(t *testi chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "plan-mode-mcp-propose-plan", ModelConfigID: model.ID, WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, @@ -11032,7 +11034,6 @@ func TestMCPServerOAuth2TokenRefresh(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "oauth2-refresh-test", ModelConfigID: model.ID, WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, @@ -11144,7 +11145,6 @@ func TestMCPServerOAuth2TokenRefreshFailureGraceful(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "graceful-degradation-test", ModelConfigID: model.ID, MCPServerIDs: []uuid.UUID{mcpConfig.ID}, @@ -11270,7 +11270,6 @@ func TestChatTemplateAllowlistEnforcement(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "allowlist-test", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{ @@ -11424,7 +11423,6 @@ func TestChatAsksUserWhenListTemplatesRequiresSelection(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "ask-template-selection-test", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{ @@ -11552,7 +11550,6 @@ func TestCreateChatImmediatelyProcessesNewChat(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "wake-test", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -11618,7 +11615,6 @@ func TestSendMessageImmediatelyProcessesWaitingChat(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "wake-send-test", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("first")}, @@ -11634,9 +11630,8 @@ func TestSendMessageImmediatelyProcessesWaitingChat(t *testing.T) { // Now send a follow-up message, which should also be // processed immediately without waiting for the acquire ticker. _, err = server.SendMessage(ctx, chatd.SendMessageOptions{ - ChatID: chat.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), - Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("second")}, + ChatID: chat.ID, + Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("second")}, }) require.NoError(t, err) @@ -11852,7 +11847,6 @@ func TestEditMessageWithModelConfigOverride(t *testing.T) { chat, err := replica.CreateChat(ctx, chatd.CreateOptions{ OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), OrganizationID: org.ID, Title: "edit-with-model-override", ModelConfigID: modelA.ID, @@ -11870,7 +11864,6 @@ func TestEditMessageWithModelConfigOverride(t *testing.T) { result, err := replica.EditMessage(ctx, chatd.EditMessageOptions{ ChatID: chat.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), EditedMessageID: initial[0].ID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edited")}, ModelConfigID: modelB.ID, @@ -11899,7 +11892,6 @@ func TestEditMessagePreservesModelConfigByDefault(t *testing.T) { chat, err := replica.CreateChat(ctx, chatd.CreateOptions{ OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), OrganizationID: org.ID, Title: "edit-preserves-model", ModelConfigID: modelA.ID, @@ -11916,7 +11908,6 @@ func TestEditMessagePreservesModelConfigByDefault(t *testing.T) { result, err := replica.EditMessage(ctx, chatd.EditMessageOptions{ ChatID: chat.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), EditedMessageID: initial[0].ID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edited")}, }) @@ -11952,7 +11943,6 @@ func TestEditMessageReasoningEffort(t *testing.T) { chat, err := replica.CreateChat(ctx, chatd.CreateOptions{ OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), OrganizationID: org.ID, Title: "edit-reasoning-effort", ModelConfigID: model.ID, @@ -11972,7 +11962,6 @@ func TestEditMessageReasoningEffort(t *testing.T) { result, err := replica.EditMessage(ctx, chatd.EditMessageOptions{ ChatID: chat.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), EditedMessageID: initial[0].ID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edited")}, ReasoningEffort: tc.requested, @@ -12003,7 +11992,6 @@ func TestEditMessageRejectsUnknownModelConfig(t *testing.T) { chat, err := replica.CreateChat(ctx, chatd.CreateOptions{ OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), OrganizationID: org.ID, Title: "edit-unknown-model", ModelConfigID: modelA.ID, @@ -12020,7 +12008,6 @@ func TestEditMessageRejectsUnknownModelConfig(t *testing.T) { _, err = replica.EditMessage(ctx, chatd.EditMessageOptions{ ChatID: chat.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), EditedMessageID: initial[0].ID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("edited")}, ModelConfigID: uuid.New(), @@ -12140,7 +12127,6 @@ func TestAdvisorGating_ExperimentDisabled(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "advisor-experiment-disabled", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{ @@ -12240,7 +12226,6 @@ func TestAdvisorGating_RootChat(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "advisor-root", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{ @@ -12419,7 +12404,6 @@ func TestAdvisorHappyPath_RootChat(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "advisor-happy-path", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{ @@ -12613,7 +12597,6 @@ func TestAdvisorGating_ChildChat(t *testing.T) { childChat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "advisor-child", ModelConfigID: model.ID, ParentChatID: uuid.NullUUID{UUID: parent.ID, Valid: true}, @@ -12696,7 +12679,6 @@ func TestAdvisorGating_PlanMode(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "advisor-plan-mode", ModelConfigID: model.ID, PlanMode: database.NullChatPlanMode{ChatPlanMode: database.ChatPlanModePlan, Valid: true}, @@ -12780,7 +12762,6 @@ func TestAdvisorGating_ExploreSubagent(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "advisor-explore", ModelConfigID: model.ID, ChatMode: database.NullChatMode{ @@ -12898,7 +12879,6 @@ func TestProviderSwitchSanitizesAndRestoresPEToolHistory(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "provider-switch-test", ModelConfigID: mA.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -12929,7 +12909,6 @@ func TestProviderSwitchSanitizesAndRestoresPEToolHistory(t *testing.T) { _, err = server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, CreatedBy: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ModelConfigID: mB.ID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("continue with B")}, }) @@ -12940,7 +12919,6 @@ func TestProviderSwitchSanitizesAndRestoresPEToolHistory(t *testing.T) { _, err = server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, CreatedBy: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ModelConfigID: mA.ID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText("back to A")}, }) diff --git a/coderd/x/chatd/generation.go b/coderd/x/chatd/generation.go index 25d553ef2cd..b9c43d8063f 100644 --- a/coderd/x/chatd/generation.go +++ b/coderd/x/chatd/generation.go @@ -651,12 +651,7 @@ func (s *taskStarter) executeLocalTools( provider = prepared.Model.Provider() modelName = prepared.Model.Model() } - // Local tool callbacks (e.g. spawn_agent, message_agent) read the - // active turn's delegated API key ID from the context to route - // subagent traffic through the AI Gateway. prepareGeneration sets it - // only on its own context, so re-derive it here for tool execution. - toolCtx := withActiveTurnAPIKeyID(ctx, prepared.ModelBuildOptions) - outcome, err := chatloop.ExecuteLocalTools(toolCtx, chatloop.ExecuteLocalToolsOptions{ + outcome, err := chatloop.ExecuteLocalTools(ctx, chatloop.ExecuteLocalToolsOptions{ Tools: prepared.Tools, ActiveTools: prepared.ActiveTools, ProviderTools: prepared.ProviderTools, diff --git a/coderd/x/chatd/generation_preparer.go b/coderd/x/chatd/generation_preparer.go index ffee255d126..06b47af4e42 100644 --- a/coderd/x/chatd/generation_preparer.go +++ b/coderd/x/chatd/generation_preparer.go @@ -80,10 +80,12 @@ func (server *Server) prepareGeneration( return generationPrepared{}, err } - modelOpts = modelBuildOptionsFromMessages(promptRows) - ctx = withActiveTurnAPIKeyID(ctx, modelOpts) + apiKeyID, err := server.ensureSyntheticAPIKeyID(ctx, chat.OwnerID) + if err != nil { + return generationPrepared{}, xerrors.Errorf("ensure synthetic API key: %w", err) + } + modelOpts = modelBuildOptions{ActiveAPIKeyID: apiKeyID} - var err error model, modelConfig, modelRoute, debugEnabled, resolvedProvider, debugModel, err = server.resolveChatModel(ctx, chat, modelOpts) if err != nil { return generationPrepared{}, err @@ -760,8 +762,12 @@ func (server *Server) deriveFinalTurnRunResult( // resolvedProvider/resolvedModel describe the model the fallback handle was // built from; they only feed the status-label fallback candidate's labels. - modelOpts := modelBuildOptionsFromMessages(promptRows) - ctx = withActiveTurnAPIKeyID(ctx, modelOpts) + apiKeyID, err := server.ensureSyntheticAPIKeyID(ctx, chat.OwnerID) + if err != nil { + logger.Warn(ctx, "derive final turn status label: ensure synthetic API key", slog.Error(err)) + return runChatResult{FinalAssistantText: finalAssistantText, TriggerMessageID: triggerMessageID, HistoryTipMessageID: historyTipMessageID} + } + modelOpts := modelBuildOptions{ActiveAPIKeyID: apiKeyID} model, _, modelRoute, _, resolvedProvider, resolvedModel, err := server.resolveChatModel(ctx, chat, modelOpts) if err != nil { // Return what we have; generateFinalTurnStatusLabel falls back to a diff --git a/coderd/x/chatd/helpers_test.go b/coderd/x/chatd/helpers_test.go index 2a5f0a42dd4..178c51b9461 100644 --- a/coderd/x/chatd/helpers_test.go +++ b/coderd/x/chatd/helpers_test.go @@ -23,12 +23,6 @@ import ( "github.com/coder/coder/v2/testutil" ) -func testAPIKeyID(t testing.TB, db database.Store, userID uuid.UUID) string { - t.Helper() - key, _ := dbgen.APIKey(t, db, database.APIKey{ID: uuid.NewString(), UserID: userID}) - return key.ID -} - type workerTestFixture struct { db database.Store pubsub dbpubsub.Pubsub diff --git a/coderd/x/chatd/integration_responses_test.go b/coderd/x/chatd/integration_responses_test.go index 6d27463f055..b35c4a77c0c 100644 --- a/coderd/x/chatd/integration_responses_test.go +++ b/coderd/x/chatd/integration_responses_test.go @@ -72,7 +72,6 @@ func TestOpenAIResponsesNoStaleWebSearchReplay(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: uniqueResponsesTitle(t, "no-stale"), ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{ @@ -87,7 +86,6 @@ func TestOpenAIResponsesNoStaleWebSearchReplay(t *testing.T) { _, err = server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, CreatedBy: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ModelConfigID: model.ID, Content: []codersdk.ChatMessagePart{ codersdk.ChatMessageText("summarize the result without searching again"), @@ -159,7 +157,6 @@ func TestOpenAIResponsesFullReplayPairsReasoningAndWebSearch(t *testing.T) { chat, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: uniqueResponsesTitle(t, "full-replay"), ModelConfigID: firstModel.ID, InitialUserContent: []codersdk.ChatMessagePart{ @@ -174,7 +171,6 @@ func TestOpenAIResponsesFullReplayPairsReasoningAndWebSearch(t *testing.T) { _, err = server.SendMessage(ctx, chatd.SendMessageOptions{ ChatID: chat.ID, CreatedBy: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ModelConfigID: secondModel.ID, Content: []codersdk.ChatMessagePart{ codersdk.ChatMessageText("summarize the result without searching again"), diff --git a/coderd/x/chatd/model_routing.go b/coderd/x/chatd/model_routing.go index 8678dffea87..f67beb17976 100644 --- a/coderd/x/chatd/model_routing.go +++ b/coderd/x/chatd/model_routing.go @@ -9,7 +9,6 @@ import ( "github.com/google/uuid" "golang.org/x/xerrors" - "github.com/coder/coder/v2/coderd/aibridge" "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" ) @@ -29,22 +28,6 @@ type modelBuildOptions struct { RecordHTTP bool } -func modelBuildOptionsFromMessages(messages []database.ChatMessage) modelBuildOptions { - apiKeyID, _ := activeTurnAPIKeyIDFromMessages(messages) - return modelBuildOptions{ActiveAPIKeyID: apiKeyID} -} - -// withActiveTurnAPIKeyID augments ctx with the active turn's delegated API -// key ID when one is known. AI Gateway routing and subagent tool callbacks -// read this value from the context to attribute requests to the correct -// turn. When no key is known, ctx is returned unchanged. -func withActiveTurnAPIKeyID(ctx context.Context, opts modelBuildOptions) context.Context { - if opts.ActiveAPIKeyID == "" { - return ctx - } - return aibridge.WithDelegatedAPIKeyID(ctx, opts.ActiveAPIKeyID) -} - func (p *Server) enabledAIProviderByID(ctx context.Context, providerID uuid.UUID) (database.AIProvider, error) { provider, err := p.db.GetAIProviderByID(ctx, providerID) if err != nil { diff --git a/coderd/x/chatd/model_routing_internal_test.go b/coderd/x/chatd/model_routing_internal_test.go index 6b02a191150..e6d9103435f 100644 --- a/coderd/x/chatd/model_routing_internal_test.go +++ b/coderd/x/chatd/model_routing_internal_test.go @@ -18,9 +18,7 @@ import ( "github.com/coder/coder/v2/coderd/aibridge" "github.com/coder/coder/v2/coderd/database" - "github.com/coder/coder/v2/coderd/database/dbgen" "github.com/coder/coder/v2/coderd/database/dbmock" - "github.com/coder/coder/v2/coderd/database/dbtestutil" "github.com/coder/coder/v2/coderd/x/chatd/chaterror" "github.com/coder/coder/v2/coderd/x/chatd/chatprovider" "github.com/coder/coder/v2/coderd/x/chatd/chattool" @@ -367,215 +365,6 @@ func TestAIGatewayModelForwardsProviderAuth(t *testing.T) { }) } -func TestActiveTurnAPIKeyIDFromMessages(t *testing.T) { - t.Parallel() - - oldKeyID := uuid.NewString() - currentKeyID := uuid.NewString() - tests := []struct { - name string - messages []database.ChatMessage - wantKey string - wantOK bool - }{ - { - name: "CurrentUserMessage", - messages: []database.ChatMessage{ - {ID: 1, Role: database.ChatMessageRoleUser, Visibility: database.ChatMessageVisibilityBoth, APIKeyID: sqlNullString(oldKeyID)}, - {ID: 2, Role: database.ChatMessageRoleAssistant, Visibility: database.ChatMessageVisibilityBoth}, - {ID: 3, Role: database.ChatMessageRoleUser, Visibility: database.ChatMessageVisibilityBoth, APIKeyID: sqlNullString(currentKeyID)}, - }, - wantKey: currentKeyID, - wantOK: true, - }, - { - name: "MissingCurrentUserAPIKeyDoesNotFallBack", - messages: []database.ChatMessage{ - {ID: 1, Role: database.ChatMessageRoleUser, Visibility: database.ChatMessageVisibilityBoth, APIKeyID: sqlNullString(oldKeyID)}, - {ID: 2, Role: database.ChatMessageRoleUser, Visibility: database.ChatMessageVisibilityBoth}, - }, - }, - { - name: "SkipsUncompressedModelOnlyUserMessages", - messages: []database.ChatMessage{ - {ID: 1, Role: database.ChatMessageRoleUser, Visibility: database.ChatMessageVisibilityBoth, APIKeyID: sqlNullString(oldKeyID)}, - {ID: 2, Role: database.ChatMessageRoleUser, Visibility: database.ChatMessageVisibilityModel, APIKeyID: sqlNullString(currentKeyID)}, - }, - wantKey: oldKeyID, - wantOK: true, - }, - { - name: "CompressedSummaryFallback", - messages: []database.ChatMessage{ - {ID: 1, Role: database.ChatMessageRoleUser, Visibility: database.ChatMessageVisibilityModel, Compressed: true, APIKeyID: sqlNullString(currentKeyID)}, - {ID: 2, Role: database.ChatMessageRoleAssistant, Visibility: database.ChatMessageVisibilityBoth}, - }, - wantKey: currentKeyID, - wantOK: true, - }, - { - name: "LatestCompressedSummaryWins", - messages: []database.ChatMessage{ - {ID: 1, Role: database.ChatMessageRoleUser, Visibility: database.ChatMessageVisibilityModel, Compressed: true, APIKeyID: sqlNullString(oldKeyID)}, - {ID: 2, Role: database.ChatMessageRoleUser, Visibility: database.ChatMessageVisibilityModel, Compressed: true, APIKeyID: sqlNullString(currentKeyID)}, - {ID: 3, Role: database.ChatMessageRoleAssistant, Visibility: database.ChatMessageVisibilityBoth}, - }, - wantKey: currentKeyID, - wantOK: true, - }, - { - name: "VisibleUserWinsOverCompressedSummary", - messages: []database.ChatMessage{ - {ID: 1, Role: database.ChatMessageRoleUser, Visibility: database.ChatMessageVisibilityModel, Compressed: true, APIKeyID: sqlNullString(oldKeyID)}, - {ID: 2, Role: database.ChatMessageRoleUser, Visibility: database.ChatMessageVisibilityBoth, APIKeyID: sqlNullString(currentKeyID)}, - }, - wantKey: currentKeyID, - wantOK: true, - }, - { - name: "MissingVisibleUserKeyDoesNotFallBackToCompressedSummary", - messages: []database.ChatMessage{ - {ID: 1, Role: database.ChatMessageRoleUser, Visibility: database.ChatMessageVisibilityModel, Compressed: true, APIKeyID: sqlNullString(oldKeyID)}, - {ID: 2, Role: database.ChatMessageRoleUser, Visibility: database.ChatMessageVisibilityBoth}, - }, - }, - { - name: "UncompressedModelOnlyUserIgnored", - messages: []database.ChatMessage{ - {ID: 1, Role: database.ChatMessageRoleUser, Visibility: database.ChatMessageVisibilityModel, APIKeyID: sqlNullString(currentKeyID)}, - }, - }, - { - name: "CompressedSummaryMissingKeyDoesNotFallBack", - messages: []database.ChatMessage{ - {ID: 1, Role: database.ChatMessageRoleUser, Visibility: database.ChatMessageVisibilityBoth, APIKeyID: sqlNullString(oldKeyID)}, - {ID: 2, Role: database.ChatMessageRoleUser, Visibility: database.ChatMessageVisibilityModel, Compressed: true}, - }, - }, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - t.Parallel() - - gotKey, gotOK := activeTurnAPIKeyIDFromMessages(tt.messages) - require.Equal(t, tt.wantOK, gotOK) - require.Equal(t, tt.wantKey, gotKey) - }) - } -} - -func TestPromptMessagesForVisibleUserPreserveActiveAPIKeyID(t *testing.T) { - t.Parallel() - - db, _ := dbtestutil.NewDB(t) - ctx := t.Context() - user := dbgen.User(t, db, database.User{}) - org := dbgen.Organization(t, db, database.Organization{}) - model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{}) - chat := dbgen.Chat(t, db, database.Chat{OrganizationID: org.ID, OwnerID: user.ID, LastModelConfigID: model.ID}) - oldKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID}) - currentKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID}) - modelOnlyKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID}) - - dbgen.ChatMessage(t, db, database.ChatMessage{ - ChatID: chat.ID, - CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true}, - Role: database.ChatMessageRoleUser, - Visibility: database.ChatMessageVisibilityBoth, - APIKeyID: sqlNullString(oldKey.ID), - }) - dbgen.ChatMessage(t, db, database.ChatMessage{ - ChatID: chat.ID, - CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true}, - Role: database.ChatMessageRoleSystem, - Visibility: database.ChatMessageVisibilityModel, - Compressed: true, - }) - dbgen.ChatMessage(t, db, database.ChatMessage{ - ChatID: chat.ID, - CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true}, - Role: database.ChatMessageRoleUser, - Visibility: database.ChatMessageVisibilityBoth, - APIKeyID: sqlNullString(currentKey.ID), - }) - dbgen.ChatMessage(t, db, database.ChatMessage{ - ChatID: chat.ID, - CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true}, - Role: database.ChatMessageRoleUser, - Visibility: database.ChatMessageVisibilityModel, - APIKeyID: sqlNullString(modelOnlyKey.ID), - }) - - messages, err := db.GetChatMessagesForPromptByChatID(ctx, chat.ID) - require.NoError(t, err) - gotKey, ok := activeTurnAPIKeyIDFromMessages(messages) - require.True(t, ok) - require.Equal(t, currentKey.ID, gotKey) -} - -func TestPromptMessagesForCompactedChatPreserveActiveAPIKeyID(t *testing.T) { - t.Parallel() - - db, _ := dbtestutil.NewDB(t) - ctx := t.Context() - user := dbgen.User(t, db, database.User{}) - org := dbgen.Organization(t, db, database.Organization{}) - model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{}) - chat := dbgen.Chat(t, db, database.Chat{OrganizationID: org.ID, OwnerID: user.ID, LastModelConfigID: model.ID}) - key, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID}) - - visibleUser := dbgen.ChatMessage(t, db, database.ChatMessage{ - ChatID: chat.ID, - CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, - ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true}, - Role: database.ChatMessageRoleUser, - Visibility: database.ChatMessageVisibilityBoth, - APIKeyID: sqlNullString(key.ID), - }) - dbgen.ChatMessage(t, db, database.ChatMessage{ - ChatID: chat.ID, - ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true}, - Role: database.ChatMessageRoleAssistant, - Visibility: database.ChatMessageVisibilityBoth, - }) - compressedSummary := dbgen.ChatMessage(t, db, database.ChatMessage{ - ChatID: chat.ID, - ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true}, - Role: database.ChatMessageRoleUser, - Visibility: database.ChatMessageVisibilityModel, - Compressed: true, - APIKeyID: sqlNullString(key.ID), - }) - afterSummary := dbgen.ChatMessage(t, db, database.ChatMessage{ - ChatID: chat.ID, - ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true}, - Role: database.ChatMessageRoleAssistant, - Visibility: database.ChatMessageVisibilityBoth, - }) - - messages, err := db.GetChatMessagesForPromptByChatID(ctx, chat.ID) - require.NoError(t, err) - - ids := make(map[int64]struct{}, len(messages)) - for _, message := range messages { - ids[message.ID] = struct{}{} - } - _, hasVisibleUser := ids[visibleUser.ID] - require.False(t, hasVisibleUser) - _, hasSummary := ids[compressedSummary.ID] - require.True(t, hasSummary) - _, hasAfterSummary := ids[afterSummary.ID] - require.True(t, hasAfterSummary) - - gotKey, ok := activeTurnAPIKeyIDFromMessages(messages) - require.True(t, ok) - require.Equal(t, key.ID, gotKey) -} - func sqlNullString(value string) sql.NullString { return sql.NullString{String: value, Valid: value != ""} } diff --git a/coderd/x/chatd/quickgen.go b/coderd/x/chatd/quickgen.go index 4657a95d006..36c102a1cac 100644 --- a/coderd/x/chatd/quickgen.go +++ b/coderd/x/chatd/quickgen.go @@ -221,11 +221,16 @@ func (p *Server) GenerateChatTitleAsync(ctx context.Context, chat database.Chat) titleCtx, stopTitleCtx := p.inflightContext(ctx) if err := p.goInflight(func() { defer stopTitleCtx() - modelOpts := modelBuildOptionsFromMessages(messages) - turnCtx := withActiveTurnAPIKeyID(titleCtx, modelOpts) + apiKeyID, err := p.ensureSyntheticAPIKeyID(titleCtx, chat.OwnerID) + if err != nil { + logger.Debug(titleCtx, "failed to ensure synthetic API key for automatic title generation", slog.Error(err)) + return + } + modelOpts := modelBuildOptions{ActiveAPIKeyID: apiKeyID} + turnCtx := titleCtx model, modelConfig, route, _, _, _, err := p.resolveChatModel(turnCtx, chat, modelOpts) if err != nil { - logger.Debug(turnCtx, "failed to resolve model for automatic title generation", + logger.Debug(titleCtx, "failed to resolve model for automatic title generation", slog.Error(err), ) return diff --git a/coderd/x/chatd/subagent.go b/coderd/x/chatd/subagent.go index aae832132d6..f2c50e3023c 100644 --- a/coderd/x/chatd/subagent.go +++ b/coderd/x/chatd/subagent.go @@ -17,7 +17,6 @@ import ( "golang.org/x/xerrors" "cdr.dev/slog/v3" - "github.com/coder/coder/v2/coderd/aibridge" "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/database/dbauthz" dbpubsub "github.com/coder/coder/v2/coderd/database/pubsub" @@ -1024,23 +1023,6 @@ func (p *Server) resolveExploreToolSnapshot( return inheritedMCPServerIDs, nil } -func (*Server) delegatedAPIKeyIDForSubagent(ctx context.Context) (string, error) { - apiKeyID, ok := aibridge.DelegatedAPIKeyIDFromContext(ctx) - if !ok || apiKeyID == "" { - return "", xerrors.New("active turn API key ID is required for subagent messages") - } - return apiKeyID, nil -} - -func (p *Server) createChildSubagentChat( - ctx context.Context, - parent database.Chat, - prompt string, - title string, -) (database.Chat, error) { - return p.createChildSubagentChatWithOptions(ctx, parent, prompt, title, childSubagentChatOptions{}) -} - func (p *Server) createChildSubagentChatWithOptions( ctx context.Context, parent database.Chat, @@ -1074,9 +1056,9 @@ func (p *Server) createChildSubagentChatWithOptions( if modelConfigID == uuid.Nil { return database.Chat{}, xerrors.New("model config is required") } - childAPIKeyID, err := p.delegatedAPIKeyIDForSubagent(ctx) + childAPIKeyID, err := p.ensureSyntheticAPIKeyID(ctx, parent.OwnerID) if err != nil { - return database.Chat{}, err + return database.Chat{}, xerrors.Errorf("ensure synthetic API key: %w", err) } childPlanMode := parent.PlanMode @@ -1218,16 +1200,10 @@ func (p *Server) sendSubagentMessage( return database.Chat{}, xerrors.Errorf("get target chat: %w", err) } - apiKeyID, err := p.delegatedAPIKeyIDForSubagent(ctx) - if err != nil { - return database.Chat{}, err - } - sendResult, err := p.SendMessage(ctx, SendMessageOptions{ ChatID: targetChatID, CreatedBy: targetChat.OwnerID, Content: []codersdk.ChatMessagePart{codersdk.ChatMessageText(message)}, - APIKeyID: apiKeyID, BusyBehavior: busyBehavior, }) if err != nil { diff --git a/coderd/x/chatd/subagent_context_internal_test.go b/coderd/x/chatd/subagent_context_internal_test.go index 9ca4d4b7955..6e20b969bc4 100644 --- a/coderd/x/chatd/subagent_context_internal_test.go +++ b/coderd/x/chatd/subagent_context_internal_test.go @@ -10,7 +10,6 @@ import ( "github.com/google/uuid" "github.com/stretchr/testify/require" - "github.com/coder/coder/v2/coderd/aibridge" "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/database/dbgen" "github.com/coder/coder/v2/coderd/database/dbtestutil" @@ -72,7 +71,6 @@ func createWorkspaceBoundParentChat( parent, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "parent-with-context", ModelConfigID: model.ID, WorkspaceID: uuid.NullUUID{UUID: ws.ID, Valid: true}, @@ -114,7 +112,6 @@ func TestSpawnComputerUseAgentInheritsPinnedContext(t *testing.T) { // (OpenAI only) that was cached before the Anthropic provider was inserted. server.configCache.InvalidateProviders() - ctx = aibridge.WithDelegatedAPIKeyID(ctx, testAPIKeyID(t, db, parentChat.OwnerID)) tools := server.subagentTools(ctx, func() database.Chat { return parentChat }, parentChat.LastModelConfigID) tool := findToolByName(tools, spawnAgentToolName) require.NotNil(t, tool) diff --git a/coderd/x/chatd/subagent_internal_test.go b/coderd/x/chatd/subagent_internal_test.go index 4f90342e013..24112baf65a 100644 --- a/coderd/x/chatd/subagent_internal_test.go +++ b/coderd/x/chatd/subagent_internal_test.go @@ -216,7 +216,6 @@ func seedInternalChatDeps( t.Helper() user := dbgen.User(t, db, database.User{}) - _ = testAPIKeyID(t, db, user.ID) org := dbgen.Organization(t, db, database.Organization{}) dbgen.OrganizationMember(t, db, database.OrganizationMember{ UserID: user.ID, @@ -268,222 +267,6 @@ func insertInternalAIProvider( }) } -func TestCreateChildSubagentChatPropagatesActiveTurnAPIKeyID(t *testing.T) { - t.Parallel() - - ctx := testutil.Context(t, testutil.WaitShort) - db, _ := dbtestutil.NewDB(t) - user := dbgen.User(t, db, database.User{}) - org := dbgen.Organization(t, db, database.Organization{}) - dbgen.OrganizationMember(t, db, database.OrganizationMember{UserID: user.ID, OrganizationID: org.ID}) - model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{}) - parent := dbgen.Chat(t, db, database.Chat{ - OrganizationID: org.ID, - OwnerID: user.ID, - LastModelConfigID: model.ID, - }) - - apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID}) - ctx = aibridge.WithDelegatedAPIKeyID(ctx, apiKey.ID) - - server := &Server{db: db, logger: slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})} - child, err := server.createChildSubagentChat(ctx, parent, "inspect the workspace", "") - require.NoError(t, err) - - messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{ChatID: child.ID}) - require.NoError(t, err) - var childUserMessage database.ChatMessage - for _, message := range messages { - if message.Role == database.ChatMessageRoleUser { - childUserMessage = message - break - } - } - require.NotZero(t, childUserMessage.ID) - require.True(t, childUserMessage.APIKeyID.Valid) - require.Equal(t, apiKey.ID, childUserMessage.APIKeyID.String) -} - -func TestSendSubagentMessagePropagatesActiveTurnAPIKeyID(t *testing.T) { - t.Parallel() - - db, ps := dbtestutil.NewDB(t) - server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{}) - ctx := chatdTestContext(t) - user, org, model := seedInternalChatDeps(t, db) - apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID}) - - parent, err := server.CreateChat(ctx, CreateOptions{ - OrganizationID: org.ID, - OwnerID: user.ID, - Title: "parent-send-subagent-key", - ModelConfigID: model.ID, - InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, - APIKeyID: apiKey.ID, - }) - require.NoError(t, err) - child, err := server.CreateChat(ctx, CreateOptions{ - OrganizationID: org.ID, - OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), - ParentChatID: uuid.NullUUID{UUID: parent.ID, Valid: true}, - RootChatID: uuid.NullUUID{UUID: parent.ID, Valid: true}, - Title: "child-send-subagent-key", - ModelConfigID: model.ID, - InitialUserContent: []codersdk.ChatMessagePart{ - codersdk.ChatMessageText("do work"), - }, - }) - require.NoError(t, err) - - setChatStatus(ctx, t, db, child.ID, database.ChatStatusWaiting, "") - - ctx = aibridge.WithDelegatedAPIKeyID(ctx, apiKey.ID) - _, err = server.sendSubagentMessage( - ctx, - parent.ID, - child.ID, - "follow up", - SendMessageBusyBehaviorInterrupt, - ) - require.NoError(t, err) - - messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{ChatID: child.ID}) - require.NoError(t, err) - var latestUserMessage database.ChatMessage - for _, message := range messages { - if message.Role == database.ChatMessageRoleUser && message.ID > latestUserMessage.ID { - latestUserMessage = message - } - } - require.NotZero(t, latestUserMessage.ID) - require.True(t, latestUserMessage.APIKeyID.Valid) - require.Equal(t, apiKey.ID, latestUserMessage.APIKeyID.String) -} - -func TestCreateChildSubagentChatRequiresActiveTurnAPIKeyIDForAIGateway(t *testing.T) { - t.Parallel() - - ctx := testutil.Context(t, testutil.WaitShort) - db, _ := dbtestutil.NewDB(t) - user := dbgen.User(t, db, database.User{}) - org := dbgen.Organization(t, db, database.Organization{}) - dbgen.OrganizationMember(t, db, database.OrganizationMember{UserID: user.ID, OrganizationID: org.ID}) - model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{}) - parent := dbgen.Chat(t, db, database.Chat{ - OrganizationID: org.ID, - OwnerID: user.ID, - LastModelConfigID: model.ID, - }) - - server := &Server{ - db: db, - logger: slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}), - } - _, err := server.createChildSubagentChat(ctx, parent, "inspect the workspace", "") - require.ErrorContains(t, err, "active turn API key ID is required for subagent messages") -} - -func TestSendSubagentMessageRequiresActiveTurnAPIKeyIDForAIGateway(t *testing.T) { - t.Parallel() - - db, ps := dbtestutil.NewDB(t) - server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{}) - ctx := chatdTestContext(t) - user, org, model := seedInternalChatDeps(t, db) - - parent, err := server.CreateChat(ctx, CreateOptions{ - OrganizationID: org.ID, - OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), - Title: "parent-send-subagent-missing-key", - ModelConfigID: model.ID, - InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, - }) - require.NoError(t, err) - child, err := server.CreateChat(ctx, CreateOptions{ - OrganizationID: org.ID, - OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), - ParentChatID: uuid.NullUUID{UUID: parent.ID, Valid: true}, - RootChatID: uuid.NullUUID{UUID: parent.ID, Valid: true}, - Title: "child-send-subagent-missing-key", - ModelConfigID: model.ID, - InitialUserContent: []codersdk.ChatMessagePart{ - codersdk.ChatMessageText("do work"), - }, - }) - require.NoError(t, err) - - setChatStatus(ctx, t, db, child.ID, database.ChatStatusWaiting, "") - _, err = server.sendSubagentMessage( - ctx, - parent.ID, - child.ID, - "follow up", - SendMessageBusyBehaviorInterrupt, - ) - require.ErrorContains(t, err, "active turn API key ID is required for subagent messages") -} - -// TestSpawnAgentUsesActiveTurnAPIKeyIDFromContext verifies that, with AI -// Gateway routing enabled, the spawn_agent tool succeeds when the active -// turn's delegated API key ID is present on the context and fails without -// it. The generation worker supplies that key by enriching the tool -// execution context with withActiveTurnAPIKeyID, derived from the prompt -// rows' model build options. This guards the regression where -// executeLocalTools passed an un-enriched context to tool callbacks, -// breaking subagent spawning under AI Gateway routing. -func TestSpawnAgentUsesActiveTurnAPIKeyIDFromContext(t *testing.T) { - t.Parallel() - - db, ps := dbtestutil.NewDB(t) - server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{}) - - ctx := chatdTestContext(t) - user, org, model := seedInternalChatDeps(t, db) - apiKey, _ := dbgen.APIKey(t, db, database.APIKey{UserID: user.ID}) - - parent, err := server.CreateChat(ctx, CreateOptions{ - OrganizationID: org.ID, - OwnerID: user.ID, - Title: "parent-active-turn-key", - ModelConfigID: model.ID, - InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, - APIKeyID: apiKey.ID, - }) - require.NoError(t, err) - parentChat, err := db.GetChatByID(ctx, parent.ID) - require.NoError(t, err) - - // The generation worker derives model build options from the prompt - // rows; this is the source executeLocalTools uses to enrich the tool - // execution context. - promptRows, err := server.db.GetChatMessagesForPromptByChatID(ctx, parentChat.ID) - require.NoError(t, err) - modelOpts := modelBuildOptionsFromMessages(promptRows) - require.Equal(t, apiKey.ID, modelOpts.ActiveAPIKeyID) - - // Without the delegated key on the context the spawn fails, matching - // the original un-enriched executeLocalTools behavior. - resp := runSpawnAgentTool(ctx, t, server, parentChat, spawnAgentArgs{ - Type: subagentTypeGeneral, - Prompt: "delegate work", - }) - require.True(t, resp.IsError, "expected error without active turn key, got: %s", resp.Content) - require.Contains(t, resp.Content, "active turn API key ID is required for subagent messages") - - // With the key on the context (as withActiveTurnAPIKeyID supplies in - // executeLocalTools), the spawn succeeds. - enrichedCtx := withActiveTurnAPIKeyID(ctx, modelOpts) - resp = runSpawnAgentTool(enrichedCtx, t, server, parentChat, spawnAgentArgs{ - Type: subagentTypeGeneral, - Prompt: "delegate work", - }) - result := requireSpawnAgentResponse(t, resp) - require.Equal(t, subagentTypeGeneral, result.SubagentType) -} - func TestResolveUserProviderAPIKeys_AIProvider(t *testing.T) { t.Parallel() @@ -868,6 +651,83 @@ func upsertInternalUserChatPersonalModelOverride( ) } +func TestCreateChildSubagentChatPersistsOwnerSyntheticAPIKeyID(t *testing.T) { + t.Parallel() + + db, ps := dbtestutil.NewDB(t) + server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{}) + + ctx := chatdTestContext(t) + user, org, model := seedInternalChatDeps(t, db) + parent := createInternalParentChat( + ctx, t, server, db, org.ID, user.ID, model.ID, "parent-child-key", + ) + + child, err := server.createChildSubagentChatWithOptions( + ctx, + parent, + "inspect the workspace", + "", + childSubagentChatOptions{}, + ) + require.NoError(t, err) + + mapping, err := db.GetChatSyntheticAPIKeyByUserID(ctx, user.ID) + require.NoError(t, err) + messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{ + ChatID: child.ID, + AfterID: 0, + }) + require.NoError(t, err) + for _, message := range messages { + if message.Role != database.ChatMessageRoleUser { + continue + } + require.True(t, message.APIKeyID.Valid) + require.Equal(t, mapping.APIKeyID, message.APIKeyID.String) + return + } + require.Fail(t, "child user message not found") +} + +func TestSendSubagentMessagePersistsOwnerSyntheticAPIKeyID(t *testing.T) { + t.Parallel() + + db, ps := dbtestutil.NewDB(t) + server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{}) + + ctx := chatdTestContext(t) + user, org, model := seedInternalChatDeps(t, db) + parent, child := createParentChildChats(ctx, t, server, user, org, model) + setChatStatus(ctx, t, db, child.ID, database.ChatStatusWaiting, "") + + _, err := server.sendSubagentMessage( + ctx, + parent.ID, + child.ID, + "follow up", + SendMessageBusyBehaviorInterrupt, + ) + require.NoError(t, err) + + mapping, err := db.GetChatSyntheticAPIKeyByUserID(ctx, user.ID) + require.NoError(t, err) + messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{ + ChatID: child.ID, + AfterID: 0, + }) + require.NoError(t, err) + var latestUserMessage database.ChatMessage + for _, message := range messages { + if message.Role == database.ChatMessageRoleUser && message.ID > latestUserMessage.ID { + latestUserMessage = message + } + } + require.NotZero(t, latestUserMessage.ID) + require.True(t, latestUserMessage.APIKeyID.Valid) + require.Equal(t, mapping.APIKeyID, latestUserMessage.APIKeyID.String) +} + func TestCreateChildSubagentChatInheritsWorkspaceBinding(t *testing.T) { t.Parallel() @@ -881,7 +741,6 @@ func TestCreateChildSubagentChatInheritsWorkspaceBinding(t *testing.T) { parent, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), WorkspaceID: uuid.NullUUID{ UUID: workspace.ID, Valid: true, @@ -903,7 +762,6 @@ func TestCreateChildSubagentChatInheritsWorkspaceBinding(t *testing.T) { parentChat, err := db.GetChatByID(ctx, parent.ID) require.NoError(t, err) - ctx = aibridge.WithDelegatedAPIKeyID(ctx, testAPIKeyID(t, server.db, parentChat.OwnerID)) child, err := server.createChildSubagentChatWithOptions(ctx, parentChat, "inspect bindings", "", childSubagentChatOptions{}) require.NoError(t, err) @@ -930,7 +788,6 @@ func createInternalParentChat( parent, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: orgID, OwnerID: userID, - APIKeyID: testAPIKeyID(t, db, userID), Title: title, ModelConfigID: modelConfigID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -943,15 +800,6 @@ func createInternalParentChat( return parentChat } -// withSubagentDelegatedKey enriches ctx with a delegated API key ID for -// subagent tool callbacks. AI Gateway routing requires this key on the -// context; tests that do not otherwise set it should call this helper -// before invoking runSpawnAgentTool or runSubagentTool with spawn_agent. -func withSubagentDelegatedKey(ctx context.Context, t *testing.T, db database.Store, ownerID uuid.UUID) context.Context { - t.Helper() - return aibridge.WithDelegatedAPIKeyID(ctx, testAPIKeyID(t, db, ownerID)) -} - func runSubagentTool( ctx context.Context, t *testing.T, @@ -1064,7 +912,6 @@ func TestCreateChildSubagentChatCopiesPlanMode(t *testing.T) { parent, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "plan-parent", ModelConfigID: model.ID, PlanMode: planMode, @@ -1078,7 +925,6 @@ func TestCreateChildSubagentChatCopiesPlanMode(t *testing.T) { require.NoError(t, err) require.Equal(t, planMode, parentChat.PlanMode) - ctx = aibridge.WithDelegatedAPIKeyID(ctx, testAPIKeyID(t, server.db, parentChat.OwnerID)) child, err := server.createChildSubagentChatWithOptions(ctx, parentChat, "inspect bindings", "", childSubagentChatOptions{}) require.NoError(t, err) @@ -1099,7 +945,6 @@ func TestSpawnAgent_GeneralInheritsParentModelWhenOmitted(t *testing.T) { ctx, t, server, db, org.ID, user.ID, model.ID, "parent-inherited-model", ) - ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID) resp := runSpawnAgentTool(ctx, t, server, parentChat, spawnAgentArgs{ Type: subagentTypeGeneral, Prompt: "delegate work", @@ -1130,7 +975,6 @@ func TestSpawnAgent_GeneralUsesConfiguredModelOverride(t *testing.T) { ctx, t, server, db, org.ID, user.ID, model.ID, "parent-general-override", ) - ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID) resp := runSpawnAgentTool(ctx, t, server, parentChat, spawnAgentArgs{ Type: subagentTypeGeneral, Prompt: "delegate general work", @@ -1323,7 +1167,6 @@ func TestSpawnAgent_GeneralHonorsPersonalModelOverrides(t *testing.T) { "parent-general-personal-override", ) - ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID) resp := runSpawnAgentTool(ctx, t, server, parentChat, spawnAgentArgs{ Type: subagentTypeGeneral, Prompt: "delegate general work", @@ -1374,7 +1217,6 @@ func TestSpawnAgent_GeneralOverrideLogsAndFallsBackWhenCredentialsUnavailable(t parent, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "parent-general-credentials-fallback", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{ @@ -1385,7 +1227,6 @@ func TestSpawnAgent_GeneralOverrideLogsAndFallsBackWhenCredentialsUnavailable(t parentChat, err := db.GetChatByID(ctx, parent.ID) require.NoError(t, err) - ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID) resp := runSpawnAgentTool(ctx, t, server, parentChat, spawnAgentArgs{ Type: subagentTypeGeneral, Prompt: "inspect provider credentials", @@ -1445,7 +1286,6 @@ func TestSpawnAgent_GeneralOverrideLogsAndFallsBackWhenProviderDisabled(t *testi parent, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "parent-general-disabled-provider-fallback", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{ @@ -1456,7 +1296,6 @@ func TestSpawnAgent_GeneralOverrideLogsAndFallsBackWhenProviderDisabled(t *testi parentChat, err := db.GetChatByID(ctx, parent.ID) require.NoError(t, err) - ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID) resp := runSpawnAgentTool(ctx, t, server, parentChat, spawnAgentArgs{ Type: subagentTypeGeneral, Prompt: "inspect disabled providers", @@ -1579,7 +1418,6 @@ func TestCreateChildSubagentChat_StoresReasoningEffortOverride(t *testing.T) { parentChat := createInternalParentChat( ctx, t, server, db, org.ID, user.ID, model.ID, "parent-effort-override", ) - ctx = aibridge.WithDelegatedAPIKeyID(ctx, testAPIKeyID(t, server.db, parentChat.OwnerID)) child, err := server.createChildSubagentChatWithOptions( ctx, parentChat, @@ -1612,7 +1450,6 @@ func TestCreateChildSubagentChat_OverrideWorksWhenParentHasNoModel(t *testing.T) // The chats table enforces a foreign key for last_model_config_id, so // use a synthetic parent value here to exercise the override path. parentChat.LastModelConfigID = uuid.Nil - ctx = aibridge.WithDelegatedAPIKeyID(ctx, testAPIKeyID(t, server.db, parentChat.OwnerID)) child, err := server.createChildSubagentChatWithOptions( ctx, parentChat, @@ -1643,7 +1480,6 @@ func TestSpawnAgent_ExploreUsesConfiguredModelOverride(t *testing.T) { ctx, t, server, db, org.ID, user.ID, model.ID, "parent-explore-override", ) - ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID) resp := runSubagentTool( ctx, t, @@ -1681,7 +1517,6 @@ func TestSpawnAgent_ExploreFallsBackToCurrentTurnModel(t *testing.T) { ctx, t, server, db, org.ID, user.ID, parentModel.ID, "parent-explore-fallback", ) - ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID) resp := runSubagentTool( ctx, t, @@ -1886,7 +1721,6 @@ func TestSpawnAgent_ExploreHonorsPersonalModelOverrides(t *testing.T) { "parent-explore-personal-override", ) - ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID) resp := runSubagentTool( ctx, t, @@ -1924,7 +1758,6 @@ func TestCreateChat_ExploreRootStartsWithoutMCPSnapshot(t *testing.T) { root, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "root-explore", ModelConfigID: model.ID, ChatMode: database.NullChatMode{ @@ -2043,7 +1876,6 @@ func TestCreateChildSubagentChatWithOptions_ExplorePersistsMCPSnapshot(t *testin t, db, user.ID, "snapshot-"+uuid.NewString(), false, ) - ctx = aibridge.WithDelegatedAPIKeyID(ctx, testAPIKeyID(t, server.db, parentChat.OwnerID)) child, err := server.createChildSubagentChatWithOptions( ctx, parentChat, @@ -2082,7 +1914,6 @@ func TestSpawnAgent_ExploreSnapshotsTurnStateParentState(t *testing.T) { parent, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "parent-turn-state-snapshot", ModelConfigID: model.ID, MCPServerIDs: []uuid.UUID{turnStartConfig.ID}, @@ -2095,7 +1926,6 @@ func TestSpawnAgent_ExploreSnapshotsTurnStateParentState(t *testing.T) { turnParent, err := db.GetChatByID(ctx, parent.ID) require.NoError(t, err) - ctx = aibridge.WithDelegatedAPIKeyID(ctx, testAPIKeyID(t, db, user.ID)) tools := server.subagentTools( ctx, func() database.Chat { return turnParent }, @@ -2163,7 +1993,6 @@ func TestSpawnAgent_ExploreFallsBackOnInvalidUUID(t *testing.T) { ctx, t, server, db, org.ID, user.ID, parentModel.ID, "parent-explore-invalid-override", ) - ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID) resp := runSubagentTool( ctx, t, @@ -2199,7 +2028,6 @@ func TestSpawnAgent_ExploreFallsBackWhenOverrideIsUnavailable(t *testing.T) { ctx, t, server, db, org.ID, user.ID, parentModel.ID, "parent-explore-disabled", ) - ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID) resp := runSubagentTool( ctx, t, @@ -2247,7 +2075,6 @@ func TestSpawnAgent_ExploreFallsBackWhenOverrideCredentialsAreUnavailable(t *tes ctx, t, server, db, org.ID, user.ID, parentModel.ID, "parent-explore-missing-user-key", ) - ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID) resp := runSubagentTool( ctx, t, @@ -2356,7 +2183,6 @@ func TestSpawnAgent_PlanModeDescriptionOmitsComputerUse(t *testing.T) { parent, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "plan-parent-description", ModelConfigID: model.ID, PlanMode: database.NullChatPlanMode{ @@ -2396,7 +2222,6 @@ func TestSpawnAgent_PlanModeRejectsComputerUse(t *testing.T) { parent, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "plan-parent-computer-use-reject", ModelConfigID: model.ID, PlanMode: database.NullChatPlanMode{ @@ -2504,7 +2329,6 @@ func TestSpawnAgent_ComputerUseRejectsMissingConfiguredProvider(t *testing.T) { server := newInternalTestServer(t, db, ps, chatprovider.ProviderAPIKeys{}) user := dbgen.User(t, db, database.User{}) - _ = testAPIKeyID(t, db, user.ID) org := dbgen.Organization(t, db, database.Organization{}) dbgen.OrganizationMember(t, db, database.OrganizationMember{ UserID: user.ID, @@ -2688,7 +2512,6 @@ func TestSpawnAgent_NotAvailableForExploreChats(t *testing.T) { exploreChat, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "root-explore", ModelConfigID: model.ID, ChatMode: database.NullChatMode{ @@ -2750,7 +2573,6 @@ func TestSubagentLifecycleToolsIncludePersistedSubagentTypeAcrossVariants(t *tes model.ID, "parent-lifecycle-"+tt.variant, ) - ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID) spawnResp := runSpawnAgentTool(ctx, t, server, parentChat, spawnAgentArgs{ Type: tt.variant, @@ -2812,7 +2634,6 @@ func TestSubagentLifecycleToolErrorsIncludePersistedSubagentType(t *testing.T) { unrelated, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "unrelated-lifecycle-parent", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("other")}, @@ -2885,7 +2706,6 @@ func TestSpawnAgent_ComputerUseUsesComputerUseModelNotParent(t *testing.T) { parent, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), WorkspaceID: uuid.NullUUID{UUID: workspace.ID, Valid: true}, BuildID: uuid.NullUUID{UUID: build.ID, Valid: true}, AgentID: uuid.NullUUID{UUID: agent.ID, Valid: true}, @@ -2898,7 +2718,6 @@ func TestSpawnAgent_ComputerUseUsesComputerUseModelNotParent(t *testing.T) { parentChat, err := db.GetChatByID(ctx, parent.ID) require.NoError(t, err) - ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID) resp := runSubagentTool( ctx, t, @@ -2952,7 +2771,6 @@ func TestSpawnAgent_ComputerUseInheritsMCPServerIDs(t *testing.T) { parent, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "parent-cu-mcp", ModelConfigID: model.ID, MCPServerIDs: parentMCPIDs, @@ -2963,7 +2781,6 @@ func TestSpawnAgent_ComputerUseInheritsMCPServerIDs(t *testing.T) { parentChat, err := db.GetChatByID(ctx, parent.ID) require.NoError(t, err) - ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID) resp := runSubagentTool( ctx, t, @@ -3014,7 +2831,6 @@ func TestCreateChildSubagentChat_InheritsMCPServerIDs(t *testing.T) { parent, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "parent-with-mcp", ModelConfigID: model.ID, MCPServerIDs: parentMCPIDs, @@ -3029,7 +2845,6 @@ func TestCreateChildSubagentChat_InheritsMCPServerIDs(t *testing.T) { "parent chat must have the MCP server IDs we set") // Spawn a child subagent chat. - ctx = aibridge.WithDelegatedAPIKeyID(ctx, testAPIKeyID(t, server.db, parentChat.OwnerID)) child, err := server.createChildSubagentChatWithOptions( ctx, parentChat, @@ -3059,7 +2874,6 @@ func TestCreateChildSubagentChat_NoMCPServersStaysEmpty(t *testing.T) { parent, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "parent-no-mcp", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -3070,7 +2884,6 @@ func TestCreateChildSubagentChat_NoMCPServersStaysEmpty(t *testing.T) { require.NoError(t, err) // Spawn a child. - ctx = aibridge.WithDelegatedAPIKeyID(ctx, testAPIKeyID(t, server.db, parentChat.OwnerID)) child, err := server.createChildSubagentChatWithOptions( ctx, parentChat, @@ -3099,7 +2912,6 @@ func TestIsSubagentDescendant(t *testing.T) { root, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "root", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("root")}, @@ -3109,7 +2921,6 @@ func TestIsSubagentDescendant(t *testing.T) { child, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ParentChatID: uuid.NullUUID{ UUID: root.ID, Valid: true, @@ -3127,7 +2938,6 @@ func TestIsSubagentDescendant(t *testing.T) { grandchild, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ParentChatID: uuid.NullUUID{ UUID: child.ID, Valid: true, @@ -3146,7 +2956,6 @@ func TestIsSubagentDescendant(t *testing.T) { unrelated, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "unrelated-root", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("unrelated")}, @@ -3156,7 +2965,6 @@ func TestIsSubagentDescendant(t *testing.T) { unrelatedChild, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ParentChatID: uuid.NullUUID{ UUID: unrelated.ID, Valid: true, @@ -3247,7 +3055,6 @@ func createParentChildChats( parent, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, server.db, user.ID), Title: "parent-" + t.Name(), ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -3257,7 +3064,6 @@ func createParentChildChats( child, err = server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, server.db, user.ID), ParentChatID: uuid.NullUUID{ UUID: parent.ID, Valid: true, @@ -3499,7 +3305,6 @@ func TestAwaitSubagentCompletion(t *testing.T) { unrelated, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "unrelated", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("other")}, @@ -4002,7 +3807,6 @@ func TestListAgents(t *testing.T) { parent, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: title, ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -4015,7 +3819,6 @@ func TestListAgents(t *testing.T) { child, err := server.CreateChat(ctx, CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), ParentChatID: uuid.NullUUID{UUID: parent.ID, Valid: true}, RootChatID: uuid.NullUUID{UUID: parent.ID, Valid: true}, Title: title, diff --git a/coderd/x/chatd/subagent_test.go b/coderd/x/chatd/subagent_test.go index f621fa579ef..a768f3487ee 100644 --- a/coderd/x/chatd/subagent_test.go +++ b/coderd/x/chatd/subagent_test.go @@ -27,7 +27,6 @@ func TestSpawnComputerUseAgent_CreatesChildWithChatMode(t *testing.T) { parent, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "parent", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -41,7 +40,6 @@ func TestSpawnComputerUseAgent_CreatesChildWithChatMode(t *testing.T) { child, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: parent.OwnerID, - APIKeyID: testAPIKeyID(t, db, parent.OwnerID), ParentChatID: uuid.NullUUID{ UUID: parent.ID, Valid: true, @@ -84,7 +82,6 @@ func TestSpawnComputerUseAgent_SystemPromptFormat(t *testing.T) { parent, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "parent", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -97,7 +94,6 @@ func TestSpawnComputerUseAgent_SystemPromptFormat(t *testing.T) { child, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: parent.OwnerID, - APIKeyID: testAPIKeyID(t, db, parent.OwnerID), ParentChatID: uuid.NullUUID{ UUID: parent.ID, Valid: true, @@ -145,7 +141,6 @@ func TestSpawnComputerUseAgent_ChildIsListedUnderParent(t *testing.T) { parent, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "parent", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -157,7 +152,6 @@ func TestSpawnComputerUseAgent_ChildIsListedUnderParent(t *testing.T) { child, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: parent.OwnerID, - APIKeyID: testAPIKeyID(t, db, parent.OwnerID), ParentChatID: uuid.NullUUID{ UUID: parent.ID, Valid: true, @@ -193,7 +187,6 @@ func TestSpawnComputerUseAgent_RootChatIDPropagation(t *testing.T) { parent, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: user.ID, - APIKeyID: testAPIKeyID(t, db, user.ID), Title: "root-parent", ModelConfigID: model.ID, InitialUserContent: []codersdk.ChatMessagePart{codersdk.ChatMessageText("hello")}, @@ -205,7 +198,6 @@ func TestSpawnComputerUseAgent_RootChatIDPropagation(t *testing.T) { child, err := server.CreateChat(ctx, chatd.CreateOptions{ OrganizationID: org.ID, OwnerID: parent.OwnerID, - APIKeyID: testAPIKeyID(t, db, parent.OwnerID), ParentChatID: uuid.NullUUID{ UUID: parent.ID, Valid: true, diff --git a/coderd/x/chatd/synthetickey.go b/coderd/x/chatd/synthetickey.go new file mode 100644 index 00000000000..eb0248cbc63 --- /dev/null +++ b/coderd/x/chatd/synthetickey.go @@ -0,0 +1,121 @@ +package chatd + +import ( + "context" + "database/sql" + "fmt" + "time" + + "github.com/google/uuid" + "golang.org/x/xerrors" + + "github.com/coder/coder/v2/coderd/apikey" + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/database/dbauthz" +) + +const ( + syntheticAPIKeyLifetime = 30 * 24 * time.Hour + syntheticAPIKeyRenewMargin = 24 * time.Hour + syntheticAPIKeyMaxAttempts = 3 +) + +func (p *Server) ensureSyntheticAPIKeyID(ctx context.Context, ownerID uuid.UUID) (string, error) { + ctx = dbauthz.AsChatdKeyMinter(ctx, ownerID) + for range syntheticAPIKeyMaxAttempts { + mapping, err := p.db.GetChatSyntheticAPIKeyByUserID(ctx, ownerID) + if err == nil { + key, keyErr := p.db.GetAPIKeyByID(ctx, mapping.APIKeyID) + switch { + case keyErr == nil && key.ExpiresAt.After(p.clock.Now().Add(syntheticAPIKeyRenewMargin)): + return key.ID, nil + case keyErr != nil && !xerrors.Is(keyErr, sql.ErrNoRows): + return "", xerrors.Errorf("get synthetic API key: %w", keyErr) + } + } else if !xerrors.Is(err, sql.ErrNoRows) { + return "", xerrors.Errorf("get synthetic API key mapping: %w", err) + } + + keyID, retry, err := p.mintSyntheticAPIKey(ctx, ownerID) + if err != nil { + return "", err + } + if !retry { + return keyID, nil + } + } + return "", xerrors.New("ensure synthetic API key: concurrent update retry limit reached") +} + +func (p *Server) mintSyntheticAPIKey(ctx context.Context, ownerID uuid.UUID) (keyID string, retry bool, err error) { + err = p.db.InTx(func(tx database.Store) error { + mapping, mappingErr := tx.GetChatSyntheticAPIKeyByUserID(ctx, ownerID) + hasMapping := mappingErr == nil + if mappingErr != nil && !xerrors.Is(mappingErr, sql.ErrNoRows) { + return xerrors.Errorf("get synthetic API key mapping: %w", mappingErr) + } + if hasMapping { + key, keyErr := tx.GetAPIKeyByID(ctx, mapping.APIKeyID) + if keyErr == nil && key.ExpiresAt.After(p.clock.Now().Add(syntheticAPIKeyRenewMargin)) { + keyID = key.ID + return nil + } + if keyErr != nil { + if xerrors.Is(keyErr, sql.ErrNoRows) { + retry = true + return nil + } + return xerrors.Errorf("get synthetic API key: %w", keyErr) + } + } + + owner, err := tx.GetUserForChatSyntheticAPIKeyByID(ctx, ownerID) + if err != nil { + return xerrors.Errorf("get synthetic API key owner: %w", err) + } + params, _, err := apikey.Generate(apikey.CreateParams{ + UserID: ownerID, + LoginType: owner.LoginType, + ExpiresAt: p.clock.Now().Add(syntheticAPIKeyLifetime), + LifetimeSeconds: int64(syntheticAPIKeyLifetime.Seconds()), + TokenName: fmt.Sprintf("%s_chat_gateway_key", ownerID), + }) + if err != nil { + return xerrors.Errorf("generate synthetic API key: %w", err) + } + key, err := tx.InsertAPIKey(ctx, params) + if err != nil { + return xerrors.Errorf("insert synthetic API key: %w", err) + } + + var rows int64 + if hasMapping { + rows, err = tx.UpdateChatSyntheticAPIKey(ctx, database.UpdateChatSyntheticAPIKeyParams{ + UserID: ownerID, + OldApiKeyID: mapping.APIKeyID, + NewApiKeyID: key.ID, + }) + } else { + rows, err = tx.InsertChatSyntheticAPIKey(ctx, database.InsertChatSyntheticAPIKeyParams{ + UserID: ownerID, + APIKeyID: key.ID, + }) + } + if err != nil { + return xerrors.Errorf("publish synthetic API key mapping: %w", err) + } + if rows == 1 { + keyID = key.ID + return nil + } + if err := tx.DeleteAPIKeyByID(ctx, key.ID); err != nil { + return xerrors.Errorf("delete losing synthetic API key: %w", err) + } + retry = true + return nil + }, nil) + if err != nil { + return "", false, err + } + return keyID, retry, nil +} diff --git a/coderd/x/chatd/synthetickey_internal_test.go b/coderd/x/chatd/synthetickey_internal_test.go new file mode 100644 index 00000000000..986479952cf --- /dev/null +++ b/coderd/x/chatd/synthetickey_internal_test.go @@ -0,0 +1,176 @@ +package chatd + +import ( + "context" + "database/sql" + "encoding/json" + "sync" + "testing" + "time" + + "github.com/google/uuid" + "github.com/stretchr/testify/require" + + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/database/dbgen" + "github.com/coder/coder/v2/coderd/database/dbtestutil" + "github.com/coder/quartz" +) + +func TestSyntheticAPIKeyLifecycle(t *testing.T) { + t.Parallel() + + db, _ := dbtestutil.NewDB(t) + user := dbgen.User(t, db, database.User{}) + server := &Server{db: db, clock: quartz.NewReal()} + + firstID, err := server.ensureSyntheticAPIKeyID(t.Context(), user.ID) + require.NoError(t, err) + secondID, err := server.ensureSyntheticAPIKeyID(t.Context(), user.ID) + require.NoError(t, err) + require.Equal(t, firstID, secondID) + + first, err := db.GetAPIKeyByID(t.Context(), firstID) + require.NoError(t, err) + require.Equal(t, user.LoginType, first.LoginType) + require.WithinDuration(t, server.clock.Now().Add(syntheticAPIKeyLifetime), first.ExpiresAt, time.Second) + + err = db.UpdateAPIKeyByID(t.Context(), database.UpdateAPIKeyByIDParams{ + ID: first.ID, + LastUsed: first.LastUsed, + ExpiresAt: server.clock.Now().Add(syntheticAPIKeyRenewMargin), + IPAddress: first.IPAddress, + }) + require.NoError(t, err) + + renewedID, err := server.ensureSyntheticAPIKeyID(t.Context(), user.ID) + require.NoError(t, err) + require.NotEqual(t, firstID, renewedID) + _, err = db.GetAPIKeyByID(t.Context(), firstID) + require.NoError(t, err, "renewal must preserve the previous key") + + require.NoError(t, db.DeleteAPIKeyByID(t.Context(), renewedID)) + _, err = db.GetChatSyntheticAPIKeyByUserID(t.Context(), user.ID) + require.ErrorIs(t, err, sql.ErrNoRows) + + recreatedID, err := server.ensureSyntheticAPIKeyID(t.Context(), user.ID) + require.NoError(t, err) + require.NotEqual(t, renewedID, recreatedID) +} + +func TestSyntheticAPIKeyConcurrentMint(t *testing.T) { + t.Parallel() + + db, _ := dbtestutil.NewDB(t) + user := dbgen.User(t, db, database.User{}) + server := &Server{db: db, clock: quartz.NewReal()} + + const workers = 8 + ids := make([]string, workers) + errs := make([]error, workers) + var wg sync.WaitGroup + for i := range workers { + wg.Add(1) + go func() { + defer wg.Done() + ids[i], errs[i] = server.ensureSyntheticAPIKeyID(t.Context(), user.ID) + }() + } + wg.Wait() + + for i := range workers { + require.NoError(t, errs[i]) + require.Equal(t, ids[0], ids[i]) + } + keys, err := db.GetAPIKeysByUserID(t.Context(), database.GetAPIKeysByUserIDParams{ + LoginType: user.LoginType, + UserID: user.ID, + IncludeExpired: true, + }) + require.NoError(t, err) + require.Len(t, keys, 1) +} + +func TestSyntheticAPIKeyDeletionDoesNotMutateChatState(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + deleteKey func(context.Context, database.Store, uuid.UUID, string) error + }{ + { + name: "individual", + deleteKey: func(ctx context.Context, db database.Store, _ uuid.UUID, keyID string) error { + return db.DeleteAPIKeyByID(ctx, keyID) + }, + }, + { + name: "all user keys", + deleteKey: func(ctx context.Context, db database.Store, userID uuid.UUID, _ string) error { + return db.DeleteAPIKeysByUserID(ctx, userID) + }, + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + db, _ := dbtestutil.NewDB(t) + user := dbgen.User(t, db, database.User{}) + org := dbgen.Organization(t, db, database.Organization{}) + model := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{}) + server := &Server{db: db, clock: quartz.NewReal()} + syntheticID, err := server.ensureSyntheticAPIKeyID(t.Context(), user.ID) + require.NoError(t, err) + + chat := dbgen.Chat(t, db, database.Chat{ + OrganizationID: org.ID, + OwnerID: user.ID, + LastModelConfigID: model.ID, + }) + message := dbgen.ChatMessage(t, db, database.ChatMessage{ + ChatID: chat.ID, + CreatedBy: uuid.NullUUID{UUID: user.ID, Valid: true}, + ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true}, + Role: database.ChatMessageRoleUser, + APIKeyID: sql.NullString{String: syntheticID, Valid: true}, + }) + queued, err := db.InsertChatQueuedMessage(t.Context(), database.InsertChatQueuedMessageParams{ + ChatID: chat.ID, + Content: json.RawMessage(`[]`), + ModelConfigID: uuid.NullUUID{UUID: model.ID, Valid: true}, + APIKeyID: sql.NullString{String: syntheticID, Valid: true}, + }) + require.NoError(t, err) + + before, err := db.GetChatByID(t.Context(), chat.ID) + require.NoError(t, err) + require.NoError(t, test.deleteKey(t.Context(), db, user.ID, syntheticID)) + + _, err = db.GetAPIKeyByID(t.Context(), syntheticID) + require.ErrorIs(t, err, sql.ErrNoRows) + _, err = db.GetChatSyntheticAPIKeyByUserID(t.Context(), user.ID) + require.ErrorIs(t, err, sql.ErrNoRows) + + after, err := db.GetChatByID(t.Context(), chat.ID) + require.NoError(t, err) + require.Equal(t, before.HistoryVersion, after.HistoryVersion) + require.Equal(t, before.QueueVersion, after.QueueVersion) + require.Equal(t, before.GenerationAttempt, after.GenerationAttempt) + + stored, err := db.GetChatMessageByID(t.Context(), message.ID) + require.NoError(t, err) + require.Equal(t, sql.NullString{String: syntheticID, Valid: true}, stored.APIKeyID) + storedQueued, err := db.GetChatQueuedMessageByID(t.Context(), database.GetChatQueuedMessageByIDParams{ + ID: queued.ID, + ChatID: chat.ID, + }) + require.NoError(t, err) + require.Equal(t, sql.NullString{String: syntheticID, Valid: true}, storedQueued.APIKeyID) + + remintedID, err := server.ensureSyntheticAPIKeyID(t.Context(), user.ID) + require.NoError(t, err) + require.NotEqual(t, syntheticID, remintedID) + }) + } +} diff --git a/coderd/x/chatd/title_override_internal_test.go b/coderd/x/chatd/title_override_internal_test.go index cdd301455f9..0c6e385792e 100644 --- a/coderd/x/chatd/title_override_internal_test.go +++ b/coderd/x/chatd/title_override_internal_test.go @@ -10,6 +10,7 @@ import ( "strings" "sync/atomic" "testing" + "time" "charm.land/fantasy" fantasyopenai "charm.land/fantasy/providers/openai" @@ -551,107 +552,75 @@ func TestResolveManualTitleModel_TitleGenerationOverrideMissingCredentials(t *te require.Equal(t, database.ChatModelConfig{}, gotConfig) } -func TestGenerateManualTitleCandidate_ActiveAPIKeyIDFallback(t *testing.T) { +func TestGenerateManualTitleCandidate_UsesSyntheticAPIKey(t *testing.T) { t.Parallel() - contextAPIKeyID := uuid.NewString() - messageAPIKeyID := uuid.NewString() - shadowedContextAPIKeyID := uuid.NewString() - tests := []struct { - name string - messageAPIKeyID string - contextAPIKeyID string - wantAPIKeyID string - wantErrContains string - }{ - { - name: "ContextFallback", - contextAPIKeyID: contextAPIKeyID, - wantAPIKeyID: contextAPIKeyID, - }, - { - name: "MessageTakesPrecedence", - messageAPIKeyID: messageAPIKeyID, - contextAPIKeyID: shadowedContextAPIKeyID, - wantAPIKeyID: messageAPIKeyID, - }, - { - name: "NoKeyAnywhereFailsClosed", - wantErrContains: "AI Gateway routing requires the active turn API key ID", - }, + ctx := testutil.Context(t, testutil.WaitShort) + ctrl := gomock.NewController(t) + db := dbmock.NewMockStore(ctrl) + logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) + chat, messages := titleOverrideTestChatAndMessages(t) + chat.OrganizationID = uuid.New() + overrideConfig := titleOverrideModelConfig("gpt-4.1", true) + providerID := uuid.New() + overrideConfig.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true} + provider := database.AIProvider{ + ID: providerID, + Name: "primary-openai", + Type: database.AIProviderTypeOpenai, + Enabled: true, } + apiKeyID := uuid.NewString() + wantTitle := "Synthetic title" + seenAPIKeyID := make(chan string, 1) + factory := &aibridgeTestFactory{rt: roundTripFunc(func(req *http.Request) (*http.Response, error) { + delegatedID, _ := aibridge.DelegatedAPIKeyIDFromContext(req.Context()) + seenAPIKeyID <- delegatedID + text := strconv.Quote(`{"title":"` + wantTitle + `"}`) + body := `{"id":"resp_test","object":"response","created_at":0,"status":"completed","model":"gpt-4.1","output":[{"id":"msg_test","type":"message","role":"assistant","content":[{"type":"output_text","text":` + text + `}]}],"usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}` + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(body)), + Request: req, + }, nil + })} - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - t.Parallel() + db.EXPECT().GetChatUsageLimitConfig(gomock.Any()).Return(database.ChatUsageLimitConfig{}, sql.ErrNoRows) + db.EXPECT().GetChatMessagesByChatIDAscPaginated(gomock.Any(), database.GetChatMessagesByChatIDAscPaginatedParams{ + ChatID: chat.ID, + AfterID: 0, + LimitVal: manualTitleMessageWindowLimit, + }).Return(messages, nil) + db.EXPECT().GetChatMessagesByChatIDDescPaginated(gomock.Any(), database.GetChatMessagesByChatIDDescPaginatedParams{ + ChatID: chat.ID, + BeforeID: 0, + LimitVal: manualTitleMessageWindowLimit, + }).Return(nil, nil) + db.EXPECT().GetChatSyntheticAPIKeyByUserID(gomock.Any(), chat.OwnerID).Return(database.ChatSyntheticApiKey{ + UserID: chat.OwnerID, + APIKeyID: apiKeyID, + }, nil) + db.EXPECT().GetAPIKeyByID(gomock.Any(), apiKeyID).Return(database.APIKey{ + ID: apiKeyID, + UserID: chat.OwnerID, + ExpiresAt: time.Now().Add(48 * time.Hour), + }, nil) + db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil) + db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil) + db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes() + db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ + ProviderID: providerID, + APIKey: "test-key", + }}, nil).AnyTimes() - ctx := testutil.Context(t, testutil.WaitShort) - if tt.contextAPIKeyID != "" { - ctx = aibridge.WithDelegatedAPIKeyID(ctx, tt.contextAPIKeyID) - } - ctrl := gomock.NewController(t) - db := dbmock.NewMockStore(ctrl) - logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}) - chat, messages := titleOverrideTestChatAndMessages(t) - chat.OrganizationID = uuid.New() - if tt.messageAPIKeyID != "" { - messages[0] = withChatMessageAPIKeyID(messages[0], tt.messageAPIKeyID) - } - overrideConfig := titleOverrideModelConfig("gpt-4.1", true) - providerID := uuid.New() - overrideConfig.AIProviderID = uuid.NullUUID{UUID: providerID, Valid: true} - provider := database.AIProvider{ - ID: providerID, - Name: "primary-openai", - Type: database.AIProviderTypeOpenai, - Enabled: true, - } - wantTitle := "Context title" - seenAPIKeyID := make(chan string, 1) - factory := &aibridgeTestFactory{rt: roundTripFunc(func(req *http.Request) (*http.Response, error) { - apiKeyID, _ := aibridge.DelegatedAPIKeyIDFromContext(req.Context()) - seenAPIKeyID <- apiKeyID - text := strconv.Quote(`{"title":"` + wantTitle + `"}`) - body := `{"id":"resp_test","object":"response","created_at":0,"status":"completed","model":"gpt-4.1","output":[{"id":"msg_test","type":"message","role":"assistant","content":[{"type":"output_text","text":` + text + `}]}],"usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}` - return &http.Response{ - StatusCode: http.StatusOK, - Header: http.Header{"Content-Type": []string{"application/json"}}, - Body: io.NopCloser(strings.NewReader(body)), - Request: req, - }, nil - })} - - db.EXPECT().GetChatUsageLimitConfig(gomock.Any()).Return(database.ChatUsageLimitConfig{}, sql.ErrNoRows) - db.EXPECT().GetChatMessagesByChatIDAscPaginated(gomock.Any(), database.GetChatMessagesByChatIDAscPaginatedParams{ - ChatID: chat.ID, - AfterID: 0, - LimitVal: manualTitleMessageWindowLimit, - }).Return(messages, nil) - db.EXPECT().GetChatMessagesByChatIDDescPaginated(gomock.Any(), database.GetChatMessagesByChatIDDescPaginatedParams{ - ChatID: chat.ID, - BeforeID: 0, - LimitVal: manualTitleMessageWindowLimit, - }).Return(nil, nil) - db.EXPECT().GetChatTitleGenerationModelOverride(gomock.Any()).Return(overrideConfig.ID.String(), nil) - db.EXPECT().GetChatModelConfigByID(gomock.Any(), overrideConfig.ID).Return(overrideConfig, nil) - db.EXPECT().GetAIProviderByID(gomock.Any(), providerID).Return(provider, nil).AnyTimes() - db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ - ProviderID: providerID, - APIKey: "test-key", - }}, nil).AnyTimes() - - server := titleOverrideTestServer(db, logger) - server.aibridgeTransportFactory = aibridgeTestFactoryPointer(factory) - title, err := server.generateManualTitleCandidate(ctx, db, chat) - if tt.wantErrContains != "" { - require.ErrorContains(t, err, tt.wantErrContains) - return - } - require.NoError(t, err) - require.Equal(t, wantTitle, title) - require.Equal(t, tt.wantAPIKeyID, testutil.RequireReceive(ctx, t, seenAPIKeyID)) - }) - } + server := titleOverrideTestServer(db, logger) + server.clock = quartz.NewReal() + server.aibridgeTransportFactory = aibridgeTestFactoryPointer(factory) + title, err := server.generateManualTitleCandidate(ctx, db, chat) + require.NoError(t, err) + require.Equal(t, wantTitle, title) + require.Equal(t, apiKeyID, testutil.RequireReceive(ctx, t, seenAPIKeyID)) } func TestResolveManualTitleModel_TitleGenerationOverrideSetUnusable(t *testing.T) { From fa35022754d5dbaa9bcbfcae5e2fe697d3147c3a Mon Sep 17 00:00:00 2001 From: Michael Suchacz <203725896+ibetitsmike@users.noreply.github.com> Date: Mon, 13 Jul 2026 18:59:14 +0000 Subject: [PATCH 2/2] refactor(coderd): drop chat synthetic key mapping table Resolve the synthetic gateway key by deterministic token name (chatd__session_token) instead of a mapping table, mirroring the provisionerd session token model. The lookup excludes login_type 'token' rows so a user-created token with the colliding name is never picked up or extended. Near-expiry keys are extended in place under a per-user advisory lock, keeping the key ID stable for in-flight generations. New keys carry a minimal scope as defense in depth. --- coderd/database/dbauthz/dbauthz.go | 25 +--- coderd/database/dbauthz/dbauthz_test.go | 24 +--- coderd/database/dbmetrics/querymetrics.go | 32 ++--- coderd/database/dbmock/dbmock.go | 60 ++------ coderd/database/dump.sql | 19 --- coderd/database/foreign_key_constraint.go | 2 - .../000546_chat_synthetic_api_keys.up.sql | 12 -- ...46_drop_chat_history_api_key_fks.down.sql} | 2 - ...00546_drop_chat_history_api_key_fks.up.sql | 5 + coderd/database/migrations/migrate_test.go | 4 +- .../000546_chat_synthetic_api_keys.up.sql | 5 - coderd/database/models.go | 7 - coderd/database/querier.go | 4 +- coderd/database/queries.sql.go | 104 ++++++-------- coderd/database/queries/apikeys.sql | 18 +++ coderd/database/queries/chats.sql | 17 --- coderd/database/unique_constraint.go | 2 - coderd/exp_chats_test.go | 7 +- coderd/x/chatd/ARCHITECTURE.md | 6 +- coderd/x/chatd/chatd_internal_test.go | 6 +- coderd/x/chatd/chatd_test.go | 59 +++++--- coderd/x/chatd/subagent_internal_test.go | 14 +- coderd/x/chatd/synthetickey.go | 130 +++++++++--------- coderd/x/chatd/synthetickey_internal_test.go | 105 ++++++++++++-- .../x/chatd/title_override_internal_test.go | 9 +- 25 files changed, 327 insertions(+), 351 deletions(-) delete mode 100644 coderd/database/migrations/000546_chat_synthetic_api_keys.up.sql rename coderd/database/migrations/{000546_chat_synthetic_api_keys.down.sql => 000546_drop_chat_history_api_key_fks.down.sql} (94%) create mode 100644 coderd/database/migrations/000546_drop_chat_history_api_key_fks.up.sql delete mode 100644 coderd/database/migrations/testdata/fixtures/000546_chat_synthetic_api_keys.up.sql diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index f04b6502286..84b212e45d3 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -3342,6 +3342,10 @@ func (q *querier) GetChatFilesByIDs(ctx context.Context, ids []uuid.UUID) ([]dat return files, nil } +func (q *querier) GetChatGatewayAPIKey(ctx context.Context, arg database.GetChatGatewayAPIKeyParams) (database.APIKey, error) { + return fetch(q.log, q.auth, q.db.GetChatGatewayAPIKey)(ctx, arg) +} + func (q *querier) GetChatGeneralModelOverride(ctx context.Context) (string, error) { if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceDeploymentConfig); err != nil { return "", err @@ -3521,13 +3525,6 @@ func (q *querier) GetChatStreamSyncRows(ctx context.Context, ids []uuid.UUID) ([ return q.db.GetChatStreamSyncRows(ctx, ids) } -func (q *querier) GetChatSyntheticAPIKeyByUserID(ctx context.Context, userID uuid.UUID) (database.ChatSyntheticApiKey, error) { - if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceApiKey.WithOwner(userID.String())); err != nil { - return database.ChatSyntheticApiKey{}, err - } - return q.db.GetChatSyntheticAPIKeyByUserID(ctx, userID) -} - func (q *querier) GetChatSystemPrompt(ctx context.Context) (string, error) { // The system prompt is a deployment-wide setting read during chat // creation by every authenticated user, so no RBAC policy check @@ -6049,13 +6046,6 @@ func (q *querier) InsertChatQueuedMessageWithCreator(ctx context.Context, arg da return q.db.InsertChatQueuedMessageWithCreator(ctx, arg) } -func (q *querier) InsertChatSyntheticAPIKey(ctx context.Context, arg database.InsertChatSyntheticAPIKeyParams) (int64, error) { - if err := q.authorizeContext(ctx, policy.ActionCreate, rbac.ResourceApiKey.WithOwner(arg.UserID.String())); err != nil { - return 0, err - } - return q.db.InsertChatSyntheticAPIKey(ctx, arg) -} - func (q *querier) InsertCryptoKey(ctx context.Context, arg database.InsertCryptoKeyParams) (database.CryptoKey, error) { if err := q.authorizeContext(ctx, policy.ActionCreate, rbac.ResourceCryptoKey); err != nil { return database.CryptoKey{}, err @@ -7447,13 +7437,6 @@ func (q *querier) UpdateChatStatus(ctx context.Context, arg database.UpdateChatS return q.db.UpdateChatStatus(ctx, arg) } -func (q *querier) UpdateChatSyntheticAPIKey(ctx context.Context, arg database.UpdateChatSyntheticAPIKeyParams) (int64, error) { - if err := q.authorizeContext(ctx, policy.ActionUpdate, rbac.ResourceApiKey.WithOwner(arg.UserID.String())); err != nil { - return 0, err - } - return q.db.UpdateChatSyntheticAPIKey(ctx, arg) -} - func (q *querier) UpdateChatTitleByID(ctx context.Context, arg database.UpdateChatTitleByIDParams) (database.Chat, error) { chat, err := q.db.GetChatByID(ctx, arg.ID) if err != nil { diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index 7b4ad1f8f81..0e1dedede6a 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -339,29 +339,17 @@ func defaultIPAddress() pqtype.Inet { } } -func (s *MethodTestSuite) TestChatSyntheticAPIKey() { +func (s *MethodTestSuite) TestChatGatewayAPIKey() { s.Run("GetUserForChatSyntheticAPIKeyByID", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { user := testutil.Fake(s.T(), faker, database.User{}) dbm.EXPECT().GetUserForChatSyntheticAPIKeyByID(gomock.Any(), user.ID).Return(user, nil).AnyTimes() check.Args(user.ID).Asserts(user, policy.ActionReadPersonal).Returns(user) })) - s.Run("GetChatSyntheticAPIKeyByUserID", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { - user := testutil.Fake(s.T(), faker, database.User{}) - mapping := testutil.Fake(s.T(), faker, database.ChatSyntheticApiKey{UserID: user.ID}) - dbm.EXPECT().GetChatSyntheticAPIKeyByUserID(gomock.Any(), user.ID).Return(mapping, nil).AnyTimes() - check.Args(user.ID).Asserts(rbac.ResourceApiKey.WithOwner(user.ID.String()), policy.ActionRead).Returns(mapping) - })) - s.Run("InsertChatSyntheticAPIKey", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { - user := testutil.Fake(s.T(), faker, database.User{}) - arg := database.InsertChatSyntheticAPIKeyParams{UserID: user.ID, APIKeyID: uuid.NewString()} - dbm.EXPECT().InsertChatSyntheticAPIKey(gomock.Any(), arg).Return(int64(1), nil).AnyTimes() - check.Args(arg).Asserts(rbac.ResourceApiKey.WithOwner(user.ID.String()), policy.ActionCreate).Returns(int64(1)) - })) - s.Run("UpdateChatSyntheticAPIKey", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { - user := testutil.Fake(s.T(), faker, database.User{}) - arg := database.UpdateChatSyntheticAPIKeyParams{UserID: user.ID, OldApiKeyID: uuid.NewString(), NewApiKeyID: uuid.NewString()} - dbm.EXPECT().UpdateChatSyntheticAPIKey(gomock.Any(), arg).Return(int64(1), nil).AnyTimes() - check.Args(arg).Asserts(rbac.ResourceApiKey.WithOwner(user.ID.String()), policy.ActionUpdate).Returns(int64(1)) + s.Run("GetChatGatewayAPIKey", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + key := testutil.Fake(s.T(), faker, database.APIKey{}) + arg := database.GetChatGatewayAPIKeyParams{UserID: key.UserID, TokenName: key.TokenName} + dbm.EXPECT().GetChatGatewayAPIKey(gomock.Any(), arg).Return(key, nil).AnyTimes() + check.Args(arg).Asserts(key, policy.ActionRead).Returns(key) })) } diff --git a/coderd/database/dbmetrics/querymetrics.go b/coderd/database/dbmetrics/querymetrics.go index 307d6908a15..798449a37e3 100644 --- a/coderd/database/dbmetrics/querymetrics.go +++ b/coderd/database/dbmetrics/querymetrics.go @@ -1601,6 +1601,14 @@ func (m queryMetricsStore) GetChatFilesByIDs(ctx context.Context, ids []uuid.UUI return r0, r1 } +func (m queryMetricsStore) GetChatGatewayAPIKey(ctx context.Context, arg database.GetChatGatewayAPIKeyParams) (database.APIKey, error) { + start := time.Now() + r0, r1 := m.s.GetChatGatewayAPIKey(ctx, arg) + m.queryLatencies.WithLabelValues("GetChatGatewayAPIKey").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetChatGatewayAPIKey").Inc() + return r0, r1 +} + func (m queryMetricsStore) GetChatGeneralModelOverride(ctx context.Context) (string, error) { start := time.Now() r0, r1 := m.s.GetChatGeneralModelOverride(ctx) @@ -1769,14 +1777,6 @@ func (m queryMetricsStore) GetChatStreamSyncRows(ctx context.Context, ids []uuid return r0, r1 } -func (m queryMetricsStore) GetChatSyntheticAPIKeyByUserID(ctx context.Context, userID uuid.UUID) (database.ChatSyntheticApiKey, error) { - start := time.Now() - r0, r1 := m.s.GetChatSyntheticAPIKeyByUserID(ctx, userID) - m.queryLatencies.WithLabelValues("GetChatSyntheticAPIKeyByUserID").Observe(time.Since(start).Seconds()) - m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetChatSyntheticAPIKeyByUserID").Inc() - return r0, r1 -} - func (m queryMetricsStore) GetChatSystemPrompt(ctx context.Context) (string, error) { start := time.Now() r0, r1 := m.s.GetChatSystemPrompt(ctx) @@ -4137,14 +4137,6 @@ func (m queryMetricsStore) InsertChatQueuedMessageWithCreator(ctx context.Contex return r0, r1 } -func (m queryMetricsStore) InsertChatSyntheticAPIKey(ctx context.Context, arg database.InsertChatSyntheticAPIKeyParams) (int64, error) { - start := time.Now() - r0, r1 := m.s.InsertChatSyntheticAPIKey(ctx, arg) - m.queryLatencies.WithLabelValues("InsertChatSyntheticAPIKey").Observe(time.Since(start).Seconds()) - m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "InsertChatSyntheticAPIKey").Inc() - return r0, r1 -} - func (m queryMetricsStore) InsertCryptoKey(ctx context.Context, arg database.InsertCryptoKeyParams) (database.CryptoKey, error) { start := time.Now() r0, r1 := m.s.InsertCryptoKey(ctx, arg) @@ -5289,14 +5281,6 @@ func (m queryMetricsStore) UpdateChatStatus(ctx context.Context, arg database.Up return r0, r1 } -func (m queryMetricsStore) UpdateChatSyntheticAPIKey(ctx context.Context, arg database.UpdateChatSyntheticAPIKeyParams) (int64, error) { - start := time.Now() - r0, r1 := m.s.UpdateChatSyntheticAPIKey(ctx, arg) - m.queryLatencies.WithLabelValues("UpdateChatSyntheticAPIKey").Observe(time.Since(start).Seconds()) - m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "UpdateChatSyntheticAPIKey").Inc() - return r0, r1 -} - func (m queryMetricsStore) UpdateChatTitleByID(ctx context.Context, arg database.UpdateChatTitleByIDParams) (database.Chat, error) { start := time.Now() r0, r1 := m.s.UpdateChatTitleByID(ctx, arg) diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index 64422a6e3d3..b6d7a0d1a24 100644 --- a/coderd/database/dbmock/dbmock.go +++ b/coderd/database/dbmock/dbmock.go @@ -2953,6 +2953,21 @@ func (mr *MockStoreMockRecorder) GetChatFilesByIDs(ctx, ids any) *gomock.Call { return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatFilesByIDs", reflect.TypeOf((*MockStore)(nil).GetChatFilesByIDs), ctx, ids) } +// GetChatGatewayAPIKey mocks base method. +func (m *MockStore) GetChatGatewayAPIKey(ctx context.Context, arg database.GetChatGatewayAPIKeyParams) (database.APIKey, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetChatGatewayAPIKey", ctx, arg) + ret0, _ := ret[0].(database.APIKey) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetChatGatewayAPIKey indicates an expected call of GetChatGatewayAPIKey. +func (mr *MockStoreMockRecorder) GetChatGatewayAPIKey(ctx, arg any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatGatewayAPIKey", reflect.TypeOf((*MockStore)(nil).GetChatGatewayAPIKey), ctx, arg) +} + // GetChatGeneralModelOverride mocks base method. func (m *MockStore) GetChatGeneralModelOverride(ctx context.Context) (string, error) { m.ctrl.T.Helper() @@ -3268,21 +3283,6 @@ func (mr *MockStoreMockRecorder) GetChatStreamSyncRows(ctx, ids any) *gomock.Cal return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatStreamSyncRows", reflect.TypeOf((*MockStore)(nil).GetChatStreamSyncRows), ctx, ids) } -// GetChatSyntheticAPIKeyByUserID mocks base method. -func (m *MockStore) GetChatSyntheticAPIKeyByUserID(ctx context.Context, userID uuid.UUID) (database.ChatSyntheticApiKey, error) { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "GetChatSyntheticAPIKeyByUserID", ctx, userID) - ret0, _ := ret[0].(database.ChatSyntheticApiKey) - ret1, _ := ret[1].(error) - return ret0, ret1 -} - -// GetChatSyntheticAPIKeyByUserID indicates an expected call of GetChatSyntheticAPIKeyByUserID. -func (mr *MockStoreMockRecorder) GetChatSyntheticAPIKeyByUserID(ctx, userID any) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatSyntheticAPIKeyByUserID", reflect.TypeOf((*MockStore)(nil).GetChatSyntheticAPIKeyByUserID), ctx, userID) -} - // GetChatSystemPrompt mocks base method. func (m *MockStore) GetChatSystemPrompt(ctx context.Context) (string, error) { m.ctrl.T.Helper() @@ -7750,21 +7750,6 @@ func (mr *MockStoreMockRecorder) InsertChatQueuedMessageWithCreator(ctx, arg any return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "InsertChatQueuedMessageWithCreator", reflect.TypeOf((*MockStore)(nil).InsertChatQueuedMessageWithCreator), ctx, arg) } -// InsertChatSyntheticAPIKey mocks base method. -func (m *MockStore) InsertChatSyntheticAPIKey(ctx context.Context, arg database.InsertChatSyntheticAPIKeyParams) (int64, error) { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "InsertChatSyntheticAPIKey", ctx, arg) - ret0, _ := ret[0].(int64) - ret1, _ := ret[1].(error) - return ret0, ret1 -} - -// InsertChatSyntheticAPIKey indicates an expected call of InsertChatSyntheticAPIKey. -func (mr *MockStoreMockRecorder) InsertChatSyntheticAPIKey(ctx, arg any) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "InsertChatSyntheticAPIKey", reflect.TypeOf((*MockStore)(nil).InsertChatSyntheticAPIKey), ctx, arg) -} - // InsertCryptoKey mocks base method. func (m *MockStore) InsertCryptoKey(ctx context.Context, arg database.InsertCryptoKeyParams) (database.CryptoKey, error) { m.ctrl.T.Helper() @@ -9965,21 +9950,6 @@ func (mr *MockStoreMockRecorder) UpdateChatStatus(ctx, arg any) *gomock.Call { return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateChatStatus", reflect.TypeOf((*MockStore)(nil).UpdateChatStatus), ctx, arg) } -// UpdateChatSyntheticAPIKey mocks base method. -func (m *MockStore) UpdateChatSyntheticAPIKey(ctx context.Context, arg database.UpdateChatSyntheticAPIKeyParams) (int64, error) { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "UpdateChatSyntheticAPIKey", ctx, arg) - ret0, _ := ret[0].(int64) - ret1, _ := ret[1].(error) - return ret0, ret1 -} - -// UpdateChatSyntheticAPIKey indicates an expected call of UpdateChatSyntheticAPIKey. -func (mr *MockStoreMockRecorder) UpdateChatSyntheticAPIKey(ctx, arg any) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateChatSyntheticAPIKey", reflect.TypeOf((*MockStore)(nil).UpdateChatSyntheticAPIKey), ctx, arg) -} - // UpdateChatTitleByID mocks base method. func (m *MockStore) UpdateChatTitleByID(ctx context.Context, arg database.UpdateChatTitleByIDParams) (database.Chat, error) { m.ctrl.T.Helper() diff --git a/coderd/database/dump.sql b/coderd/database/dump.sql index e36ad5bb568..2313ebf1df8 100644 --- a/coderd/database/dump.sql +++ b/coderd/database/dump.sql @@ -2033,13 +2033,6 @@ CREATE SEQUENCE chat_queued_messages_id_seq ALTER SEQUENCE chat_queued_messages_id_seq OWNED BY chat_queued_messages.id; -CREATE TABLE chat_synthetic_api_keys ( - user_id uuid NOT NULL, - api_key_id text NOT NULL, - created_at timestamp with time zone DEFAULT now() NOT NULL, - updated_at timestamp with time zone DEFAULT now() NOT NULL -); - CREATE TABLE chat_usage_limit_config ( id bigint NOT NULL, singleton boolean DEFAULT true NOT NULL, @@ -4308,12 +4301,6 @@ ALTER TABLE ONLY chat_model_configs ALTER TABLE ONLY chat_queued_messages ADD CONSTRAINT chat_queued_messages_pkey PRIMARY KEY (id); -ALTER TABLE ONLY chat_synthetic_api_keys - ADD CONSTRAINT chat_synthetic_api_keys_api_key_id_key UNIQUE (api_key_id); - -ALTER TABLE ONLY chat_synthetic_api_keys - ADD CONSTRAINT chat_synthetic_api_keys_pkey PRIMARY KEY (user_id); - ALTER TABLE ONLY chat_usage_limit_config ADD CONSTRAINT chat_usage_limit_config_pkey PRIMARY KEY (id); @@ -5170,12 +5157,6 @@ ALTER TABLE ONLY chat_model_configs ALTER TABLE ONLY chat_queued_messages ADD CONSTRAINT chat_queued_messages_chat_id_fkey FOREIGN KEY (chat_id) REFERENCES chats(id) ON DELETE CASCADE; -ALTER TABLE ONLY chat_synthetic_api_keys - ADD CONSTRAINT chat_synthetic_api_keys_api_key_id_fkey FOREIGN KEY (api_key_id) REFERENCES api_keys(id) ON DELETE CASCADE; - -ALTER TABLE ONLY chat_synthetic_api_keys - ADD CONSTRAINT chat_synthetic_api_keys_user_id_fkey FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE; - ALTER TABLE ONLY chats ADD CONSTRAINT chats_agent_id_fkey FOREIGN KEY (agent_id) REFERENCES workspace_agents(id) ON DELETE SET NULL; diff --git a/coderd/database/foreign_key_constraint.go b/coderd/database/foreign_key_constraint.go index 94bb3f0ed22..75c99671c63 100644 --- a/coderd/database/foreign_key_constraint.go +++ b/coderd/database/foreign_key_constraint.go @@ -30,8 +30,6 @@ const ( ForeignKeyChatModelConfigsCreatedBy ForeignKeyConstraint = "chat_model_configs_created_by_fkey" // ALTER TABLE ONLY chat_model_configs ADD CONSTRAINT chat_model_configs_created_by_fkey FOREIGN KEY (created_by) REFERENCES users(id); ForeignKeyChatModelConfigsUpdatedBy ForeignKeyConstraint = "chat_model_configs_updated_by_fkey" // ALTER TABLE ONLY chat_model_configs ADD CONSTRAINT chat_model_configs_updated_by_fkey FOREIGN KEY (updated_by) REFERENCES users(id); ForeignKeyChatQueuedMessagesChatID ForeignKeyConstraint = "chat_queued_messages_chat_id_fkey" // ALTER TABLE ONLY chat_queued_messages ADD CONSTRAINT chat_queued_messages_chat_id_fkey FOREIGN KEY (chat_id) REFERENCES chats(id) ON DELETE CASCADE; - ForeignKeyChatSyntheticAPIKeysAPIKeyID ForeignKeyConstraint = "chat_synthetic_api_keys_api_key_id_fkey" // ALTER TABLE ONLY chat_synthetic_api_keys ADD CONSTRAINT chat_synthetic_api_keys_api_key_id_fkey FOREIGN KEY (api_key_id) REFERENCES api_keys(id) ON DELETE CASCADE; - ForeignKeyChatSyntheticAPIKeysUserID ForeignKeyConstraint = "chat_synthetic_api_keys_user_id_fkey" // ALTER TABLE ONLY chat_synthetic_api_keys ADD CONSTRAINT chat_synthetic_api_keys_user_id_fkey FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE; ForeignKeyChatsAgentID ForeignKeyConstraint = "chats_agent_id_fkey" // ALTER TABLE ONLY chats ADD CONSTRAINT chats_agent_id_fkey FOREIGN KEY (agent_id) REFERENCES workspace_agents(id) ON DELETE SET NULL; ForeignKeyChatsBuildID ForeignKeyConstraint = "chats_build_id_fkey" // ALTER TABLE ONLY chats ADD CONSTRAINT chats_build_id_fkey FOREIGN KEY (build_id) REFERENCES workspace_builds(id) ON DELETE SET NULL; ForeignKeyChatsLastModelConfigID ForeignKeyConstraint = "chats_last_model_config_id_fkey" // ALTER TABLE ONLY chats ADD CONSTRAINT chats_last_model_config_id_fkey FOREIGN KEY (last_model_config_id) REFERENCES chat_model_configs(id); diff --git a/coderd/database/migrations/000546_chat_synthetic_api_keys.up.sql b/coderd/database/migrations/000546_chat_synthetic_api_keys.up.sql deleted file mode 100644 index bde888b0dcc..00000000000 --- a/coderd/database/migrations/000546_chat_synthetic_api_keys.up.sql +++ /dev/null @@ -1,12 +0,0 @@ -CREATE TABLE chat_synthetic_api_keys ( - user_id uuid NOT NULL PRIMARY KEY REFERENCES users(id) ON DELETE CASCADE, - api_key_id text NOT NULL UNIQUE REFERENCES api_keys(id) ON DELETE CASCADE, - created_at timestamptz NOT NULL DEFAULT now(), - updated_at timestamptz NOT NULL DEFAULT now() -); - -ALTER TABLE chat_messages -DROP CONSTRAINT chat_messages_api_key_id_fkey; - -ALTER TABLE chat_queued_messages -DROP CONSTRAINT chat_queued_messages_api_key_id_fkey; diff --git a/coderd/database/migrations/000546_chat_synthetic_api_keys.down.sql b/coderd/database/migrations/000546_drop_chat_history_api_key_fks.down.sql similarity index 94% rename from coderd/database/migrations/000546_chat_synthetic_api_keys.down.sql rename to coderd/database/migrations/000546_drop_chat_history_api_key_fks.down.sql index 88c6ffc34da..a93f8bf6dde 100644 --- a/coderd/database/migrations/000546_chat_synthetic_api_keys.down.sql +++ b/coderd/database/migrations/000546_drop_chat_history_api_key_fks.down.sql @@ -1,5 +1,3 @@ -DROP TABLE chat_synthetic_api_keys; - UPDATE chat_messages SET api_key_id = NULL WHERE api_key_id IS NOT NULL diff --git a/coderd/database/migrations/000546_drop_chat_history_api_key_fks.up.sql b/coderd/database/migrations/000546_drop_chat_history_api_key_fks.up.sql new file mode 100644 index 00000000000..1a0831c9aca --- /dev/null +++ b/coderd/database/migrations/000546_drop_chat_history_api_key_fks.up.sql @@ -0,0 +1,5 @@ +ALTER TABLE chat_messages +DROP CONSTRAINT chat_messages_api_key_id_fkey; + +ALTER TABLE chat_queued_messages +DROP CONSTRAINT chat_queued_messages_api_key_id_fkey; diff --git a/coderd/database/migrations/migrate_test.go b/coderd/database/migrations/migrate_test.go index 2e5be77104e..11897bec83c 100644 --- a/coderd/database/migrations/migrate_test.go +++ b/coderd/database/migrations/migrate_test.go @@ -1693,13 +1693,13 @@ func TestMigration000546ChatHistoryAPIKeyConstraints(t *testing.T) { } } - upSQL, err := os.ReadFile("000546_chat_synthetic_api_keys.up.sql") + upSQL, err := os.ReadFile("000546_drop_chat_history_api_key_fks.up.sql") require.NoError(t, err) _, err = sqlDB.ExecContext(ctx, string(upSQL)) require.NoError(t, err) assertConstraintCount(t, 0) - downSQL, err := os.ReadFile("000546_chat_synthetic_api_keys.down.sql") + downSQL, err := os.ReadFile("000546_drop_chat_history_api_key_fks.down.sql") require.NoError(t, err) _, err = sqlDB.ExecContext(ctx, string(downSQL)) require.NoError(t, err) diff --git a/coderd/database/migrations/testdata/fixtures/000546_chat_synthetic_api_keys.up.sql b/coderd/database/migrations/testdata/fixtures/000546_chat_synthetic_api_keys.up.sql deleted file mode 100644 index e6d0257e8ee..00000000000 --- a/coderd/database/migrations/testdata/fixtures/000546_chat_synthetic_api_keys.up.sql +++ /dev/null @@ -1,5 +0,0 @@ -INSERT INTO chat_synthetic_api_keys (user_id, api_key_id) -VALUES ( - '30095c71-380b-457a-8995-97b8ee6e5307', - 'fixture-api-key' -); diff --git a/coderd/database/models.go b/coderd/database/models.go index a0e88e5bb8b..73c11d14680 100644 --- a/coderd/database/models.go +++ b/coderd/database/models.go @@ -5158,13 +5158,6 @@ type ChatQueuedMessage struct { ReasoningEffort NullChatReasoningEffort `db:"reasoning_effort" json:"reasoning_effort"` } -type ChatSyntheticApiKey struct { - UserID uuid.UUID `db:"user_id" json:"user_id"` - APIKeyID string `db:"api_key_id" json:"api_key_id"` - CreatedAt time.Time `db:"created_at" json:"created_at"` - UpdatedAt time.Time `db:"updated_at" json:"updated_at"` -} - type ChatTable struct { ID uuid.UUID `db:"id" json:"id"` OwnerID uuid.UUID `db:"owner_id" json:"owner_id"` diff --git a/coderd/database/querier.go b/coderd/database/querier.go index 6cd2625d4a1..1866216a403 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -433,6 +433,7 @@ type sqlcQuerier interface { // loading file content. GetChatFileMetadataByChatID(ctx context.Context, chatID uuid.UUID) ([]GetChatFileMetadataByChatIDRow, error) GetChatFilesByIDs(ctx context.Context, ids []uuid.UUID) ([]ChatFile, error) + GetChatGatewayAPIKey(ctx context.Context, arg GetChatGatewayAPIKeyParams) (APIKey, error) GetChatGeneralModelOverride(ctx context.Context) (string, error) GetChatHeartbeat(ctx context.Context, arg GetChatHeartbeatParams) (ChatHeartbeat, error) // GetChatIncludeDefaultSystemPrompt preserves the legacy default @@ -471,7 +472,6 @@ type sqlcQuerier interface { // A value of 0 disables chat purging entirely. GetChatRetentionDays(ctx context.Context) (int32, error) GetChatStreamSyncRows(ctx context.Context, ids []uuid.UUID) ([]GetChatStreamSyncRowsRow, error) - GetChatSyntheticAPIKeyByUserID(ctx context.Context, userID uuid.UUID) (ChatSyntheticApiKey, error) GetChatSystemPrompt(ctx context.Context) (string, error) // GetChatSystemPromptConfig returns both chat system prompt settings in a // single read to avoid torn reads between separate site-config lookups. @@ -1065,7 +1065,6 @@ type sqlcQuerier interface { // sequence) and an explicit created_by reference. Use this when the // queued-message creator differs from the chat owner. InsertChatQueuedMessageWithCreator(ctx context.Context, arg InsertChatQueuedMessageWithCreatorParams) (ChatQueuedMessage, error) - InsertChatSyntheticAPIKey(ctx context.Context, arg InsertChatSyntheticAPIKeyParams) (int64, error) InsertCryptoKey(ctx context.Context, arg InsertCryptoKeyParams) (CryptoKey, error) InsertCustomRole(ctx context.Context, arg InsertCustomRoleParams) (CustomRole, error) InsertDBCryptKey(ctx context.Context, arg InsertDBCryptKeyParams) error @@ -1402,7 +1401,6 @@ type sqlcQuerier interface { // assigned by trigger from the current snapshot_version. UpdateChatRetryState(ctx context.Context, arg UpdateChatRetryStateParams) (Chat, error) UpdateChatStatus(ctx context.Context, arg UpdateChatStatusParams) (Chat, error) - UpdateChatSyntheticAPIKey(ctx context.Context, arg UpdateChatSyntheticAPIKeyParams) (int64, error) UpdateChatTitleByID(ctx context.Context, arg UpdateChatTitleByIDParams) (Chat, error) UpdateChatWorkspaceBinding(ctx context.Context, arg UpdateChatWorkspaceBindingParams) (Chat, error) UpdateCryptoKeyDeletesAt(ctx context.Context, arg UpdateCryptoKeyDeletesAtParams) (CryptoKey, error) diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index de9ddbce507..11e94d5923c 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -3163,6 +3163,51 @@ func (q *sqlQuerier) GetAPIKeysLastUsedAfter(ctx context.Context, lastUsed time. return items, nil } +const getChatGatewayAPIKey = `-- name: GetChatGatewayAPIKey :one +SELECT + id, hashed_secret, user_id, last_used, expires_at, created_at, updated_at, login_type, lifetime_seconds, ip_address, token_name, scopes, allow_list +FROM + api_keys +WHERE + user_id = $1 AND + token_name = $2 AND + -- Token names are unvalidated user input, so a user could create a token + -- with the chat gateway name. Excluding login_type 'token' ensures chatd + -- never picks up (and extends) a real bearer token. Synthetic gateway + -- keys are minted with the owner's login type, which is never 'token'. + login_type != 'token' +ORDER BY + created_at ASC, id ASC +LIMIT + 1 +` + +type GetChatGatewayAPIKeyParams struct { + UserID uuid.UUID `db:"user_id" json:"user_id"` + TokenName string `db:"token_name" json:"token_name"` +} + +func (q *sqlQuerier) GetChatGatewayAPIKey(ctx context.Context, arg GetChatGatewayAPIKeyParams) (APIKey, error) { + row := q.db.QueryRowContext(ctx, getChatGatewayAPIKey, arg.UserID, arg.TokenName) + var i APIKey + err := row.Scan( + &i.ID, + &i.HashedSecret, + &i.UserID, + &i.LastUsed, + &i.ExpiresAt, + &i.CreatedAt, + &i.UpdatedAt, + &i.LoginType, + &i.LifetimeSeconds, + &i.IPAddress, + &i.TokenName, + &i.Scopes, + &i.AllowList, + ) + return i, err +} + const insertAPIKey = `-- name: InsertAPIKey :one INSERT INTO api_keys ( @@ -8202,24 +8247,6 @@ func (q *sqlQuerier) GetChatStreamSyncRows(ctx context.Context, ids []uuid.UUID) return items, nil } -const getChatSyntheticAPIKeyByUserID = `-- name: GetChatSyntheticAPIKeyByUserID :one -SELECT user_id, api_key_id, created_at, updated_at -FROM chat_synthetic_api_keys -WHERE user_id = $1::uuid -` - -func (q *sqlQuerier) GetChatSyntheticAPIKeyByUserID(ctx context.Context, userID uuid.UUID) (ChatSyntheticApiKey, error) { - row := q.db.QueryRowContext(ctx, getChatSyntheticAPIKeyByUserID, userID) - var i ChatSyntheticApiKey - err := row.Scan( - &i.UserID, - &i.APIKeyID, - &i.CreatedAt, - &i.UpdatedAt, - ) - return i, err -} - const getChatUsageLimitConfig = `-- name: GetChatUsageLimitConfig :one SELECT id, singleton, enabled, default_limit_micros, period, created_at, updated_at FROM chat_usage_limit_config WHERE singleton = TRUE LIMIT 1 ` @@ -10003,25 +10030,6 @@ func (q *sqlQuerier) InsertChatQueuedMessageWithCreator(ctx context.Context, arg return i, err } -const insertChatSyntheticAPIKey = `-- name: InsertChatSyntheticAPIKey :execrows -INSERT INTO chat_synthetic_api_keys (user_id, api_key_id) -VALUES ($1::uuid, $2::text) -ON CONFLICT (user_id) DO NOTHING -` - -type InsertChatSyntheticAPIKeyParams struct { - UserID uuid.UUID `db:"user_id" json:"user_id"` - APIKeyID string `db:"api_key_id" json:"api_key_id"` -} - -func (q *sqlQuerier) InsertChatSyntheticAPIKey(ctx context.Context, arg InsertChatSyntheticAPIKeyParams) (int64, error) { - result, err := q.db.ExecContext(ctx, insertChatSyntheticAPIKey, arg.UserID, arg.APIKeyID) - if err != nil { - return 0, err - } - return result.RowsAffected() -} - const isChatHeartbeatStale = `-- name: IsChatHeartbeatStale :one SELECT NOT EXISTS ( SELECT 1 FROM chat_heartbeats @@ -12207,28 +12215,6 @@ func (q *sqlQuerier) UpdateChatStatus(ctx context.Context, arg UpdateChatStatusP return i, err } -const updateChatSyntheticAPIKey = `-- name: UpdateChatSyntheticAPIKey :execrows -UPDATE chat_synthetic_api_keys -SET api_key_id = $1::text, - updated_at = NOW() -WHERE user_id = $2::uuid - AND api_key_id = $3::text -` - -type UpdateChatSyntheticAPIKeyParams struct { - NewApiKeyID string `db:"new_api_key_id" json:"new_api_key_id"` - UserID uuid.UUID `db:"user_id" json:"user_id"` - OldApiKeyID string `db:"old_api_key_id" json:"old_api_key_id"` -} - -func (q *sqlQuerier) UpdateChatSyntheticAPIKey(ctx context.Context, arg UpdateChatSyntheticAPIKeyParams) (int64, error) { - result, err := q.db.ExecContext(ctx, updateChatSyntheticAPIKey, arg.NewApiKeyID, arg.UserID, arg.OldApiKeyID) - if err != nil { - return 0, err - } - return result.RowsAffected() -} - const updateChatTitleByID = `-- name: UpdateChatTitleByID :one WITH updated_chat AS ( UPDATE diff --git a/coderd/database/queries/apikeys.sql b/coderd/database/queries/apikeys.sql index 2b197255fb3..90e7610cf06 100644 --- a/coderd/database/queries/apikeys.sql +++ b/coderd/database/queries/apikeys.sql @@ -21,6 +21,24 @@ WHERE LIMIT 1; +-- name: GetChatGatewayAPIKey :one +SELECT + * +FROM + api_keys +WHERE + user_id = @user_id AND + token_name = @token_name AND + -- Token names are unvalidated user input, so a user could create a token + -- with the chat gateway name. Excluding login_type 'token' ensures chatd + -- never picks up (and extends) a real bearer token. Synthetic gateway + -- keys are minted with the owner's login type, which is never 'token'. + login_type != 'token' +ORDER BY + created_at ASC, id ASC +LIMIT + 1; + -- name: GetAPIKeysLastUsedAfter :many SELECT * FROM api_keys WHERE last_used > $1; diff --git a/coderd/database/queries/chats.sql b/coderd/database/queries/chats.sql index 31b0050708d..51bd02e34bd 100644 --- a/coderd/database/queries/chats.sql +++ b/coderd/database/queries/chats.sql @@ -3008,20 +3008,3 @@ LEFT JOIN to_archive t ON t.id = a.id -- created_at ASC flows through to dbpurge's digest truncation; see -- buildDigestData in dbpurge.go for the tradeoff rationale. ORDER BY (a.root_chat_id IS NULL) DESC, a.owner_id ASC, a.created_at ASC, a.id ASC; - --- name: GetChatSyntheticAPIKeyByUserID :one -SELECT * -FROM chat_synthetic_api_keys -WHERE user_id = @user_id::uuid; - --- name: InsertChatSyntheticAPIKey :execrows -INSERT INTO chat_synthetic_api_keys (user_id, api_key_id) -VALUES (@user_id::uuid, @api_key_id::text) -ON CONFLICT (user_id) DO NOTHING; - --- name: UpdateChatSyntheticAPIKey :execrows -UPDATE chat_synthetic_api_keys -SET api_key_id = @new_api_key_id::text, - updated_at = NOW() -WHERE user_id = @user_id::uuid - AND api_key_id = @old_api_key_id::text; diff --git a/coderd/database/unique_constraint.go b/coderd/database/unique_constraint.go index aa954b5cf0c..4b1a4376f2d 100644 --- a/coderd/database/unique_constraint.go +++ b/coderd/database/unique_constraint.go @@ -32,8 +32,6 @@ const ( UniqueChatMessagesPkey UniqueConstraint = "chat_messages_pkey" // ALTER TABLE ONLY chat_messages ADD CONSTRAINT chat_messages_pkey PRIMARY KEY (id); UniqueChatModelConfigsPkey UniqueConstraint = "chat_model_configs_pkey" // ALTER TABLE ONLY chat_model_configs ADD CONSTRAINT chat_model_configs_pkey PRIMARY KEY (id); UniqueChatQueuedMessagesPkey UniqueConstraint = "chat_queued_messages_pkey" // ALTER TABLE ONLY chat_queued_messages ADD CONSTRAINT chat_queued_messages_pkey PRIMARY KEY (id); - UniqueChatSyntheticAPIKeysAPIKeyIDKey UniqueConstraint = "chat_synthetic_api_keys_api_key_id_key" // ALTER TABLE ONLY chat_synthetic_api_keys ADD CONSTRAINT chat_synthetic_api_keys_api_key_id_key UNIQUE (api_key_id); - UniqueChatSyntheticAPIKeysPkey UniqueConstraint = "chat_synthetic_api_keys_pkey" // ALTER TABLE ONLY chat_synthetic_api_keys ADD CONSTRAINT chat_synthetic_api_keys_pkey PRIMARY KEY (user_id); UniqueChatUsageLimitConfigPkey UniqueConstraint = "chat_usage_limit_config_pkey" // ALTER TABLE ONLY chat_usage_limit_config ADD CONSTRAINT chat_usage_limit_config_pkey PRIMARY KEY (id); UniqueChatUsageLimitConfigSingletonKey UniqueConstraint = "chat_usage_limit_config_singleton_key" // ALTER TABLE ONLY chat_usage_limit_config ADD CONSTRAINT chat_usage_limit_config_singleton_key UNIQUE (singleton); UniqueChatsPkey UniqueConstraint = "chats_pkey" // ALTER TABLE ONLY chats ADD CONSTRAINT chats_pkey PRIMARY KEY (id); diff --git a/coderd/exp_chats_test.go b/coderd/exp_chats_test.go index d48ab69db2c..93cab6479d0 100644 --- a/coderd/exp_chats_test.go +++ b/coderd/exp_chats_test.go @@ -9613,9 +9613,12 @@ func TestManualTitleEndpointsPassOwnerSyntheticAPIKeyToAIGateway(t *testing.T) { }) require.NoError(t, tt.call(ctx, client, chat.ID)) - mapping, err := db.GetChatSyntheticAPIKeyByUserID(dbauthz.AsSystemRestricted(ctx), firstUser.UserID) + gatewayKey, err := db.GetChatGatewayAPIKey(dbauthz.AsSystemRestricted(ctx), database.GetChatGatewayAPIKeyParams{ + UserID: firstUser.UserID, + TokenName: chatd.GatewayTokenName(firstUser.UserID), + }) require.NoError(t, err) - require.Equal(t, mapping.APIKeyID, testutil.RequireReceive(ctx, t, seenAPIKeyID)) + require.Equal(t, gatewayKey.ID, testutil.RequireReceive(ctx, t, seenAPIKeyID)) }) } } diff --git a/coderd/x/chatd/ARCHITECTURE.md b/coderd/x/chatd/ARCHITECTURE.md index 119d154412c..dee9660132c 100644 --- a/coderd/x/chatd/ARCHITECTURE.md +++ b/coderd/x/chatd/ARCHITECTURE.md @@ -9,13 +9,13 @@ Chatd has 4 main pieces: # Gateway attribution keys -Chatd attributes AI Gateway requests with a synthetic API key owned by the chat owner. The mapping is stored in `chat_synthetic_api_keys`, with one active key per user. All chatd AI Gateway attribution resolves the key from `chats.owner_id`. Callers do not provide the key ID. +Chatd attributes AI Gateway requests with a synthetic API key owned by the chat owner, one key per user. There is no mapping table: the key is found in `api_keys` by its deterministic token name, `chatd__session_token`, excluding `login_type = 'token'` rows. Token names are unvalidated user input, so the login type filter ensures chatd never picks up (or extends) a real bearer token a user created with the colliding name. Synthetic keys are minted with the owner's login type, which is never `'token'`. All chatd AI Gateway attribution resolves the key from `chats.owner_id`; callers do not provide the key ID. -Synthetic keys expire after 30 days and renew when less than 24 hours remain. Renewal replaces the mapping but does not delete the previous key, because an in-flight request may still use it. Concurrent mints use conditional insert/update operations, and losing keys are deleted. The generated token is discarded, so the stored key cannot be used as a bearer credential. +Synthetic keys expire after 30 days. When less than 24 hours remain, chatd extends the expiry of the existing row in place instead of replacing it, because an in-flight generation may have already delegated the current key ID to the gateway. The key ID is therefore stable for the lifetime of the user. Mints and extensions are serialized with a per-user advisory lock, since the partial unique index on token names only covers `login_type = 'token'` rows. The generated token is discarded, so the stored key cannot be used as a bearer credential, and it carries a minimal scope as defense in depth. The legacy `api_key_id` columns on messages and queued messages are still stamped with the synthetic key for rolling deployment compatibility, but they no longer have foreign keys to `api_keys`. They are not the source of gateway routing. Stale IDs are harmless because chatd resolves attribution from `chats.owner_id`. -Deleting a synthetic key removes its mapping without updating chat messages, queued messages, or their version fields. Deleting all user API keys or resetting a password therefore causes chatd to mint a replacement on the next request without mutating history. User suspension and deletion still block delegated gateway authorization. +Deleting a synthetic key (password reset, explicit key deletion, dbpurge of long-expired keys) does not touch chat messages, queued messages, or their version fields. Chatd mints a replacement on the next request without mutating history. User suspension and deletion still block delegated gateway authorization. # Core state machine diff --git a/coderd/x/chatd/chatd_internal_test.go b/coderd/x/chatd/chatd_internal_test.go index dacc9116b9b..9f97a56b138 100644 --- a/coderd/x/chatd/chatd_internal_test.go +++ b/coderd/x/chatd/chatd_internal_test.go @@ -858,8 +858,7 @@ func TestRegenerateChatTitle_PersistsAndBroadcasts(t *testing.T) { db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes() db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), gomock.Any()).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes() db.EXPECT().GetChatUsageLimitConfig(gomock.Any()).Return(database.ChatUsageLimitConfig{}, sql.ErrNoRows) - db.EXPECT().GetChatSyntheticAPIKeyByUserID(gomock.Any(), ownerID).Return(database.ChatSyntheticApiKey{UserID: ownerID, APIKeyID: activeAPIKeyID}, nil) - db.EXPECT().GetAPIKeyByID(gomock.Any(), activeAPIKeyID).Return(database.APIKey{ID: activeAPIKeyID, UserID: ownerID, ExpiresAt: time.Now().Add(48 * time.Hour)}, nil) + db.EXPECT().GetChatGatewayAPIKey(gomock.Any(), database.GetChatGatewayAPIKeyParams{UserID: ownerID, TokenName: GatewayTokenName(ownerID)}).Return(database.APIKey{ID: activeAPIKeyID, UserID: ownerID, ExpiresAt: time.Now().Add(48 * time.Hour)}, nil) db.EXPECT().GetChatMessagesByChatIDAscPaginated( gomock.Any(), database.GetChatMessagesByChatIDAscPaginatedParams{ @@ -1011,8 +1010,7 @@ func TestRegenerateChatTitle_SkipsPersistWhenTitleChangedConcurrently(t *testing db.EXPECT().GetAIProviderKeysByProviderID(gomock.Any(), providerID).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes() db.EXPECT().GetAIProviderKeysByProviderIDs(gomock.Any(), gomock.Any()).Return([]database.AIProviderKey{{ProviderID: providerID, APIKey: "test-key"}}, nil).AnyTimes() db.EXPECT().GetChatUsageLimitConfig(gomock.Any()).Return(database.ChatUsageLimitConfig{}, sql.ErrNoRows) - db.EXPECT().GetChatSyntheticAPIKeyByUserID(gomock.Any(), ownerID).Return(database.ChatSyntheticApiKey{UserID: ownerID, APIKeyID: activeAPIKeyID}, nil) - db.EXPECT().GetAPIKeyByID(gomock.Any(), activeAPIKeyID).Return(database.APIKey{ID: activeAPIKeyID, UserID: ownerID, ExpiresAt: time.Now().Add(48 * time.Hour)}, nil) + db.EXPECT().GetChatGatewayAPIKey(gomock.Any(), database.GetChatGatewayAPIKeyParams{UserID: ownerID, TokenName: GatewayTokenName(ownerID)}).Return(database.APIKey{ID: activeAPIKeyID, UserID: ownerID, ExpiresAt: time.Now().Add(48 * time.Hour)}, nil) db.EXPECT().GetChatMessagesByChatIDAscPaginated( gomock.Any(), database.GetChatMessagesByChatIDAscPaginatedParams{ diff --git a/coderd/x/chatd/chatd_test.go b/coderd/x/chatd/chatd_test.go index d843105cbe1..e0b2c668e5e 100644 --- a/coderd/x/chatd/chatd_test.go +++ b/coderd/x/chatd/chatd_test.go @@ -1628,9 +1628,12 @@ func TestCreateChatPersistsSyntheticAPIKeyIDOnInitialUserMessage(t *testing.T) { require.Len(t, messages, 1) require.Equal(t, database.ChatMessageRoleUser, messages[0].Role) require.True(t, messages[0].APIKeyID.Valid) - mapping, err := db.GetChatSyntheticAPIKeyByUserID(ctx, user.ID) + gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{ + UserID: user.ID, + TokenName: chatd.GatewayTokenName(user.ID), + }) require.NoError(t, err) - require.Equal(t, mapping.APIKeyID, messages[0].APIKeyID.String) + require.Equal(t, gatewayKey.ID, messages[0].APIKeyID.String) } func TestSendMessagePersistsSyntheticAPIKeyIDOnUserMessage(t *testing.T) { @@ -1659,14 +1662,17 @@ func TestSendMessagePersistsSyntheticAPIKeyIDOnUserMessage(t *testing.T) { require.NoError(t, err) require.False(t, result.Queued) require.True(t, result.Message.APIKeyID.Valid) - mapping, err := db.GetChatSyntheticAPIKeyByUserID(ctx, user.ID) + gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{ + UserID: user.ID, + TokenName: chatd.GatewayTokenName(user.ID), + }) require.NoError(t, err) - require.Equal(t, mapping.APIKeyID, result.Message.APIKeyID.String) + require.Equal(t, gatewayKey.ID, result.Message.APIKeyID.String) stored, err := db.GetChatMessageByID(ctx, result.Message.ID) require.NoError(t, err) require.True(t, stored.APIKeyID.Valid) - require.Equal(t, mapping.APIKeyID, stored.APIKeyID.String) + require.Equal(t, gatewayKey.ID, stored.APIKeyID.String) } func TestSendMessagePersistsSyntheticAPIKeyIDOnQueuedUserMessage(t *testing.T) { @@ -1705,16 +1711,19 @@ func TestSendMessagePersistsSyntheticAPIKeyIDOnQueuedUserMessage(t *testing.T) { require.True(t, result.Queued) require.NotNil(t, result.QueuedMessage) - mapping, err := db.GetChatSyntheticAPIKeyByUserID(ctx, user.ID) + gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{ + UserID: user.ID, + TokenName: chatd.GatewayTokenName(user.ID), + }) require.NoError(t, err) require.True(t, result.QueuedMessage.APIKeyID.Valid) - require.Equal(t, mapping.APIKeyID, result.QueuedMessage.APIKeyID.String) + require.Equal(t, gatewayKey.ID, result.QueuedMessage.APIKeyID.String) queued, err := db.GetChatQueuedMessages(ctx, chat.ID) require.NoError(t, err) require.Len(t, queued, 1) require.True(t, queued[0].APIKeyID.Valid) - require.Equal(t, mapping.APIKeyID, queued[0].APIKeyID.String) + require.Equal(t, gatewayKey.ID, queued[0].APIKeyID.String) } func TestEditMessagePersistsSyntheticAPIKeyIDOnReplacement(t *testing.T) { @@ -1749,15 +1758,18 @@ func TestEditMessagePersistsSyntheticAPIKeyIDOnReplacement(t *testing.T) { }) require.NoError(t, err) - mapping, err := db.GetChatSyntheticAPIKeyByUserID(ctx, user.ID) + gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{ + UserID: user.ID, + TokenName: chatd.GatewayTokenName(user.ID), + }) require.NoError(t, err) require.True(t, result.Message.APIKeyID.Valid) - require.Equal(t, mapping.APIKeyID, result.Message.APIKeyID.String) + require.Equal(t, gatewayKey.ID, result.Message.APIKeyID.String) stored, err := db.GetChatMessageByID(ctx, result.Message.ID) require.NoError(t, err) require.True(t, stored.APIKeyID.Valid) - require.Equal(t, mapping.APIKeyID, stored.APIKeyID.String) + require.Equal(t, gatewayKey.ID, stored.APIKeyID.String) } func TestSendMessageQueueBehaviorQueuesWhenBusy(t *testing.T) { @@ -5447,7 +5459,10 @@ func TestActiveServer_RoutingPreservesAPIKeyAfterCompaction(t *testing.T) { }, }) require.NoError(t, err) - mapping, err := db.GetChatSyntheticAPIKeyByUserID(ctx, user.ID) + gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{ + UserID: user.ID, + TokenName: chatd.GatewayTokenName(user.ID), + }) require.NoError(t, err) contextContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{{ Type: codersdk.ChatMessagePartTypeContextFile, @@ -5460,7 +5475,7 @@ func TestActiveServer_RoutingPreservesAPIKeyAfterCompaction(t *testing.T) { require.NoError(t, err) _, err = db.InsertChatMessages(ctx, chatd.BuildSingleUserChatMessageInsertParams( chat.ID, - mapping.APIKeyID, + gatewayKey.ID, contextContent, database.ChatMessageVisibilityBoth, model.ID, @@ -5489,14 +5504,14 @@ func TestActiveServer_RoutingPreservesAPIKeyAfterCompaction(t *testing.T) { compressed := compressedChatSummarizedMessages(t, append(promptMessages, messages...)) require.Len(t, compressed.summaries, 1) require.True(t, compressed.summaries[0].APIKeyID.Valid) - require.Equal(t, mapping.APIKeyID, compressed.summaries[0].APIKeyID.String) + require.Equal(t, gatewayKey.ID, compressed.summaries[0].APIKeyID.String) requests := factory.RequestsSnapshot() require.NotEmpty(t, requests) for _, req := range requests { require.Equal(t, provider.Name, req.ProviderName) require.Equal(t, aibridge.SourceAgents, req.Source) - require.Equal(t, mapping.APIKeyID, req.APIKeyID) + require.Equal(t, gatewayKey.ID, req.APIKeyID) require.Equal(t, "sk-user-aibridge", req.Request.Header.Get("X-Api-Key")) require.Equal(t, "delegated", req.Request.Header.Get(aibridge.HeaderCoderToken)) } @@ -10140,7 +10155,10 @@ func TestProcessChat_RoutingUsesDelegatedAPIKey(t *testing.T) { }, }) require.NoError(t, err) - mapping, err := db.GetChatSyntheticAPIKeyByUserID(ctx, user.ID) + gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{ + UserID: user.ID, + TokenName: chatd.GatewayTokenName(user.ID), + }) require.NoError(t, err) _, events, cancel, ok := creator.Subscribe(ctx, chat.ID, nil, 0) @@ -10167,7 +10185,7 @@ func TestProcessChat_RoutingUsesDelegatedAPIKey(t *testing.T) { for _, req := range requests { require.Equal(t, provider.Name, req.ProviderName) require.Equal(t, aibridge.SourceAgents, req.Source) - require.Equal(t, mapping.APIKeyID, req.APIKeyID) + require.Equal(t, gatewayKey.ID, req.APIKeyID) require.Equal(t, "Bearer sk-user-aibridge", req.Request.Header.Get("Authorization")) require.Empty(t, req.Request.Header.Get("X-Api-Key")) require.Equal(t, "delegated", req.Request.Header.Get(aibridge.HeaderCoderToken)) @@ -10205,7 +10223,10 @@ func TestProcessChat_RoutingPreservesAPIKeyAfterWorkspaceContext(t *testing.T) { }, }) require.NoError(t, err) - mapping, err := db.GetChatSyntheticAPIKeyByUserID(ctx, user.ID) + gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{ + UserID: user.ID, + TokenName: chatd.GatewayTokenName(user.ID), + }) require.NoError(t, err) const contextText = "# Project instructions\nAlways keep routing metadata." @@ -10243,7 +10264,7 @@ func TestProcessChat_RoutingPreservesAPIKeyAfterWorkspaceContext(t *testing.T) { for _, req := range requests { require.Equal(t, provider.Name, req.ProviderName) require.Equal(t, aibridge.SourceAgents, req.Source) - require.Equal(t, mapping.APIKeyID, req.APIKeyID) + require.Equal(t, gatewayKey.ID, req.APIKeyID) require.Equal(t, "Bearer sk-user-aibridge", req.Request.Header.Get("Authorization")) require.Equal(t, "delegated", req.Request.Header.Get(aibridge.HeaderCoderToken)) } diff --git a/coderd/x/chatd/subagent_internal_test.go b/coderd/x/chatd/subagent_internal_test.go index 24112baf65a..896d658106d 100644 --- a/coderd/x/chatd/subagent_internal_test.go +++ b/coderd/x/chatd/subagent_internal_test.go @@ -672,7 +672,10 @@ func TestCreateChildSubagentChatPersistsOwnerSyntheticAPIKeyID(t *testing.T) { ) require.NoError(t, err) - mapping, err := db.GetChatSyntheticAPIKeyByUserID(ctx, user.ID) + gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{ + UserID: user.ID, + TokenName: GatewayTokenName(user.ID), + }) require.NoError(t, err) messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{ ChatID: child.ID, @@ -684,7 +687,7 @@ func TestCreateChildSubagentChatPersistsOwnerSyntheticAPIKeyID(t *testing.T) { continue } require.True(t, message.APIKeyID.Valid) - require.Equal(t, mapping.APIKeyID, message.APIKeyID.String) + require.Equal(t, gatewayKey.ID, message.APIKeyID.String) return } require.Fail(t, "child user message not found") @@ -710,7 +713,10 @@ func TestSendSubagentMessagePersistsOwnerSyntheticAPIKeyID(t *testing.T) { ) require.NoError(t, err) - mapping, err := db.GetChatSyntheticAPIKeyByUserID(ctx, user.ID) + gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{ + UserID: user.ID, + TokenName: GatewayTokenName(user.ID), + }) require.NoError(t, err) messages, err := db.GetChatMessagesByChatID(ctx, database.GetChatMessagesByChatIDParams{ ChatID: child.ID, @@ -725,7 +731,7 @@ func TestSendSubagentMessagePersistsOwnerSyntheticAPIKeyID(t *testing.T) { } require.NotZero(t, latestUserMessage.ID) require.True(t, latestUserMessage.APIKeyID.Valid) - require.Equal(t, mapping.APIKeyID, latestUserMessage.APIKeyID.String) + require.Equal(t, gatewayKey.ID, latestUserMessage.APIKeyID.String) } func TestCreateChildSubagentChatInheritsWorkspaceBinding(t *testing.T) { diff --git a/coderd/x/chatd/synthetickey.go b/coderd/x/chatd/synthetickey.go index eb0248cbc63..550419e2b9a 100644 --- a/coderd/x/chatd/synthetickey.go +++ b/coderd/x/chatd/synthetickey.go @@ -17,56 +17,70 @@ import ( const ( syntheticAPIKeyLifetime = 30 * 24 * time.Hour syntheticAPIKeyRenewMargin = 24 * time.Hour - syntheticAPIKeyMaxAttempts = 3 ) +// GatewayTokenName returns the deterministic token name of the synthetic +// gateway key for a user. The name is the lookup key: no mapping table exists, +// so attribution resolves the key by (user_id, token_name, login_type != +// 'token'). +func GatewayTokenName(ownerID uuid.UUID) string { + return fmt.Sprintf("chatd_%s_session_token", ownerID) +} + +// ensureSyntheticAPIKeyID returns the ID of the synthetic gateway key for the +// given user, minting or extending it as needed. The key ID is stable for the +// lifetime of the user: near-expiry keys are extended in place rather than +// replaced, because an in-flight generation may have already delegated the +// current key ID to the gateway. func (p *Server) ensureSyntheticAPIKeyID(ctx context.Context, ownerID uuid.UUID) (string, error) { ctx = dbauthz.AsChatdKeyMinter(ctx, ownerID) - for range syntheticAPIKeyMaxAttempts { - mapping, err := p.db.GetChatSyntheticAPIKeyByUserID(ctx, ownerID) - if err == nil { - key, keyErr := p.db.GetAPIKeyByID(ctx, mapping.APIKeyID) - switch { - case keyErr == nil && key.ExpiresAt.After(p.clock.Now().Add(syntheticAPIKeyRenewMargin)): - return key.ID, nil - case keyErr != nil && !xerrors.Is(keyErr, sql.ErrNoRows): - return "", xerrors.Errorf("get synthetic API key: %w", keyErr) - } - } else if !xerrors.Is(err, sql.ErrNoRows) { - return "", xerrors.Errorf("get synthetic API key mapping: %w", err) - } - - keyID, retry, err := p.mintSyntheticAPIKey(ctx, ownerID) - if err != nil { - return "", err - } - if !retry { - return keyID, nil - } + key, err := p.db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{ + UserID: ownerID, + TokenName: GatewayTokenName(ownerID), + }) + switch { + case err == nil && key.ExpiresAt.After(p.clock.Now().Add(syntheticAPIKeyRenewMargin)): + return key.ID, nil + case err != nil && !xerrors.Is(err, sql.ErrNoRows): + return "", xerrors.Errorf("get synthetic API key: %w", err) } - return "", xerrors.New("ensure synthetic API key: concurrent update retry limit reached") + return p.mintSyntheticAPIKey(ctx, ownerID) } -func (p *Server) mintSyntheticAPIKey(ctx context.Context, ownerID uuid.UUID) (keyID string, retry bool, err error) { - err = p.db.InTx(func(tx database.Store) error { - mapping, mappingErr := tx.GetChatSyntheticAPIKeyByUserID(ctx, ownerID) - hasMapping := mappingErr == nil - if mappingErr != nil && !xerrors.Is(mappingErr, sql.ErrNoRows) { - return xerrors.Errorf("get synthetic API key mapping: %w", mappingErr) +// mintSyntheticAPIKey extends or mints the synthetic gateway key under a +// per-user advisory lock. The lock serializes concurrent mints because the +// partial unique index on token names only covers login_type 'token' rows, so +// nothing else prevents duplicate synthetic keys. +func (p *Server) mintSyntheticAPIKey(ctx context.Context, ownerID uuid.UUID) (string, error) { + tokenName := GatewayTokenName(ownerID) + var keyID string + err := p.db.InTx(func(tx database.Store) error { + err := tx.AcquireLock(ctx, database.GenLockID("chatd_gateway_key:"+ownerID.String())) + if err != nil { + return xerrors.Errorf("acquire chat gateway key lock: %w", err) } - if hasMapping { - key, keyErr := tx.GetAPIKeyByID(ctx, mapping.APIKeyID) - if keyErr == nil && key.ExpiresAt.After(p.clock.Now().Add(syntheticAPIKeyRenewMargin)) { - keyID = key.ID + key, err := tx.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{ + UserID: ownerID, + TokenName: tokenName, + }) + if err == nil { + keyID = key.ID + if key.ExpiresAt.After(p.clock.Now().Add(syntheticAPIKeyRenewMargin)) { return nil } - if keyErr != nil { - if xerrors.Is(keyErr, sql.ErrNoRows) { - retry = true - return nil - } - return xerrors.Errorf("get synthetic API key: %w", keyErr) + err = tx.UpdateAPIKeyByID(ctx, database.UpdateAPIKeyByIDParams{ + ID: key.ID, + LastUsed: key.LastUsed, + ExpiresAt: p.clock.Now().Add(syntheticAPIKeyLifetime), + IPAddress: key.IPAddress, + }) + if err != nil { + return xerrors.Errorf("extend synthetic API key: %w", err) } + return nil + } + if !xerrors.Is(err, sql.ErrNoRows) { + return xerrors.Errorf("get synthetic API key: %w", err) } owner, err := tx.GetUserForChatSyntheticAPIKeyByID(ctx, ownerID) @@ -78,44 +92,24 @@ func (p *Server) mintSyntheticAPIKey(ctx context.Context, ownerID uuid.UUID) (ke LoginType: owner.LoginType, ExpiresAt: p.clock.Now().Add(syntheticAPIKeyLifetime), LifetimeSeconds: int64(syntheticAPIKeyLifetime.Seconds()), - TokenName: fmt.Sprintf("%s_chat_gateway_key", ownerID), + TokenName: tokenName, + // The key only attributes gateway requests; the secret is + // discarded, so it is never usable as a bearer credential. The + // minimal scope is defense in depth on top of that. + Scopes: database.APIKeyScopes{database.ApiKeyScopeApiKeyRead}, }) if err != nil { return xerrors.Errorf("generate synthetic API key: %w", err) } - key, err := tx.InsertAPIKey(ctx, params) + inserted, err := tx.InsertAPIKey(ctx, params) if err != nil { return xerrors.Errorf("insert synthetic API key: %w", err) } - - var rows int64 - if hasMapping { - rows, err = tx.UpdateChatSyntheticAPIKey(ctx, database.UpdateChatSyntheticAPIKeyParams{ - UserID: ownerID, - OldApiKeyID: mapping.APIKeyID, - NewApiKeyID: key.ID, - }) - } else { - rows, err = tx.InsertChatSyntheticAPIKey(ctx, database.InsertChatSyntheticAPIKeyParams{ - UserID: ownerID, - APIKeyID: key.ID, - }) - } - if err != nil { - return xerrors.Errorf("publish synthetic API key mapping: %w", err) - } - if rows == 1 { - keyID = key.ID - return nil - } - if err := tx.DeleteAPIKeyByID(ctx, key.ID); err != nil { - return xerrors.Errorf("delete losing synthetic API key: %w", err) - } - retry = true + keyID = inserted.ID return nil }, nil) if err != nil { - return "", false, err + return "", err } - return keyID, retry, nil + return keyID, nil } diff --git a/coderd/x/chatd/synthetickey_internal_test.go b/coderd/x/chatd/synthetickey_internal_test.go index 986479952cf..fa63cf24e46 100644 --- a/coderd/x/chatd/synthetickey_internal_test.go +++ b/coderd/x/chatd/synthetickey_internal_test.go @@ -14,9 +14,17 @@ import ( "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/database/dbgen" "github.com/coder/coder/v2/coderd/database/dbtestutil" + "github.com/coder/coder/v2/coderd/database/dbtime" "github.com/coder/quartz" ) +func getGatewayKey(ctx context.Context, db database.Store, userID uuid.UUID) (database.APIKey, error) { + return db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{ + UserID: userID, + TokenName: GatewayTokenName(userID), + }) +} + func TestSyntheticAPIKeyLifecycle(t *testing.T) { t.Parallel() @@ -33,29 +41,110 @@ func TestSyntheticAPIKeyLifecycle(t *testing.T) { first, err := db.GetAPIKeyByID(t.Context(), firstID) require.NoError(t, err) require.Equal(t, user.LoginType, first.LoginType) + require.Equal(t, GatewayTokenName(user.ID), first.TokenName) + require.Equal(t, database.APIKeyScopes{database.ApiKeyScopeApiKeyRead}, first.Scopes) require.WithinDuration(t, server.clock.Now().Add(syntheticAPIKeyLifetime), first.ExpiresAt, time.Second) + // Within the renew margin the key is extended in place: the ID stays + // stable because in-flight generations may have delegated it already. err = db.UpdateAPIKeyByID(t.Context(), database.UpdateAPIKeyByIDParams{ ID: first.ID, LastUsed: first.LastUsed, - ExpiresAt: server.clock.Now().Add(syntheticAPIKeyRenewMargin), + ExpiresAt: server.clock.Now().Add(time.Hour), IPAddress: first.IPAddress, }) require.NoError(t, err) renewedID, err := server.ensureSyntheticAPIKeyID(t.Context(), user.ID) require.NoError(t, err) - require.NotEqual(t, firstID, renewedID) - _, err = db.GetAPIKeyByID(t.Context(), firstID) - require.NoError(t, err, "renewal must preserve the previous key") + require.Equal(t, firstID, renewedID) + renewed, err := db.GetAPIKeyByID(t.Context(), renewedID) + require.NoError(t, err) + require.WithinDuration(t, server.clock.Now().Add(syntheticAPIKeyLifetime), renewed.ExpiresAt, time.Second) + + // A fully expired key is extended the same way. + err = db.UpdateAPIKeyByID(t.Context(), database.UpdateAPIKeyByIDParams{ + ID: first.ID, + LastUsed: first.LastUsed, + ExpiresAt: server.clock.Now().Add(-time.Hour), + IPAddress: first.IPAddress, + }) + require.NoError(t, err) - require.NoError(t, db.DeleteAPIKeyByID(t.Context(), renewedID)) - _, err = db.GetChatSyntheticAPIKeyByUserID(t.Context(), user.ID) + revivedID, err := server.ensureSyntheticAPIKeyID(t.Context(), user.ID) + require.NoError(t, err) + require.Equal(t, firstID, revivedID) + revived, err := db.GetAPIKeyByID(t.Context(), revivedID) + require.NoError(t, err) + require.WithinDuration(t, server.clock.Now().Add(syntheticAPIKeyLifetime), revived.ExpiresAt, time.Second) + + // External deletion (password reset, dbpurge) causes a remint. + require.NoError(t, db.DeleteAPIKeyByID(t.Context(), firstID)) + _, err = getGatewayKey(t.Context(), db, user.ID) require.ErrorIs(t, err, sql.ErrNoRows) recreatedID, err := server.ensureSyntheticAPIKeyID(t.Context(), user.ID) require.NoError(t, err) - require.NotEqual(t, renewedID, recreatedID) + require.NotEqual(t, firstID, recreatedID) +} + +func TestSyntheticAPIKeyIgnoresUserTokenCollision(t *testing.T) { + t.Parallel() + + db, _ := dbtestutil.NewDB(t) + user := dbgen.User(t, db, database.User{}) + server := &Server{db: db, clock: quartz.NewReal()} + + // Token names are unvalidated user input, so a user can create a token + // named exactly like the synthetic gateway key. It must never be picked + // up or extended; its near-margin expiry would otherwise trigger the + // extension path. + collisionExpiry := server.clock.Now().Add(time.Hour).UTC() + collision, _ := dbgen.APIKey(t, db, database.APIKey{ + UserID: user.ID, + LoginType: database.LoginTypeToken, + TokenName: GatewayTokenName(user.ID), + ExpiresAt: collisionExpiry, + }) + + syntheticID, err := server.ensureSyntheticAPIKeyID(t.Context(), user.ID) + require.NoError(t, err) + require.NotEqual(t, collision.ID, syntheticID) + + synthetic, err := db.GetAPIKeyByID(t.Context(), syntheticID) + require.NoError(t, err) + require.Equal(t, user.LoginType, synthetic.LoginType) + + unchanged, err := db.GetAPIKeyByID(t.Context(), collision.ID) + require.NoError(t, err) + require.WithinDuration(t, collisionExpiry, unchanged.ExpiresAt, time.Millisecond) +} + +func TestSyntheticAPIKeySurvivesSuspension(t *testing.T) { + t.Parallel() + + db, _ := dbtestutil.NewDB(t) + user := dbgen.User(t, db, database.User{}) + server := &Server{db: db, clock: quartz.NewReal()} + + keyID, err := server.ensureSyntheticAPIKeyID(t.Context(), user.ID) + require.NoError(t, err) + + _, err = db.UpdateUserStatus(t.Context(), database.UpdateUserStatusParams{ + ID: user.ID, + Status: database.UserStatusSuspended, + UpdatedAt: dbtime.Now(), + }) + require.NoError(t, err) + + // Suspension does not delete the key. Delegated gateway authorization + // rejects suspended owners at request time instead; see the + // aibridgedserver IsAuthorized tests for that rejection. + _, err = db.GetAPIKeyByID(t.Context(), keyID) + require.NoError(t, err) + sameID, err := server.ensureSyntheticAPIKeyID(t.Context(), user.ID) + require.NoError(t, err) + require.Equal(t, keyID, sameID) } func TestSyntheticAPIKeyConcurrentMint(t *testing.T) { @@ -149,7 +238,7 @@ func TestSyntheticAPIKeyDeletionDoesNotMutateChatState(t *testing.T) { _, err = db.GetAPIKeyByID(t.Context(), syntheticID) require.ErrorIs(t, err, sql.ErrNoRows) - _, err = db.GetChatSyntheticAPIKeyByUserID(t.Context(), user.ID) + _, err = getGatewayKey(t.Context(), db, user.ID) require.ErrorIs(t, err, sql.ErrNoRows) after, err := db.GetChatByID(t.Context(), chat.ID) diff --git a/coderd/x/chatd/title_override_internal_test.go b/coderd/x/chatd/title_override_internal_test.go index 0c6e385792e..9ed08a47e56 100644 --- a/coderd/x/chatd/title_override_internal_test.go +++ b/coderd/x/chatd/title_override_internal_test.go @@ -597,11 +597,10 @@ func TestGenerateManualTitleCandidate_UsesSyntheticAPIKey(t *testing.T) { BeforeID: 0, LimitVal: manualTitleMessageWindowLimit, }).Return(nil, nil) - db.EXPECT().GetChatSyntheticAPIKeyByUserID(gomock.Any(), chat.OwnerID).Return(database.ChatSyntheticApiKey{ - UserID: chat.OwnerID, - APIKeyID: apiKeyID, - }, nil) - db.EXPECT().GetAPIKeyByID(gomock.Any(), apiKeyID).Return(database.APIKey{ + db.EXPECT().GetChatGatewayAPIKey(gomock.Any(), database.GetChatGatewayAPIKeyParams{ + UserID: chat.OwnerID, + TokenName: GatewayTokenName(chat.OwnerID), + }).Return(database.APIKey{ ID: apiKeyID, UserID: chat.OwnerID, ExpiresAt: time.Now().Add(48 * time.Hour),