diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index 5a18c5fb8f30e..84b212e45d308 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 @@ -3315,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 @@ -5072,6 +5103,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 diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index d3263972e6b1e..0e1dedede6aeb 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -339,6 +339,20 @@ func defaultIPAddress() pqtype.Inet { } } +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("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) + })) +} + 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 +7495,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 16ba3084ebd86..798449a37e338 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) @@ -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) diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index 67e42764a8873..b6d7a0d1a2449 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() @@ -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() diff --git a/coderd/database/dump.sql b/coderd/database/dump.sql index a8e815ec25aaa..2313ebf1df887 100644 --- a/coderd/database/dump.sql +++ b/coderd/database/dump.sql @@ -5139,9 +5139,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,9 +5154,6 @@ 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; diff --git a/coderd/database/foreign_key_constraint.go b/coderd/database/foreign_key_constraint.go index 4b40dcf679e03..75c99671c63ea 100644 --- a/coderd/database/foreign_key_constraint.go +++ b/coderd/database/foreign_key_constraint.go @@ -24,13 +24,11 @@ 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; 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; diff --git a/coderd/database/migrations/000546_drop_chat_history_api_key_fks.down.sql b/coderd/database/migrations/000546_drop_chat_history_api_key_fks.down.sql new file mode 100644 index 0000000000000..a93f8bf6ddee4 --- /dev/null +++ b/coderd/database/migrations/000546_drop_chat_history_api_key_fks.down.sql @@ -0,0 +1,25 @@ +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_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 0000000000000..1a0831c9aca49 --- /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 ae5b2edc48926..11897bec83c5e 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_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_drop_chat_history_api_key_fks.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/querier.go b/coderd/database/querier.go index cd0a2bfa7fb70..1866216a40372 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 @@ -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 diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index a65852cf8de56..11e94d5923cfc 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 ( @@ -29857,6 +29902,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/apikeys.sql b/coderd/database/queries/apikeys.sql index 2b197255fb363..90e7610cf06db 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/users.sql b/coderd/database/queries/users.sql index 92dc26a4d7d64..3c79b4052256d 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/exp_chats.go b/coderd/exp_chats.go index 4632384683495..39f6673ab3032 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 209e51d1e7131..93cab6479d0b7 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,12 @@ func TestManualTitleEndpointsPassCallerAPIKeyToAIGateway(t *testing.T) { }) require.NoError(t, tt.call(ctx, client, chat.ID)) - require.Equal(t, wantAPIKeyID, testutil.RequireReceive(ctx, t, seenAPIKeyID)) + gatewayKey, err := db.GetChatGatewayAPIKey(dbauthz.AsSystemRestricted(ctx), database.GetChatGatewayAPIKeyParams{ + UserID: firstUser.UserID, + TokenName: chatd.GatewayTokenName(firstUser.UserID), + }) + require.NoError(t, err) + require.Equal(t, gatewayKey.ID, testutil.RequireReceive(ctx, t, seenAPIKeyID)) }) } } diff --git a/coderd/rbac/authz.go b/coderd/rbac/authz.go index 0be013849418c..58086159d9de5 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 d532d1b8448ce..dee9660132cd2 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, 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. 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 (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 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 fa2012b2adece..c110b016662e1 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 b51d7c55c6d77..a271016ab82f5 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 916df0f80fb88..9f97a56b138e5 100644 --- a/coderd/x/chatd/chatd_internal_test.go +++ b/coderd/x/chatd/chatd_internal_test.go @@ -858,6 +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().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{ @@ -1009,6 +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().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 53ca0b9078bb8..e0b2c668e5e73 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,15 @@ 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) + gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{ + UserID: user.ID, + TokenName: chatd.GatewayTokenName(user.ID), + }) + require.NoError(t, err) + require.Equal(t, gatewayKey.ID, messages[0].APIKeyID.String) } -func TestSendMessagePersistsAPIKeyIDOnUserMessage(t *testing.T) { +func TestSendMessagePersistsSyntheticAPIKeyIDOnUserMessage(t *testing.T) { t.Parallel() db, ps := dbtestutil.NewDB(t) @@ -1649,32 +1644,132 @@ 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) + gatewayKey, err := db.GetChatGatewayAPIKey(ctx, database.GetChatGatewayAPIKeyParams{ + UserID: user.ID, + TokenName: chatd.GatewayTokenName(user.ID), + }) + require.NoError(t, err) + 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, apiKey.ID, stored.APIKeyID.String) + require.Equal(t, gatewayKey.ID, 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) + + 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, 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, gatewayKey.ID, 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) + + 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, 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, gatewayKey.ID, stored.APIKeyID.String) } func TestSendMessageQueueBehaviorQueuesWhenBusy(t *testing.T) { @@ -1689,7 +1784,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 +1802,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 +1860,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 +1920,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 +1960,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 +1997,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 +2076,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 +2173,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 +2183,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 +2192,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 +2338,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 +2348,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 +2361,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 +2444,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 +2474,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 +2507,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 +2515,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 +2534,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 +2555,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 +2581,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 +2630,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 +2687,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 +2719,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 +3038,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 +3157,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 +3289,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 +3410,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 +3424,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 +3515,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 +3529,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 +3629,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 +3726,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 +3739,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 +3825,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 +3879,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 +3939,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 +3988,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 +4080,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 +4250,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 +4367,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 +4509,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 +4637,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 +5105,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 +5272,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 +5430,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 +5454,16 @@ 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) + 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, ContextFileAgentID: uuid.NullUUID{UUID: dbAgent.ID, Valid: true}, @@ -5421,7 +5475,7 @@ func TestActiveServer_RoutingPreservesAPIKeyAfterCompaction(t *testing.T) { require.NoError(t, err) _, err = db.InsertChatMessages(ctx, chatd.BuildSingleUserChatMessageInsertParams( chat.ID, - apiKey.ID, + gatewayKey.ID, contextContent, database.ChatMessageVisibilityBoth, model.ID, @@ -5450,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, apiKey.ID, 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, apiKey.ID, 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)) } @@ -5520,7 +5574,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 +5666,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 +5793,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 +5811,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 +5993,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 +6097,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 +6141,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 +6179,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 +6231,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 +6311,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 +6373,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 +6454,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 +6619,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 +7054,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 +7141,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 +7248,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 +7284,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 +7443,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 +7520,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 +7592,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 +7648,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 +8474,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 +8778,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 +8808,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 +9078,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 +9200,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 +9287,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 +9387,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 +9445,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 +9506,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 +9566,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 +9799,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 +9965,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 +10069,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 +10100,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 +10115,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 +10142,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 +10150,16 @@ 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) + 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) require.True(t, ok) @@ -10161,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, apiKey.ID, 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)) @@ -10184,7 +10208,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 +10218,16 @@ 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) + 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." // Workspace context is sourced from the agent's pinned snapshot. Seed it so @@ -10236,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, apiKey.ID, 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)) } @@ -10274,7 +10302,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 +10365,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 +10548,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 +10706,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 +10836,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 +11055,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 +11166,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 +11291,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 +11444,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 +11571,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 +11636,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 +11651,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 +11868,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 +11885,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 +11913,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 +11929,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 +11964,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 +11983,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 +12013,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 +12029,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 +12148,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 +12247,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 +12425,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 +12618,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 +12700,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 +12783,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 +12900,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 +12930,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 +12940,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 25d553ef2cd90..b9c43d8063f97 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 ffee255d12663..06b47af4e4218 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 2a5f0a42dd4ca..178c51b94612e 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 6d27463f0552e..b35c4a77c0c1f 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 8678dffea8753..f67beb17976b1 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 6b02a191150b8..e6d9103435f82 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 4657a95d006f6..36c102a1cac97 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 aae832132d61c..f2c50e3023cc0 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 9ca4d4b795520..6e20b969bc480 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 4f90342e0130c..896d658106dbe 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,89 @@ 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) + + 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, + 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, gatewayKey.ID, 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) + + 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, + 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, gatewayKey.ID, latestUserMessage.APIKeyID.String) +} + func TestCreateChildSubagentChatInheritsWorkspaceBinding(t *testing.T) { t.Parallel() @@ -881,7 +747,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 +768,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 +794,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 +806,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 +918,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 +931,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 +951,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 +981,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 +1173,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 +1223,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 +1233,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 +1292,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 +1302,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 +1424,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 +1456,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 +1486,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 +1523,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 +1727,6 @@ func TestSpawnAgent_ExploreHonorsPersonalModelOverrides(t *testing.T) { "parent-explore-personal-override", ) - ctx = withSubagentDelegatedKey(ctx, t, db, parentChat.OwnerID) resp := runSubagentTool( ctx, t, @@ -1924,7 +1764,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 +1882,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 +1920,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 +1932,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 +1999,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 +2034,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 +2081,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 +2189,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 +2228,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 +2335,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 +2518,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 +2579,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 +2640,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 +2712,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 +2724,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 +2777,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 +2787,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 +2837,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 +2851,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 +2880,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 +2890,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 +2918,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 +2927,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 +2944,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 +2962,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 +2971,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 +3061,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 +3070,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 +3311,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 +3813,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 +3825,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 f621fa579ef2a..a768f3487ee62 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 0000000000000..550419e2b9a53 --- /dev/null +++ b/coderd/x/chatd/synthetickey.go @@ -0,0 +1,115 @@ +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 +) + +// 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) + 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 p.mintSyntheticAPIKey(ctx, ownerID) +} + +// 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) + } + 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 + } + 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) + 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: 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) + } + inserted, err := tx.InsertAPIKey(ctx, params) + if err != nil { + return xerrors.Errorf("insert synthetic API key: %w", err) + } + keyID = inserted.ID + return nil + }, nil) + if err != nil { + return "", err + } + return keyID, nil +} diff --git a/coderd/x/chatd/synthetickey_internal_test.go b/coderd/x/chatd/synthetickey_internal_test.go new file mode 100644 index 0000000000000..fa63cf24e4600 --- /dev/null +++ b/coderd/x/chatd/synthetickey_internal_test.go @@ -0,0 +1,265 @@ +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/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() + + 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.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(time.Hour), + IPAddress: first.IPAddress, + }) + require.NoError(t, err) + + renewedID, err := server.ensureSyntheticAPIKeyID(t.Context(), user.ID) + require.NoError(t, err) + 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) + + 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, 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) { + 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 = getGatewayKey(t.Context(), db, 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 cdd301455f981..9ed08a47e5657 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,74 @@ 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().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), + }, 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) {