feat: limit concurrent chat agents with pooled admission (#27902)

Limits concurrent chat generation on capped deployments to 5 root chats
and 10 delegated subagent chats. The pools are deployment-wide and
independent, so delegated work can continue while root capacity is full.

The default caps live in AGPL code. Enterprise contributes only a
licensing unlock, so unlicensed deployments stay capped and cannot fail
open. Licensed deployments are uncapped while Agent Hours usage stays
below an explicit hard limit. Deployments without a hard limit remain
uncapped, and reaching the Agent Hours allocation only triggers
warnings.

Admission happens before a worker takes chat ownership. Capped
deployments serialize admission across replicas with a
transaction-scoped advisory lock and derive active and queued state from
current ownership plus fresh runner heartbeats, rather than persisted
queue markers or per-replica state. The acquisition query returns a
bounded, pool-interleaved candidate set instead of ranking the whole
backlog; a migration replaces the acquisition index with a pool-aware
one. Refused chats stay running but unowned, and interrupt requests
bypass admission so users can stop queued or over-cap chats.

The single-chat API derives `queued_for_capacity` from live pool state;
list endpoints do not report it. The UI polls that value every 5 seconds
while a chat is running and shows a callout when the chat is waiting for
capacity.

Updates the administrator documentation and deployment-wide Prometheus
gauges for active and queued agents. Replica-level values must be
aggregated with `max`, not `sum`.

> Mux updated this PR on Mike's behalf.
This commit is contained in:
Michael Suchacz
2026-08-18 16:55:43 +02:00
committed by GitHub
parent cb0a9ebbbf
commit 119f2b1dd9
51 changed files with 2187 additions and 277 deletions
+4
View File
@@ -17518,6 +17518,10 @@ const docTemplate = `{
"plan_mode": {
"$ref": "#/definitions/codersdk.ChatPlanMode"
},
"queued_for_capacity": {
"description": "QueuedForCapacity reports that the chat is waiting for a concurrent\nagent slot. Single-chat reads derive it; list responses leave it false.",
"type": "boolean"
},
"root_chat_id": {
"type": "string",
"format": "uuid"
+4
View File
@@ -15755,6 +15755,10 @@
"plan_mode": {
"$ref": "#/definitions/codersdk.ChatPlanMode"
},
"queued_for_capacity": {
"description": "QueuedForCapacity reports that the chat is waiting for a concurrent\nagent slot. Single-chat reads derive it; list responses leave it false.",
"type": "boolean"
},
"root_chat_id": {
"type": "string",
"format": "uuid"
+3
View File
@@ -273,6 +273,8 @@ type Options struct {
// Set by enterprise for HA deployments. Nil uses chatd's local
// in-process channel dialer.
ChatStreamPartsDialer chatd.StreamPartsDialer
// Nil keeps the default chat agent caps active.
ChatAgentCapacityUnlock chatd.AgentCapacityUnlock
// ChatProviderAPIKeys overrides deployment-derived provider keys.
// Test harnesses use this to route chat models to local providers.
ChatProviderAPIKeys *chatprovider.ProviderAPIKeys
@@ -941,6 +943,7 @@ func New(options *Options) *API {
HookDispatcher: hookDispatcher,
UsageTracker: options.WorkspaceUsageTracker,
PrometheusRegistry: options.PrometheusRegistry,
AgentCapacityUnlock: options.ChatAgentCapacityUnlock,
OIDCTokenSource: oidcMCPSrc,
NotificationsEnqueuer: options.NotificationsEnqueuer,
Auditor: &api.Auditor,
+2 -5
View File
@@ -757,11 +757,8 @@ func TestChat_AllFieldsPopulated(t *testing.T) {
v := reflect.ValueOf(got)
typ := v.Type()
// HasUnread is populated by ChatRowsWithChildren (which joins the
// read-cursor query), not by Chat. Warnings is a transient
// field populated by handlers, not the converter. Both are
// expected to remain zero here.
skip := map[string]bool{"HasUnread": true, "Warnings": true}
// These fields are set outside db2sdk.Chat and intentionally remain zero.
skip := map[string]bool{"HasUnread": true, "Warnings": true, "QueuedForCapacity": true}
for i := range typ.NumField() {
field := typ.Field(i)
if skip[field.Name] {
+23
View File
@@ -1964,6 +1964,20 @@ func (q *querier) CountAuditLogs(ctx context.Context, arg database.CountAuditLog
return q.db.CountAuthorizedAuditLogs(ctx, arg, prep)
}
func (q *querier) CountChatCapacityActiveByPool(ctx context.Context, arg database.CountChatCapacityActiveByPoolParams) (database.CountChatCapacityActiveByPoolRow, error) {
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceChat); err != nil {
return database.CountChatCapacityActiveByPoolRow{}, err
}
return q.db.CountChatCapacityActiveByPool(ctx, arg)
}
func (q *querier) CountChatCapacityQueuedByPool(ctx context.Context, staleSeconds int32) (database.CountChatCapacityQueuedByPoolRow, error) {
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceChat); err != nil {
return database.CountChatCapacityQueuedByPoolRow{}, err
}
return q.db.CountChatCapacityQueuedByPool(ctx, staleSeconds)
}
func (q *querier) CountChatQueuedMessages(ctx context.Context, chatID uuid.UUID) (int64, error) {
_, err := q.GetChatByID(ctx, chatID)
if err != nil {
@@ -3483,6 +3497,15 @@ func (q *querier) GetChatPlanModeInstructions(ctx context.Context) (string, erro
return q.db.GetChatPlanModeInstructions(ctx)
}
func (q *querier) GetChatQueuedForCapacity(ctx context.Context, arg database.GetChatQueuedForCapacityParams) (bool, error) {
// The pool-fullness derivation counts other users' chats, so require
// deployment-wide chat read rather than per-chat authorization.
if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceChat); err != nil {
return false, err
}
return q.db.GetChatQueuedForCapacity(ctx, arg)
}
func (q *querier) GetChatQueuedMessageByID(ctx context.Context, arg database.GetChatQueuedMessageByIDParams) (database.ChatQueuedMessage, error) {
_, err := q.GetChatByID(ctx, arg.ChatID)
if err != nil {
+17
View File
@@ -608,6 +608,23 @@ func (s *MethodTestSuite) TestChats() {
dbm.EXPECT().GetChatWorkerAcquisitionCandidates(gomock.Any(), arg).Return([]database.GetChatWorkerAcquisitionCandidatesRow{row}, nil).AnyTimes()
check.Args(arg).Asserts(rbac.ResourceChat, policy.ActionUpdate).Returns([]database.GetChatWorkerAcquisitionCandidatesRow{row})
}))
s.Run("CountChatCapacityActiveByPool", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
arg := database.CountChatCapacityActiveByPoolParams{ExcludeChatID: uuid.New(), StaleSeconds: 30}
row := database.CountChatCapacityActiveByPoolRow{ActiveRootCount: 1, ActiveSubagentCount: 2}
dbm.EXPECT().CountChatCapacityActiveByPool(gomock.Any(), arg).Return(row, nil).AnyTimes()
check.Args(arg).Asserts(rbac.ResourceChat, policy.ActionRead).Returns(row)
}))
s.Run("CountChatCapacityQueuedByPool", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
const staleSeconds = int32(30)
row := database.CountChatCapacityQueuedByPoolRow{QueuedRootCount: 3, QueuedSubagentCount: 4}
dbm.EXPECT().CountChatCapacityQueuedByPool(gomock.Any(), staleSeconds).Return(row, nil).AnyTimes()
check.Args(staleSeconds).Asserts(rbac.ResourceChat, policy.ActionRead).Returns(row)
}))
s.Run("GetChatQueuedForCapacity", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
arg := database.GetChatQueuedForCapacityParams{ChatID: uuid.New(), StaleSeconds: 30, RootCapacity: 5, SubagentCapacity: 10}
dbm.EXPECT().GetChatQueuedForCapacity(gomock.Any(), arg).Return(true, nil).AnyTimes()
check.Args(arg).Asserts(rbac.ResourceChat, policy.ActionRead).Returns(true)
}))
s.Run("GetChatsByIDsForRunnerSync", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) {
ids := []uuid.UUID{uuid.New(), uuid.New()}
chat := testutil.Fake(s.T(), faker, database.Chat{ID: ids[0]})
+24
View File
@@ -321,6 +321,22 @@ func (m queryMetricsStore) CountAuditLogs(ctx context.Context, arg database.Coun
return r0, r1
}
func (m queryMetricsStore) CountChatCapacityActiveByPool(ctx context.Context, arg database.CountChatCapacityActiveByPoolParams) (database.CountChatCapacityActiveByPoolRow, error) {
start := time.Now()
r0, r1 := m.s.CountChatCapacityActiveByPool(ctx, arg)
m.queryLatencies.WithLabelValues("CountChatCapacityActiveByPool").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "CountChatCapacityActiveByPool").Inc()
return r0, r1
}
func (m queryMetricsStore) CountChatCapacityQueuedByPool(ctx context.Context, staleSeconds int32) (database.CountChatCapacityQueuedByPoolRow, error) {
start := time.Now()
r0, r1 := m.s.CountChatCapacityQueuedByPool(ctx, staleSeconds)
m.queryLatencies.WithLabelValues("CountChatCapacityQueuedByPool").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "CountChatCapacityQueuedByPool").Inc()
return r0, r1
}
func (m queryMetricsStore) CountChatQueuedMessages(ctx context.Context, chatID uuid.UUID) (int64, error) {
start := time.Now()
r0, r1 := m.s.CountChatQueuedMessages(ctx, chatID)
@@ -1705,6 +1721,14 @@ func (m queryMetricsStore) GetChatPlanModeInstructions(ctx context.Context) (str
return r0, r1
}
func (m queryMetricsStore) GetChatQueuedForCapacity(ctx context.Context, arg database.GetChatQueuedForCapacityParams) (bool, error) {
start := time.Now()
r0, r1 := m.s.GetChatQueuedForCapacity(ctx, arg)
m.queryLatencies.WithLabelValues("GetChatQueuedForCapacity").Observe(time.Since(start).Seconds())
m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetChatQueuedForCapacity").Inc()
return r0, r1
}
func (m queryMetricsStore) GetChatQueuedMessageByID(ctx context.Context, arg database.GetChatQueuedMessageByIDParams) (database.ChatQueuedMessage, error) {
start := time.Now()
r0, r1 := m.s.GetChatQueuedMessageByID(ctx, arg)
+45
View File
@@ -483,6 +483,36 @@ func (mr *MockStoreMockRecorder) CountAuthorizedConnectionLogs(ctx, arg, prepare
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CountAuthorizedConnectionLogs", reflect.TypeOf((*MockStore)(nil).CountAuthorizedConnectionLogs), ctx, arg, prepared)
}
// CountChatCapacityActiveByPool mocks base method.
func (m *MockStore) CountChatCapacityActiveByPool(ctx context.Context, arg database.CountChatCapacityActiveByPoolParams) (database.CountChatCapacityActiveByPoolRow, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "CountChatCapacityActiveByPool", ctx, arg)
ret0, _ := ret[0].(database.CountChatCapacityActiveByPoolRow)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// CountChatCapacityActiveByPool indicates an expected call of CountChatCapacityActiveByPool.
func (mr *MockStoreMockRecorder) CountChatCapacityActiveByPool(ctx, arg any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CountChatCapacityActiveByPool", reflect.TypeOf((*MockStore)(nil).CountChatCapacityActiveByPool), ctx, arg)
}
// CountChatCapacityQueuedByPool mocks base method.
func (m *MockStore) CountChatCapacityQueuedByPool(ctx context.Context, staleSeconds int32) (database.CountChatCapacityQueuedByPoolRow, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "CountChatCapacityQueuedByPool", ctx, staleSeconds)
ret0, _ := ret[0].(database.CountChatCapacityQueuedByPoolRow)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// CountChatCapacityQueuedByPool indicates an expected call of CountChatCapacityQueuedByPool.
func (mr *MockStoreMockRecorder) CountChatCapacityQueuedByPool(ctx, staleSeconds any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CountChatCapacityQueuedByPool", reflect.TypeOf((*MockStore)(nil).CountChatCapacityQueuedByPool), ctx, staleSeconds)
}
// CountChatQueuedMessages mocks base method.
func (m *MockStore) CountChatQueuedMessages(ctx context.Context, chatID uuid.UUID) (int64, error) {
m.ctrl.T.Helper()
@@ -3165,6 +3195,21 @@ func (mr *MockStoreMockRecorder) GetChatPlanModeInstructions(ctx any) *gomock.Ca
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatPlanModeInstructions", reflect.TypeOf((*MockStore)(nil).GetChatPlanModeInstructions), ctx)
}
// GetChatQueuedForCapacity mocks base method.
func (m *MockStore) GetChatQueuedForCapacity(ctx context.Context, arg database.GetChatQueuedForCapacityParams) (bool, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetChatQueuedForCapacity", ctx, arg)
ret0, _ := ret[0].(bool)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetChatQueuedForCapacity indicates an expected call of GetChatQueuedForCapacity.
func (mr *MockStoreMockRecorder) GetChatQueuedForCapacity(ctx, arg any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetChatQueuedForCapacity", reflect.TypeOf((*MockStore)(nil).GetChatQueuedForCapacity), ctx, arg)
}
// GetChatQueuedMessageByID mocks base method.
func (m *MockStore) GetChatQueuedMessageByID(ctx context.Context, arg database.GetChatQueuedMessageByIDParams) (database.ChatQueuedMessage, error) {
m.ctrl.T.Helper()
+1 -1
View File
@@ -4842,7 +4842,7 @@ CREATE INDEX idx_chats_title_fts ON chats USING gin (to_tsvector('simple'::regco
COMMENT ON INDEX idx_chats_title_fts IS 'Used for full text search. Defined over all rows of the chats table.';
CREATE INDEX idx_chats_worker_acquisition_candidates ON chats USING btree (status, updated_at, id) WHERE (archived = false);
CREATE INDEX idx_chats_worker_acquisition_candidates ON chats USING btree (((parent_chat_id IS NULL)), status, updated_at, id) WHERE (archived = false);
CREATE INDEX idx_chats_workspace ON chats USING btree (workspace_id);
+1
View File
@@ -17,6 +17,7 @@ const (
LockIDBoundaryUsageStats
LockIDAIProvidersEnvSeed
LockIDChatModelConfigWrites
LockIDChatCapacityAdmission
)
// GenLockID generates a unique and consistent lock ID from a given string.
@@ -0,0 +1,4 @@
DROP INDEX idx_chats_worker_acquisition_candidates;
CREATE INDEX idx_chats_worker_acquisition_candidates ON chats
(status, updated_at, id)
WHERE archived = false;
@@ -0,0 +1,4 @@
DROP INDEX idx_chats_worker_acquisition_candidates;
CREATE INDEX idx_chats_worker_acquisition_candidates ON chats
((parent_chat_id IS NULL), status, updated_at, id)
WHERE archived = false;
+8 -11
View File
@@ -93,6 +93,9 @@ type sqlcQuerier interface {
CleanupDeletedMCPServerIDsFromChats(ctx context.Context) error
CountAIBridgeSessions(ctx context.Context, arg CountAIBridgeSessionsParams) (int64, error)
CountAuditLogs(ctx context.Context, arg CountAuditLogsParams) (int64, error)
// Excluding the candidate keeps ownership takeover capacity-neutral.
CountChatCapacityActiveByPool(ctx context.Context, arg CountChatCapacityActiveByPoolParams) (CountChatCapacityActiveByPoolRow, error)
CountChatCapacityQueuedByPool(ctx context.Context, staleSeconds int32) (CountChatCapacityQueuedByPoolRow, error)
// Cheap queue-length check used by ChatMachine.Update when deciding
// whether the chat is in a "1" sub-state.
CountChatQueuedMessages(ctx context.Context, chatID uuid.UUID) (int64, error)
@@ -477,6 +480,8 @@ type sqlcQuerier interface {
// personal chat model overrides. It defaults to false when unset.
GetChatPersonalModelOverridesEnabled(ctx context.Context) (bool, error)
GetChatPlanModeInstructions(ctx context.Context) (string, error)
// Pool fullness distinguishes capacity waits from worker pickup delays.
GetChatQueuedForCapacity(ctx context.Context, arg GetChatQueuedForCapacityParams) (bool, error)
GetChatQueuedMessageByID(ctx context.Context, arg GetChatQueuedMessageByIDParams) (ChatQueuedMessage, error)
// Returns the queue head (lowest position, then lowest id).
GetChatQueuedMessageHead(ctx context.Context, chatID uuid.UUID) (ChatQueuedMessage, error)
@@ -507,17 +512,9 @@ type sqlcQuerier interface {
// jsonb_array_elements never raises "cannot extract elements from a
// scalar". Backed by idx_chat_messages_user_prompts.
GetChatUserPromptsByChatID(ctx context.Context, arg GetChatUserPromptsByChatIDParams) ([]GetChatUserPromptsByChatIDRow, error)
// Returns chats that workers may try to acquire. Candidates must be:
// - in a worker-runnable execution status;
// - unarchived; and
// - missing ownership, carrying inconsistent ownership, or lacking a
// fresh heartbeat for the assigned runner.
//
// Missing ownership is worker_id IS NULL. Inconsistent ownership is
// runner_id IS NULL while worker_id is set. Stale ownership is no
// heartbeat row for (chat_id, runner_id), or one older than
// @stale_seconds by database time. Candidates are ordered by oldest
// updated_at first so workers drain stale runnable chats predictably.
// Returns a bounded, pool-interleaved set of chats that workers may acquire.
// Interrupting chats finish active work first. Requires-action chats follow so
// their runner can enforce the action deadline before new generations start.
GetChatWorkerAcquisitionCandidates(ctx context.Context, arg GetChatWorkerAcquisitionCandidatesParams) ([]GetChatWorkerAcquisitionCandidatesRow, error)
// Returns the global TTL for chat workspaces as a Go duration string.
// Returns "0s" (disabled) when no value has been configured.
+183 -142
View File
@@ -7048,6 +7048,69 @@ func (q *sqlQuerier) BatchUpsertChatHeartbeats(ctx context.Context, arg BatchUps
return err
}
const countChatCapacityActiveByPool = `-- name: CountChatCapacityActiveByPool :one
SELECT
COUNT(*) FILTER (WHERE c.parent_chat_id IS NULL)::bigint AS active_root_count,
COUNT(*) FILTER (WHERE c.parent_chat_id IS NOT NULL)::bigint AS active_subagent_count
FROM chat_heartbeats hb
JOIN chats c
ON c.id = hb.chat_id
AND c.runner_id = hb.runner_id
WHERE c.worker_id IS NOT NULL
AND c.id != $1::uuid
AND hb.heartbeat_at > NOW() - (INTERVAL '1 second' * $2::int)
`
type CountChatCapacityActiveByPoolParams struct {
ExcludeChatID uuid.UUID `db:"exclude_chat_id" json:"exclude_chat_id"`
StaleSeconds int32 `db:"stale_seconds" json:"stale_seconds"`
}
type CountChatCapacityActiveByPoolRow struct {
ActiveRootCount int64 `db:"active_root_count" json:"active_root_count"`
ActiveSubagentCount int64 `db:"active_subagent_count" json:"active_subagent_count"`
}
// Excluding the candidate keeps ownership takeover capacity-neutral.
func (q *sqlQuerier) CountChatCapacityActiveByPool(ctx context.Context, arg CountChatCapacityActiveByPoolParams) (CountChatCapacityActiveByPoolRow, error) {
row := q.db.QueryRowContext(ctx, countChatCapacityActiveByPool, arg.ExcludeChatID, arg.StaleSeconds)
var i CountChatCapacityActiveByPoolRow
err := row.Scan(&i.ActiveRootCount, &i.ActiveSubagentCount)
return i, err
}
const countChatCapacityQueuedByPool = `-- name: CountChatCapacityQueuedByPool :one
SELECT
COUNT(*) FILTER (WHERE c.parent_chat_id IS NULL)::bigint AS queued_root_count,
COUNT(*) FILTER (WHERE c.parent_chat_id IS NOT NULL)::bigint AS queued_subagent_count
FROM chats c
WHERE c.status = 'running'::chat_status
AND c.archived = false
AND (
c.worker_id IS NULL
OR c.runner_id IS NULL
OR NOT EXISTS (
SELECT 1
FROM chat_heartbeats hb
WHERE hb.chat_id = c.id
AND hb.runner_id = c.runner_id
AND hb.heartbeat_at > NOW() - (INTERVAL '1 second' * $1::int)
)
)
`
type CountChatCapacityQueuedByPoolRow struct {
QueuedRootCount int64 `db:"queued_root_count" json:"queued_root_count"`
QueuedSubagentCount int64 `db:"queued_subagent_count" json:"queued_subagent_count"`
}
func (q *sqlQuerier) CountChatCapacityQueuedByPool(ctx context.Context, staleSeconds int32) (CountChatCapacityQueuedByPoolRow, error) {
row := q.db.QueryRowContext(ctx, countChatCapacityQueuedByPool, staleSeconds)
var i CountChatCapacityQueuedByPoolRow
err := row.Scan(&i.QueuedRootCount, &i.QueuedSubagentCount)
return i, err
}
const countChatQueuedMessages = `-- name: CountChatQueuedMessages :one
SELECT COUNT(*)::bigint AS count
FROM chat_queued_messages
@@ -8507,6 +8570,62 @@ func (q *sqlQuerier) GetChatModelConfigsForTelemetry(ctx context.Context) ([]Get
return items, nil
}
const getChatQueuedForCapacity = `-- name: GetChatQueuedForCapacity :one
WITH active AS (
SELECT
COUNT(*) FILTER (WHERE a.parent_chat_id IS NULL)::bigint AS root_count,
COUNT(*) FILTER (WHERE a.parent_chat_id IS NOT NULL)::bigint AS subagent_count
FROM chat_heartbeats hb
JOIN chats a
ON a.id = hb.chat_id
AND a.runner_id = hb.runner_id
WHERE a.worker_id IS NOT NULL
AND hb.heartbeat_at > NOW() - (INTERVAL '1 second' * $1::int)
)
SELECT (
c.status = 'running'::chat_status
AND c.archived = false
AND (
c.worker_id IS NULL
OR c.runner_id IS NULL
OR NOT EXISTS (
SELECT 1
FROM chat_heartbeats hb
WHERE hb.chat_id = c.id
AND hb.runner_id = c.runner_id
AND hb.heartbeat_at > NOW() - (INTERVAL '1 second' * $1::int)
)
)
AND CASE
WHEN c.parent_chat_id IS NULL THEN active.root_count >= $2::bigint
ELSE active.subagent_count >= $3::bigint
END
)::boolean AS queued_for_capacity
FROM chats c
CROSS JOIN active
WHERE c.id = $4::uuid
`
type GetChatQueuedForCapacityParams struct {
StaleSeconds int32 `db:"stale_seconds" json:"stale_seconds"`
RootCapacity int64 `db:"root_capacity" json:"root_capacity"`
SubagentCapacity int64 `db:"subagent_capacity" json:"subagent_capacity"`
ChatID uuid.UUID `db:"chat_id" json:"chat_id"`
}
// Pool fullness distinguishes capacity waits from worker pickup delays.
func (q *sqlQuerier) GetChatQueuedForCapacity(ctx context.Context, arg GetChatQueuedForCapacityParams) (bool, error) {
row := q.db.QueryRowContext(ctx, getChatQueuedForCapacity,
arg.StaleSeconds,
arg.RootCapacity,
arg.SubagentCapacity,
arg.ChatID,
)
var queued_for_capacity bool
err := row.Scan(&queued_for_capacity)
return queued_for_capacity, err
}
const getChatQueuedMessageByID = `-- name: GetChatQueuedMessageByID :one
SELECT id, chat_id, content, created_at, model_config_id, position, created_by, reasoning_effort FROM chat_queued_messages
WHERE id = $1::bigint AND chat_id = $2::uuid
@@ -8759,108 +8878,80 @@ func (q *sqlQuerier) GetChatUserPromptsByChatID(ctx context.Context, arg GetChat
}
const getChatWorkerAcquisitionCandidates = `-- name: GetChatWorkerAcquisitionCandidates :many
WITH candidate_partitions AS (
SELECT true AS is_root, 'interrupting'::chat_status AS status, 0 AS status_priority, 0 AS pool_priority
UNION ALL
SELECT false, 'interrupting'::chat_status, 0, 1
UNION ALL
SELECT true, 'requires_action'::chat_status, 1, 0
UNION ALL
SELECT false, 'requires_action'::chat_status, 1, 1
UNION ALL
SELECT true, 'running'::chat_status, 2, 0
UNION ALL
SELECT false, 'running'::chat_status, 2, 1
),
candidates AS (
SELECT
candidate.id,
candidate_partitions.status_priority,
candidate_partitions.pool_priority,
ROW_NUMBER() OVER (
PARTITION BY candidate_partitions.status_priority, candidate_partitions.is_root
ORDER BY candidate.updated_at ASC, candidate.id ASC
) AS pool_position
FROM candidate_partitions
CROSS JOIN LATERAL (
SELECT chats.id, chats.updated_at
FROM chats
WHERE (chats.parent_chat_id IS NULL) = candidate_partitions.is_root
AND chats.status = candidate_partitions.status
AND chats.archived = false
AND (
chats.worker_id IS NULL
OR chats.runner_id IS NULL
OR NOT EXISTS (
SELECT 1
FROM chat_heartbeats current_lease
WHERE current_lease.chat_id = chats.id
AND current_lease.runner_id = chats.runner_id
AND current_lease.heartbeat_at > NOW() - (INTERVAL '1 second' * $2::int)
)
)
ORDER BY chats.updated_at ASC, chats.id ASC
LIMIT $1::int
) candidate
)
SELECT
chats_expanded.id, chats_expanded.owner_id, chats_expanded.workspace_id, chats_expanded.title, chats_expanded.status, chats_expanded.worker_id, chats_expanded.started_at, chats_expanded.heartbeat_at, chats_expanded.created_at, chats_expanded.updated_at, chats_expanded.parent_chat_id, chats_expanded.root_chat_id, chats_expanded.last_model_config_id, chats_expanded.last_reasoning_effort, chats_expanded.archived, chats_expanded.last_error, chats_expanded.mode, chats_expanded.mcp_server_ids, chats_expanded.labels, chats_expanded.build_id, chats_expanded.agent_id, chats_expanded.pin_order, chats_expanded.last_read_message_id, chats_expanded.dynamic_tools, chats_expanded.organization_id, chats_expanded.plan_mode, chats_expanded.client_type, chats_expanded.last_turn_summary, chats_expanded.summary, chats_expanded.summary_generated_at, chats_expanded.snapshot_version, chats_expanded.history_version, chats_expanded.queue_version, chats_expanded.generation_attempt, chats_expanded.retry_state, chats_expanded.retry_state_version, chats_expanded.runner_id, chats_expanded.requires_action_deadline_at, chats_expanded.user_acl, chats_expanded.group_acl, chats_expanded.owner_username, chats_expanded.owner_name, chats_expanded.context_aggregate_hash, chats_expanded.context_dirty_since, chats_expanded.context_dirty_resources, chats_expanded.context_error, chats_expanded.compaction_requested_at,
chat_heartbeats.heartbeat_at AS current_heartbeat_at,
NOT EXISTS (
SELECT 1
FROM chat_heartbeats current_lease
WHERE current_lease.chat_id = chats_expanded.id
AND current_lease.runner_id = chats_expanded.runner_id
AND current_lease.heartbeat_at > NOW() - (INTERVAL '1 second' * $1::int)
) AS heartbeat_stale
FROM chats_expanded
LEFT JOIN chat_heartbeats
ON chat_heartbeats.chat_id = chats_expanded.id
AND chat_heartbeats.runner_id = chats_expanded.runner_id
WHERE
chats_expanded.status IN ('running'::chat_status, 'interrupting'::chat_status, 'requires_action'::chat_status)
AND chats_expanded.archived = false
AND (
chats_expanded.worker_id IS NULL
OR chats_expanded.runner_id IS NULL
OR NOT EXISTS (
SELECT 1
FROM chat_heartbeats current_lease
WHERE current_lease.chat_id = chats_expanded.id
AND current_lease.runner_id = chats_expanded.runner_id
AND current_lease.heartbeat_at > NOW() - (INTERVAL '1 second' * $1::int)
)
)
ORDER BY chats_expanded.updated_at ASC, chats_expanded.id ASC
LIMIT $2::int
chats.id,
chats.status,
chats.parent_chat_id
FROM candidates
JOIN chats ON chats.id = candidates.id
ORDER BY
candidates.status_priority ASC,
candidates.pool_position ASC,
candidates.pool_priority ASC,
chats.id ASC
LIMIT $1::int
`
type GetChatWorkerAcquisitionCandidatesParams struct {
StaleSeconds int32 `db:"stale_seconds" json:"stale_seconds"`
LimitCount int32 `db:"limit_count" json:"limit_count"`
StaleSeconds int32 `db:"stale_seconds" json:"stale_seconds"`
}
type GetChatWorkerAcquisitionCandidatesRow struct {
ID uuid.UUID `db:"id" json:"id"`
OwnerID uuid.UUID `db:"owner_id" json:"owner_id"`
WorkspaceID uuid.NullUUID `db:"workspace_id" json:"workspace_id"`
Title string `db:"title" json:"title"`
Status ChatStatus `db:"status" json:"status"`
WorkerID uuid.NullUUID `db:"worker_id" json:"worker_id"`
StartedAt sql.NullTime `db:"started_at" json:"started_at"`
HeartbeatAt sql.NullTime `db:"heartbeat_at" json:"heartbeat_at"`
CreatedAt time.Time `db:"created_at" json:"created_at"`
UpdatedAt time.Time `db:"updated_at" json:"updated_at"`
ParentChatID uuid.NullUUID `db:"parent_chat_id" json:"parent_chat_id"`
RootChatID uuid.NullUUID `db:"root_chat_id" json:"root_chat_id"`
LastModelConfigID uuid.UUID `db:"last_model_config_id" json:"last_model_config_id"`
LastReasoningEffort NullChatReasoningEffort `db:"last_reasoning_effort" json:"last_reasoning_effort"`
Archived bool `db:"archived" json:"archived"`
LastError pqtype.NullRawMessage `db:"last_error" json:"last_error"`
Mode NullChatMode `db:"mode" json:"mode"`
MCPServerIDs []uuid.UUID `db:"mcp_server_ids" json:"mcp_server_ids"`
Labels StringMap `db:"labels" json:"labels"`
BuildID uuid.NullUUID `db:"build_id" json:"build_id"`
AgentID uuid.NullUUID `db:"agent_id" json:"agent_id"`
PinOrder int32 `db:"pin_order" json:"pin_order"`
LastReadMessageID sql.NullInt64 `db:"last_read_message_id" json:"last_read_message_id"`
DynamicTools pqtype.NullRawMessage `db:"dynamic_tools" json:"dynamic_tools"`
OrganizationID uuid.UUID `db:"organization_id" json:"organization_id"`
PlanMode NullChatPlanMode `db:"plan_mode" json:"plan_mode"`
ClientType ChatClientType `db:"client_type" json:"client_type"`
LastTurnSummary sql.NullString `db:"last_turn_summary" json:"last_turn_summary"`
Summary sql.NullString `db:"summary" json:"summary"`
SummaryGeneratedAt sql.NullTime `db:"summary_generated_at" json:"summary_generated_at"`
SnapshotVersion int64 `db:"snapshot_version" json:"snapshot_version"`
HistoryVersion int64 `db:"history_version" json:"history_version"`
QueueVersion int64 `db:"queue_version" json:"queue_version"`
GenerationAttempt int64 `db:"generation_attempt" json:"generation_attempt"`
RetryState pqtype.NullRawMessage `db:"retry_state" json:"retry_state"`
RetryStateVersion int64 `db:"retry_state_version" json:"retry_state_version"`
RunnerID uuid.NullUUID `db:"runner_id" json:"runner_id"`
RequiresActionDeadlineAt sql.NullTime `db:"requires_action_deadline_at" json:"requires_action_deadline_at"`
UserACL ChatACL `db:"user_acl" json:"user_acl"`
GroupACL ChatACL `db:"group_acl" json:"group_acl"`
OwnerUsername string `db:"owner_username" json:"owner_username"`
OwnerName string `db:"owner_name" json:"owner_name"`
ContextAggregateHash []byte `db:"context_aggregate_hash" json:"context_aggregate_hash"`
ContextDirtySince sql.NullTime `db:"context_dirty_since" json:"context_dirty_since"`
ContextDirtyResources pqtype.NullRawMessage `db:"context_dirty_resources" json:"context_dirty_resources"`
ContextError string `db:"context_error" json:"context_error"`
CompactionRequestedAt sql.NullTime `db:"compaction_requested_at" json:"compaction_requested_at"`
CurrentHeartbeatAt sql.NullTime `db:"current_heartbeat_at" json:"current_heartbeat_at"`
HeartbeatStale bool `db:"heartbeat_stale" json:"heartbeat_stale"`
ID uuid.UUID `db:"id" json:"id"`
Status ChatStatus `db:"status" json:"status"`
ParentChatID uuid.NullUUID `db:"parent_chat_id" json:"parent_chat_id"`
}
// Returns chats that workers may try to acquire. Candidates must be:
// - in a worker-runnable execution status;
// - unarchived; and
// - missing ownership, carrying inconsistent ownership, or lacking a
// fresh heartbeat for the assigned runner.
//
// Missing ownership is worker_id IS NULL. Inconsistent ownership is
// runner_id IS NULL while worker_id is set. Stale ownership is no
// heartbeat row for (chat_id, runner_id), or one older than
// @stale_seconds by database time. Candidates are ordered by oldest
// updated_at first so workers drain stale runnable chats predictably.
// Returns a bounded, pool-interleaved set of chats that workers may acquire.
// Interrupting chats finish active work first. Requires-action chats follow so
// their runner can enforce the action deadline before new generations start.
func (q *sqlQuerier) GetChatWorkerAcquisitionCandidates(ctx context.Context, arg GetChatWorkerAcquisitionCandidatesParams) ([]GetChatWorkerAcquisitionCandidatesRow, error) {
rows, err := q.db.QueryContext(ctx, getChatWorkerAcquisitionCandidates, arg.StaleSeconds, arg.LimitCount)
rows, err := q.db.QueryContext(ctx, getChatWorkerAcquisitionCandidates, arg.LimitCount, arg.StaleSeconds)
if err != nil {
return nil, err
}
@@ -8868,57 +8959,7 @@ func (q *sqlQuerier) GetChatWorkerAcquisitionCandidates(ctx context.Context, arg
var items []GetChatWorkerAcquisitionCandidatesRow
for rows.Next() {
var i GetChatWorkerAcquisitionCandidatesRow
if err := rows.Scan(
&i.ID,
&i.OwnerID,
&i.WorkspaceID,
&i.Title,
&i.Status,
&i.WorkerID,
&i.StartedAt,
&i.HeartbeatAt,
&i.CreatedAt,
&i.UpdatedAt,
&i.ParentChatID,
&i.RootChatID,
&i.LastModelConfigID,
&i.LastReasoningEffort,
&i.Archived,
&i.LastError,
&i.Mode,
pq.Array(&i.MCPServerIDs),
&i.Labels,
&i.BuildID,
&i.AgentID,
&i.PinOrder,
&i.LastReadMessageID,
&i.DynamicTools,
&i.OrganizationID,
&i.PlanMode,
&i.ClientType,
&i.LastTurnSummary,
&i.Summary,
&i.SummaryGeneratedAt,
&i.SnapshotVersion,
&i.HistoryVersion,
&i.QueueVersion,
&i.GenerationAttempt,
&i.RetryState,
&i.RetryStateVersion,
&i.RunnerID,
&i.RequiresActionDeadlineAt,
&i.UserACL,
&i.GroupACL,
&i.OwnerUsername,
&i.OwnerName,
&i.ContextAggregateHash,
&i.ContextDirtySince,
&i.ContextDirtyResources,
&i.ContextError,
&i.CompactionRequestedAt,
&i.CurrentHeartbeatAt,
&i.HeartbeatStale,
); err != nil {
if err := rows.Scan(&i.ID, &i.Status, &i.ParentChatID); err != nil {
return nil, err
}
items = append(items, i)
+125 -39
View File
@@ -2303,46 +2303,64 @@ WHERE chat_id = @chat_id::uuid
AND content::jsonb @> '[{"type": "context-file"}]';
-- name: GetChatWorkerAcquisitionCandidates :many
-- Returns chats that workers may try to acquire. Candidates must be:
-- - in a worker-runnable execution status;
-- - unarchived; and
-- - missing ownership, carrying inconsistent ownership, or lacking a
-- fresh heartbeat for the assigned runner.
--
-- Missing ownership is worker_id IS NULL. Inconsistent ownership is
-- runner_id IS NULL while worker_id is set. Stale ownership is no
-- heartbeat row for (chat_id, runner_id), or one older than
-- @stale_seconds by database time. Candidates are ordered by oldest
-- updated_at first so workers drain stale runnable chats predictably.
-- Returns a bounded, pool-interleaved set of chats that workers may acquire.
-- Interrupting chats finish active work first. Requires-action chats follow so
-- their runner can enforce the action deadline before new generations start.
WITH candidate_partitions AS (
SELECT true AS is_root, 'interrupting'::chat_status AS status, 0 AS status_priority, 0 AS pool_priority
UNION ALL
SELECT false, 'interrupting'::chat_status, 0, 1
UNION ALL
SELECT true, 'requires_action'::chat_status, 1, 0
UNION ALL
SELECT false, 'requires_action'::chat_status, 1, 1
UNION ALL
SELECT true, 'running'::chat_status, 2, 0
UNION ALL
SELECT false, 'running'::chat_status, 2, 1
),
candidates AS (
SELECT
candidate.id,
candidate_partitions.status_priority,
candidate_partitions.pool_priority,
ROW_NUMBER() OVER (
PARTITION BY candidate_partitions.status_priority, candidate_partitions.is_root
ORDER BY candidate.updated_at ASC, candidate.id ASC
) AS pool_position
FROM candidate_partitions
CROSS JOIN LATERAL (
SELECT chats.id, chats.updated_at
FROM chats
WHERE (chats.parent_chat_id IS NULL) = candidate_partitions.is_root
AND chats.status = candidate_partitions.status
AND chats.archived = false
AND (
chats.worker_id IS NULL
OR chats.runner_id IS NULL
OR NOT EXISTS (
SELECT 1
FROM chat_heartbeats current_lease
WHERE current_lease.chat_id = chats.id
AND current_lease.runner_id = chats.runner_id
AND current_lease.heartbeat_at > NOW() - (INTERVAL '1 second' * @stale_seconds::int)
)
)
ORDER BY chats.updated_at ASC, chats.id ASC
LIMIT @limit_count::int
) candidate
)
SELECT
chats_expanded.*,
chat_heartbeats.heartbeat_at AS current_heartbeat_at,
NOT EXISTS (
SELECT 1
FROM chat_heartbeats current_lease
WHERE current_lease.chat_id = chats_expanded.id
AND current_lease.runner_id = chats_expanded.runner_id
AND current_lease.heartbeat_at > NOW() - (INTERVAL '1 second' * @stale_seconds::int)
) AS heartbeat_stale
FROM chats_expanded
LEFT JOIN chat_heartbeats
ON chat_heartbeats.chat_id = chats_expanded.id
AND chat_heartbeats.runner_id = chats_expanded.runner_id
WHERE
chats_expanded.status IN ('running'::chat_status, 'interrupting'::chat_status, 'requires_action'::chat_status)
AND chats_expanded.archived = false
AND (
chats_expanded.worker_id IS NULL
OR chats_expanded.runner_id IS NULL
OR NOT EXISTS (
SELECT 1
FROM chat_heartbeats current_lease
WHERE current_lease.chat_id = chats_expanded.id
AND current_lease.runner_id = chats_expanded.runner_id
AND current_lease.heartbeat_at > NOW() - (INTERVAL '1 second' * @stale_seconds::int)
)
)
ORDER BY chats_expanded.updated_at ASC, chats_expanded.id ASC
chats.id,
chats.status,
chats.parent_chat_id
FROM candidates
JOIN chats ON chats.id = candidates.id
ORDER BY
candidates.status_priority ASC,
candidates.pool_position ASC,
candidates.pool_priority ASC,
chats.id ASC
LIMIT @limit_count::int;
-- name: GetChatsByIDsForRunnerSync :many
@@ -2800,3 +2818,71 @@ LEFT JOIN to_archive t ON t.id = a.id
-- created_at ASC flows through to dbpurge's digest truncation; see
-- buildDigestData in dbpurge.go for the tradeoff rationale.
ORDER BY (a.root_chat_id IS NULL) DESC, a.owner_id ASC, a.created_at ASC, a.id ASC;
-- name: CountChatCapacityActiveByPool :one
-- Excluding the candidate keeps ownership takeover capacity-neutral.
SELECT
COUNT(*) FILTER (WHERE c.parent_chat_id IS NULL)::bigint AS active_root_count,
COUNT(*) FILTER (WHERE c.parent_chat_id IS NOT NULL)::bigint AS active_subagent_count
FROM chat_heartbeats hb
JOIN chats c
ON c.id = hb.chat_id
AND c.runner_id = hb.runner_id
WHERE c.worker_id IS NOT NULL
AND c.id != @exclude_chat_id::uuid
AND hb.heartbeat_at > NOW() - (INTERVAL '1 second' * @stale_seconds::int);
-- name: CountChatCapacityQueuedByPool :one
SELECT
COUNT(*) FILTER (WHERE c.parent_chat_id IS NULL)::bigint AS queued_root_count,
COUNT(*) FILTER (WHERE c.parent_chat_id IS NOT NULL)::bigint AS queued_subagent_count
FROM chats c
WHERE c.status = 'running'::chat_status
AND c.archived = false
AND (
c.worker_id IS NULL
OR c.runner_id IS NULL
OR NOT EXISTS (
SELECT 1
FROM chat_heartbeats hb
WHERE hb.chat_id = c.id
AND hb.runner_id = c.runner_id
AND hb.heartbeat_at > NOW() - (INTERVAL '1 second' * @stale_seconds::int)
)
);
-- name: GetChatQueuedForCapacity :one
-- Pool fullness distinguishes capacity waits from worker pickup delays.
WITH active AS (
SELECT
COUNT(*) FILTER (WHERE a.parent_chat_id IS NULL)::bigint AS root_count,
COUNT(*) FILTER (WHERE a.parent_chat_id IS NOT NULL)::bigint AS subagent_count
FROM chat_heartbeats hb
JOIN chats a
ON a.id = hb.chat_id
AND a.runner_id = hb.runner_id
WHERE a.worker_id IS NOT NULL
AND hb.heartbeat_at > NOW() - (INTERVAL '1 second' * @stale_seconds::int)
)
SELECT (
c.status = 'running'::chat_status
AND c.archived = false
AND (
c.worker_id IS NULL
OR c.runner_id IS NULL
OR NOT EXISTS (
SELECT 1
FROM chat_heartbeats hb
WHERE hb.chat_id = c.id
AND hb.runner_id = c.runner_id
AND hb.heartbeat_at > NOW() - (INTERVAL '1 second' * @stale_seconds::int)
)
)
AND CASE
WHEN c.parent_chat_id IS NULL THEN active.root_count >= @root_capacity::bigint
ELSE active.subagent_count >= @subagent_capacity::bigint
END
)::boolean AS queued_for_capacity
FROM chats c
CROSS JOIN active
WHERE c.id = @chat_id::uuid;
+12
View File
@@ -1621,6 +1621,18 @@ func (api *API) getChat(rw http.ResponseWriter, r *http.Request) {
sdkChat := db2sdk.Chat(chat, diffStatus, chatFiles)
if api.chatDaemon != nil {
queued, err := api.chatDaemon.ChatQueuedForCapacity(ctx, chat)
if err != nil {
api.Logger.Error(ctx, "failed to derive chat queued-for-capacity state",
slog.F("chat_id", chat.ID),
slog.Error(err),
)
} else {
sdkChat.QueuedForCapacity = queued
}
}
// Enrich the lightweight context summary with the chat's pinned
// resources (metadata only). This detail is computed on read and only
// attached on the single-chat GET; list and watch payloads stay
+4
View File
@@ -947,6 +947,10 @@ The abandon chat goroutine is responsible for abandoning the chat. It is spawned
When the manager cleans up a runner, the runner must cancel all goroutines it has spawned and unsubscribe from pubsub.
## Concurrent agent limiter
By default, chatd runs up to five top-level chats and ten subagent chats at once. Each limit applies across the entire deployment. Enterprise deployments can remove these limits when their plan permits it. Extra chats wait for capacity, but users can still interrupt active chats.
## Auto-archive loop
The worker periodically archives old, unused chats.
+85
View File
@@ -0,0 +1,85 @@
package chatd
import (
"context"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/database/dbauthz"
)
const (
defaultMaxConcurrentRootAgents = int64(5)
defaultMaxConcurrentSubagents = int64(10)
)
// AgentCapacityLimiter controls chat admission and reports the current per-pool limits.
type AgentCapacityLimiter interface {
// Admit runs inside the acquisition transaction so its serialization
// extends through the ownership write. Refused chats remain unowned.
Admit(ctx context.Context, store database.Store, chat database.Chat) (bool, error)
Limits() (limits AgentCapacityLimits, capped bool)
}
// AgentCapacityUnlock reports whether the default chat agent caps are disabled.
type AgentCapacityUnlock interface {
Unlocked() bool
}
// AgentCapacityLimits defines concurrent-agent limits for root and subagent pools.
type AgentCapacityLimits struct {
Root int64
Subagent int64
}
type agentCapacityLimiter struct {
unlock AgentCapacityUnlock
staleSeconds int32
rootCapacity int64
subagentCapacity int64
}
func newAgentCapacityLimiter(unlock AgentCapacityUnlock, staleSeconds int32) *agentCapacityLimiter {
return &agentCapacityLimiter{
unlock: unlock,
staleSeconds: staleSeconds,
rootCapacity: defaultMaxConcurrentRootAgents,
subagentCapacity: defaultMaxConcurrentSubagents,
}
}
func (a *agentCapacityLimiter) Admit(ctx context.Context, store database.Store, chat database.Chat) (bool, error) {
//nolint:gocritic // Capacity accounting is chatd-internal state.
ctx = dbauthz.AsChatd(ctx)
if a.unlocked() || chat.Status != database.ChatStatusRunning {
return true, nil
}
// The transaction lock remains held through the caller's ownership write,
// preventing replicas from over-admitting the pool.
if err := store.AcquireLock(ctx, database.LockIDChatCapacityAdmission); err != nil {
return false, err
}
counts, err := store.CountChatCapacityActiveByPool(ctx, database.CountChatCapacityActiveByPoolParams{
ExcludeChatID: chat.ID,
StaleSeconds: a.staleSeconds,
})
if err != nil {
return false, err
}
used, capacity := counts.ActiveRootCount, a.rootCapacity
if chat.ParentChatID.Valid {
used, capacity = counts.ActiveSubagentCount, a.subagentCapacity
}
return used < capacity, nil
}
func (a *agentCapacityLimiter) Limits() (AgentCapacityLimits, bool) {
return AgentCapacityLimits{
Root: a.rootCapacity,
Subagent: a.subagentCapacity,
}, !a.unlocked()
}
func (a *agentCapacityLimiter) unlocked() bool {
return a.unlock != nil && a.unlock.Unlocked()
}
@@ -0,0 +1,487 @@
package chatd
import (
"context"
"sync"
"testing"
"github.com/google/uuid"
"github.com/prometheus/client_golang/prometheus"
promtestutil "github.com/prometheus/client_golang/prometheus/testutil"
"github.com/stretchr/testify/require"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/x/chatd/chatstate"
"github.com/coder/coder/v2/testutil"
)
type fakeAdmission struct {
mu sync.Mutex
refused map[uuid.UUID]bool
refuseFn func(database.Chat) bool
admitCalls int
admitted []uuid.UUID
limits AgentCapacityLimits
uncapped bool
}
func newFakeAdmission() *fakeAdmission {
return &fakeAdmission{
refused: make(map[uuid.UUID]bool),
limits: AgentCapacityLimits{Root: 1, Subagent: 1},
}
}
func (f *fakeAdmission) Limits() (AgentCapacityLimits, bool) {
return f.limits, !f.uncapped
}
func (f *fakeAdmission) refuse(chatID uuid.UUID) {
f.mu.Lock()
defer f.mu.Unlock()
f.refused[chatID] = true
}
func (f *fakeAdmission) allow(chatID uuid.UUID) {
f.mu.Lock()
defer f.mu.Unlock()
delete(f.refused, chatID)
}
func (f *fakeAdmission) Admit(_ context.Context, _ database.Store, chat database.Chat) (bool, error) {
f.mu.Lock()
defer f.mu.Unlock()
f.admitCalls++
if f.refused[chat.ID] || (f.refuseFn != nil && f.refuseFn(chat)) {
return false, nil
}
f.admitted = append(f.admitted, chat.ID)
return true, nil
}
func (f *fakeAdmission) admittedOrder() []uuid.UUID {
f.mu.Lock()
defer f.mu.Unlock()
return append([]uuid.UUID(nil), f.admitted...)
}
func (f *fakeAdmission) admitCallCount() int {
f.mu.Lock()
defer f.mu.Unlock()
return f.admitCalls
}
func TestWorker_AdmissionRefusalDoesNotAcquireChat(t *testing.T) {
t.Parallel()
f := newWorkerTestFixture(t)
recording := newRecordingPubsub(f.pubsub)
starter := newRecordingTaskStarter()
admission := newFakeAdmission()
opts := testOptions(t, f, starter)
opts.Pubsub = recording
opts.AgentCapacityLimiter = admission
chat := f.createRunningChat(t)
admission.refuse(chat.ID)
startWorker(t, opts)
require.Eventually(t, func() bool {
return admission.admitCallCount() > 0
}, testutil.WaitLong, testutil.IntervalFast)
starter.assertNoCall(t)
// The recorder wraps only worker pubsub, so an ownership hint here would
// prove a refusal can wake workers into an immediate retry loop.
require.Empty(t, recording.ownershipMessages(t))
}
func TestWorker_InterruptingSortsBeforeRunning(t *testing.T) {
t.Parallel()
f := newWorkerTestFixture(t)
ctx := testutil.Context(t, testutil.WaitLong)
running := []database.Chat{f.createRunningChat(t), f.createRunningChat(t)}
interrupting := f.createRunningChat(t)
interruptChat(t, f, interrupting.ID)
requiresAction := f.createRequiresActionChat(t)
_, err := f.sqlDB.ExecContext(ctx, `
UPDATE chats
SET updated_at = NOW() - INTERVAL '1 hour'
WHERE id IN ($1, $2)
`, running[0].ID, running[1].ID)
require.NoError(t, err)
rows, err := f.db.GetChatWorkerAcquisitionCandidates(ctx, database.GetChatWorkerAcquisitionCandidatesParams{
StaleSeconds: 30,
LimitCount: 2,
})
require.NoError(t, err)
require.Len(t, rows, 2)
require.Equal(t, interrupting.ID, rows[0].ID)
require.Equal(t, requiresAction.ID, rows[1].ID)
}
func TestWorker_AcquisitionCandidatesInterleavePools(t *testing.T) {
t.Parallel()
f := newWorkerTestFixture(t)
ctx := testutil.Context(t, testutil.WaitLong)
rootOlder := f.createRunningChat(t)
rootNewer := f.createRunningChat(t)
subOlder := f.createRunningSubagentChat(t, rootOlder.ID)
subNewer := f.createRunningSubagentChat(t, rootOlder.ID)
_, err := f.sqlDB.ExecContext(ctx, `
UPDATE chats
SET updated_at = CASE id
WHEN $1 THEN NOW() - INTERVAL '4 hours'
WHEN $2 THEN NOW() - INTERVAL '3 hours'
WHEN $3 THEN NOW() - INTERVAL '2 hours'
WHEN $4 THEN NOW() - INTERVAL '1 hour'
END
WHERE id IN ($1, $2, $3, $4)
`, rootOlder.ID, subOlder.ID, rootNewer.ID, subNewer.ID)
require.NoError(t, err)
rows, err := f.db.GetChatWorkerAcquisitionCandidates(ctx, database.GetChatWorkerAcquisitionCandidatesParams{
StaleSeconds: 30,
LimitCount: 4,
})
require.NoError(t, err)
require.Len(t, rows, 4)
require.Equal(t, []uuid.UUID{rootOlder.ID, subOlder.ID, rootNewer.ID, subNewer.ID}, []uuid.UUID{
rows[0].ID,
rows[1].ID,
rows[2].ID,
rows[3].ID,
})
}
func TestWorker_MessageBumpSendsChatToQueueBack(t *testing.T) {
t.Parallel()
f := newWorkerTestFixture(t)
ctx := testutil.Context(t, testutil.WaitLong)
older := f.createRunningChat(t)
newer := f.createRunningChat(t)
_, err := f.sqlDB.ExecContext(ctx, `
UPDATE chats
SET updated_at = CASE id
WHEN $1 THEN NOW() - INTERVAL '1 hour'
WHEN $2 THEN NOW() - INTERVAL '30 minutes'
END
WHERE id IN ($1, $2)
`, older.ID, newer.ID)
require.NoError(t, err)
machine := chatstate.NewChatMachine(f.db, f.pubsub, older.ID)
require.NoError(t, machine.Update(ctx, func(tx *chatstate.Tx, _ database.Store) error {
_, err := tx.SendMessage(chatstate.SendMessageInput{
Message: userTextMessage(t, "move me", f.user.ID, f.model.ID, f.apiKey.ID),
BusyBehavior: chatstate.BusyBehaviorQueue,
})
return err
}))
rows, err := f.db.GetChatWorkerAcquisitionCandidates(ctx, database.GetChatWorkerAcquisitionCandidatesParams{
StaleSeconds: 30,
LimitCount: 10,
})
require.NoError(t, err)
require.GreaterOrEqual(t, len(rows), 2)
require.Equal(t, newer.ID, rows[0].ID)
require.Equal(t, older.ID, rows[1].ID)
}
func newRootRefusingAdmission() *fakeAdmission {
admission := newFakeAdmission()
admission.refuseFn = func(chat database.Chat) bool {
return chat.Status == database.ChatStatusRunning && !chat.ParentChatID.Valid
}
return admission
}
func TestWorker_FullPoolDoesNotStarveOtherPool(t *testing.T) {
t.Parallel()
f := newWorkerTestFixture(t)
starter := newRecordingTaskStarter()
opts := testOptions(t, f, starter)
opts.AgentCapacityLimiter = newRootRefusingAdmission()
roots := make([]database.Chat, 0, 2*int(opts.AcquisitionBatchSize)+5)
for range cap(roots) {
roots = append(roots, f.createRunningChat(t))
}
sub := f.createRunningSubagentChat(t, roots[0].ID)
startWorker(t, opts)
call := starter.waitCall(t, taskKindGeneration, sub.ID)
require.Equal(t, sub.ID, call.input.ChatID)
}
func TestWorker_BatchSizeOneCannotHideAPool(t *testing.T) {
t.Parallel()
f := newWorkerTestFixture(t)
starter := newRecordingTaskStarter()
opts := testOptions(t, f, starter)
opts.AgentCapacityLimiter = newRootRefusingAdmission()
opts.AcquisitionBatchSize = 1
roots := []database.Chat{f.createRunningChat(t), f.createRunningChat(t)}
sub := f.createRunningSubagentChat(t, roots[0].ID)
startWorker(t, opts)
call := starter.waitCall(t, taskKindGeneration, sub.ID)
require.Equal(t, sub.ID, call.input.ChatID)
}
func TestWorker_FullPoolSkipsRefusalsAfterFirst(t *testing.T) {
t.Parallel()
f := newWorkerTestFixture(t)
starter := newRecordingTaskStarter()
opts := testOptions(t, f, starter)
admission := newRootRefusingAdmission()
opts.AgentCapacityLimiter = admission
opts.AcquisitionBatchSize = 2
for range 5 {
f.createRunningChat(t)
}
startWorker(t, opts)
require.Eventually(t, func() bool {
return admission.admitCallCount() == 1
}, testutil.WaitLong, testutil.IntervalFast)
starter.assertNoCall(t)
require.Equal(t, 1, admission.admitCallCount(),
"a full pool must be skipped after one refusal, not re-refused per chat")
}
func TestWorker_AdmissionAdmitsInUpdatedAtOrder(t *testing.T) {
t.Parallel()
f := newWorkerTestFixture(t)
starter := newRecordingTaskStarter()
opts := testOptions(t, f, starter)
admission := newFakeAdmission()
opts.AgentCapacityLimiter = admission
older := f.createRunningChat(t)
newer := f.createRunningChat(t)
// Back-to-back inserts can collide at timestamp resolution, which would
// leave FIFO order to the random UUID tiebreak.
ctx := testutil.Context(t, testutil.WaitLong)
_, err := f.sqlDB.ExecContext(ctx,
"UPDATE chats SET updated_at = NOW() - INTERVAL '1 hour' WHERE id = $1", older.ID)
require.NoError(t, err)
admission.refuse(older.ID)
admission.refuse(newer.ID)
worker := startWorker(t, opts)
require.Eventually(t, func() bool {
return admission.admitCallCount() == 1
}, testutil.WaitLong, testutil.IntervalFast)
admission.allow(older.ID)
admission.allow(newer.ID)
worker.Wake()
// Runner goroutines race task starts, so wait for both without
// ordering and assert the worker's serial admission order instead.
starter.waitCall(t, taskKindGeneration, uuid.Nil)
starter.waitCall(t, taskKindGeneration, uuid.Nil)
require.Equal(t, []uuid.UUID{older.ID, newer.ID}, admission.admittedOrder(),
"the longer-waiting chat must admit first")
}
func TestWorker_InterruptClaimsCapacityQueuedChat(t *testing.T) {
t.Parallel()
f := newWorkerTestFixture(t)
starter := newRecordingTaskStarter()
opts := testOptions(t, f, starter)
admission := newFakeAdmission()
admission.refuseFn = func(chat database.Chat) bool {
return chat.Status == database.ChatStatusRunning
}
opts.AgentCapacityLimiter = admission
chat := f.createRunningChat(t)
worker := startWorker(t, opts)
require.Eventually(t, func() bool {
return admission.admitCallCount() > 0
}, testutil.WaitLong, testutil.IntervalFast)
interruptChat(t, f, chat.ID)
worker.Wake()
call := starter.waitCall(t, taskKindInterrupt, chat.ID)
require.Equal(t, chat.ID, call.input.ChatID)
}
func TestWorker_CapacityMetricsUseFreshOwnership(t *testing.T) {
t.Parallel()
f := newWorkerTestFixture(t)
metrics := newCapacityMetrics(prometheus.NewRegistry())
opts := testOptions(t, f, newRecordingTaskStarter())
opts.CapacityMetrics = metrics
opts.AgentCapacityLimiter = newFakeAdmission()
ctx := testutil.Context(t, testutil.WaitLong)
occupied := f.createRunningChat(t)
acquireChat(t, f, occupied.ID, uuid.New(), uuid.New())
f.createRunningChat(t)
worker, err := newChatWorker(newUnstartedServer(t, f.pubsub, f.db), opts)
require.NoError(t, err)
worker.refreshCapacityMetrics(ctx)
require.Equal(t, float64(1), promtestutil.ToFloat64(metrics.active.WithLabelValues("root")))
require.Equal(t, float64(1), promtestutil.ToFloat64(metrics.queued.WithLabelValues("root")))
forceExecutionState(t, f, occupied.ID, database.ChatStatusWaiting, true)
worker.refreshCapacityMetrics(ctx)
require.Equal(t, float64(1), promtestutil.ToFloat64(metrics.active.WithLabelValues("root")))
require.Equal(t, float64(1), promtestutil.ToFloat64(metrics.queued.WithLabelValues("root")))
}
func TestGetChatQueuedForCapacity(t *testing.T) {
t.Parallel()
queued := func(t *testing.T, f *workerTestFixture, chatID uuid.UUID, rootCap, subagentCap int64) bool {
t.Helper()
ctx := testutil.Context(t, testutil.WaitLong)
got, err := f.db.GetChatQueuedForCapacity(ctx, database.GetChatQueuedForCapacityParams{
ChatID: chatID,
StaleSeconds: 30,
RootCapacity: rootCap,
SubagentCapacity: subagentCap,
})
require.NoError(t, err)
return got
}
t.Run("PoolNotFull", func(t *testing.T) {
t.Parallel()
f := newWorkerTestFixture(t)
chat := f.createRunningChat(t)
require.False(t, queued(t, f, chat.ID, 1, 1))
})
t.Run("PoolFull", func(t *testing.T) {
t.Parallel()
f := newWorkerTestFixture(t)
occupied := f.createRunningChat(t)
acquireChat(t, f, occupied.ID, uuid.New(), uuid.New())
chat := f.createRunningChat(t)
require.True(t, queued(t, f, chat.ID, 1, 1))
})
t.Run("IncompleteOwnershipIsQueued", func(t *testing.T) {
t.Parallel()
f := newWorkerTestFixture(t)
occupied := f.createRunningChat(t)
acquireChat(t, f, occupied.ID, uuid.New(), uuid.New())
chat := f.createRunningChat(t)
acquireChat(t, f, chat.ID, uuid.New(), uuid.New())
_, err := f.sqlDB.ExecContext(testutil.Context(t, testutil.WaitLong), `UPDATE chats SET worker_id = NULL WHERE id = $1`, chat.ID)
require.NoError(t, err)
require.True(t, queued(t, f, chat.ID, 1, 1))
})
t.Run("OwnedChatIsNotQueued", func(t *testing.T) {
t.Parallel()
f := newWorkerTestFixture(t)
occupied := f.createRunningChat(t)
acquireChat(t, f, occupied.ID, uuid.New(), uuid.New())
require.False(t, queued(t, f, occupied.ID, 1, 1))
})
t.Run("NonRunningChatIsNotQueued", func(t *testing.T) {
t.Parallel()
f := newWorkerTestFixture(t)
occupied := f.createRunningChat(t)
acquireChat(t, f, occupied.ID, uuid.New(), uuid.New())
requiresAction := f.createRequiresActionChat(t)
require.False(t, queued(t, f, requiresAction.ID, 1, 1))
})
t.Run("PoolsAreIndependent", func(t *testing.T) {
t.Parallel()
f := newWorkerTestFixture(t)
occupied := f.createRunningChat(t)
acquireChat(t, f, occupied.ID, uuid.New(), uuid.New())
sub := f.createRunningSubagentChat(t, occupied.ID)
require.False(t, queued(t, f, sub.ID, 1, 1),
"a full root pool must not mark subagents queued")
})
}
func TestServer_ChatQueuedForCapacity(t *testing.T) {
t.Parallel()
f := newWorkerTestFixture(t)
ctx := testutil.Context(t, testutil.WaitLong)
occupied := f.createRunningChat(t)
acquireChat(t, f, occupied.ID, uuid.New(), uuid.New())
for range 4 {
chat := f.createRunningChat(t)
acquireChat(t, f, chat.ID, uuid.New(), uuid.New())
}
waiting := f.createRunningChat(t)
server := newUnstartedServer(t, f.pubsub, f.db)
queued, err := server.ChatQueuedForCapacity(ctx, waiting)
require.NoError(t, err)
require.True(t, queued, "AGPL deployments must enforce the default root capacity")
uncapped := newFakeAdmission()
uncapped.uncapped = true
server.agentCapacityLimiter = uncapped
queued, err = server.ChatQueuedForCapacity(ctx, waiting)
require.NoError(t, err)
require.False(t, queued, "uncapped deployments must never report queued")
server.agentCapacityLimiter = newFakeAdmission()
queued, err = server.ChatQueuedForCapacity(ctx, waiting)
require.NoError(t, err)
require.True(t, queued)
queued, err = server.ChatQueuedForCapacity(ctx, occupied)
require.NoError(t, err)
require.False(t, queued, "owned chats are active, not queued")
}
func TestChatCapacityCountsByPool(t *testing.T) {
t.Parallel()
f := newWorkerTestFixture(t)
ctx := testutil.Context(t, testutil.WaitLong)
owned := f.createRunningChat(t)
acquireChat(t, f, owned.ID, uuid.New(), uuid.New())
f.createRunningChat(t)
f.createRunningSubagentChat(t, owned.ID)
incomplete := f.createRunningChat(t)
acquireChat(t, f, incomplete.ID, uuid.New(), uuid.New())
_, err := f.sqlDB.ExecContext(ctx, `UPDATE chats SET worker_id = NULL WHERE id = $1`, incomplete.ID)
require.NoError(t, err)
active, err := f.db.CountChatCapacityActiveByPool(ctx, database.CountChatCapacityActiveByPoolParams{StaleSeconds: 30})
require.NoError(t, err)
require.EqualValues(t, 1, active.ActiveRootCount)
require.EqualValues(t, 0, active.ActiveSubagentCount)
queued, err := f.db.CountChatCapacityQueuedByPool(ctx, 30)
require.NoError(t, err)
require.EqualValues(t, 2, queued.QueuedRootCount)
require.EqualValues(t, 1, queued.QueuedSubagentCount)
active, err = f.db.CountChatCapacityActiveByPool(ctx, database.CountChatCapacityActiveByPoolParams{
ExcludeChatID: owned.ID,
StaleSeconds: 30,
})
require.NoError(t, err)
require.EqualValues(t, 0, active.ActiveRootCount, "the excluded chat must not count as active")
}
@@ -0,0 +1,266 @@
package chatd
import (
"errors"
"sync"
"sync/atomic"
"testing"
"github.com/google/uuid"
"github.com/stretchr/testify/require"
"golang.org/x/xerrors"
"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/pubsub"
"github.com/coder/coder/v2/coderd/x/chatd/chatstate"
"github.com/coder/coder/v2/testutil"
)
type admissionFixture struct {
db database.Store
ps pubsub.Pubsub
owner database.User
org database.Organization
modelConfig database.ChatModelConfig
}
func newAdmissionFixture(t *testing.T) admissionFixture {
t.Helper()
db, ps := dbtestutil.NewDB(t)
owner := dbgen.User(t, db, database.User{})
org := dbgen.Organization(t, db, database.Organization{})
modelConfig := dbgen.ChatModelConfig(t, db, database.ChatModelConfig{})
return admissionFixture{db: db, ps: ps, owner: owner, org: org, modelConfig: modelConfig}
}
func (f admissionFixture) chat(t *testing.T, seed database.Chat) database.Chat {
t.Helper()
seed.OwnerID = f.owner.ID
seed.OrganizationID = f.org.ID
seed.LastModelConfigID = f.modelConfig.ID
if seed.Status == "" {
seed.Status = database.ChatStatusRunning
}
return dbgen.Chat(t, f.db, seed)
}
func (f admissionFixture) occupy(t *testing.T, chatID uuid.UUID) {
t.Helper()
ctx := testutil.Context(t, testutil.WaitShort)
machine := chatstate.NewChatMachine(f.db, f.ps, chatID)
require.NoError(t, machine.Update(ctx, func(tx *chatstate.Tx, _ database.Store) error {
_, err := tx.Acquire(chatstate.AcquireInput{WorkerID: uuid.New(), RunnerID: uuid.New()})
return err
}))
}
func (f admissionFixture) occupiedRoot(t *testing.T) database.Chat {
t.Helper()
chat := f.chat(t, database.Chat{})
f.occupy(t, chat.ID)
return chat
}
func (f admissionFixture) occupiedSubagent(t *testing.T, root database.Chat) database.Chat {
t.Helper()
chat := f.chat(t, database.Chat{
ParentChatID: uuid.NullUUID{UUID: root.ID, Valid: true},
RootChatID: uuid.NullUUID{UUID: root.ID, Valid: true},
})
f.occupy(t, chat.ID)
return chat
}
func testAdmission() *agentCapacityLimiter {
a := newAgentCapacityLimiter(nil, 30)
a.rootCapacity = 2
a.subagentCapacity = 2
return a
}
func TestAdmission_RootPoolCap(t *testing.T) {
t.Parallel()
f := newAdmissionFixture(t)
ctx := testutil.Context(t, testutil.WaitLong)
a := testAdmission()
root := f.occupiedRoot(t)
f.occupiedRoot(t)
admitted, err := a.Admit(ctx, f.db, f.chat(t, database.Chat{}))
require.NoError(t, err)
require.False(t, admitted, "third root must be refused at capacity 2")
subagent := f.chat(t, database.Chat{
ParentChatID: uuid.NullUUID{UUID: root.ID, Valid: true},
RootChatID: uuid.NullUUID{UUID: root.ID, Valid: true},
})
admitted, err = a.Admit(ctx, f.db, subagent)
require.NoError(t, err)
require.True(t, admitted)
}
func TestAdmission_SubagentPoolCap(t *testing.T) {
t.Parallel()
f := newAdmissionFixture(t)
ctx := testutil.Context(t, testutil.WaitLong)
a := testAdmission()
root := f.occupiedRoot(t)
f.occupiedSubagent(t, root)
f.occupiedSubagent(t, root)
subagent := f.chat(t, database.Chat{
ParentChatID: uuid.NullUUID{UUID: root.ID, Valid: true},
RootChatID: uuid.NullUUID{UUID: root.ID, Valid: true},
})
admitted, err := a.Admit(ctx, f.db, subagent)
require.NoError(t, err)
require.False(t, admitted, "third subagent must be refused at capacity 2")
admitted, err = a.Admit(ctx, f.db, f.chat(t, database.Chat{}))
require.NoError(t, err)
require.True(t, admitted, "a full subagent pool must not refuse roots")
}
func TestAdmission_InterruptingBypassesCap(t *testing.T) {
t.Parallel()
f := newAdmissionFixture(t)
ctx := testutil.Context(t, testutil.WaitLong)
a := testAdmission()
f.occupiedRoot(t)
f.occupiedRoot(t)
interrupting := f.chat(t, database.Chat{Status: database.ChatStatusInterrupting})
admitted, err := a.Admit(ctx, f.db, interrupting)
require.NoError(t, err)
require.True(t, admitted, "interrupting chats must always be acquirable")
}
func TestAdmission_RequiresActionBypassesCap(t *testing.T) {
t.Parallel()
f := newAdmissionFixture(t)
ctx := testutil.Context(t, testutil.WaitLong)
a := testAdmission()
f.occupiedRoot(t)
f.occupiedRoot(t)
requiresAction := f.chat(t, database.Chat{Status: database.ChatStatusRequiresAction})
admitted, err := a.Admit(ctx, f.db, requiresAction)
require.NoError(t, err)
require.True(t, admitted, "requires_action chats hold no slot and need their runner")
}
func TestAdmission_TakeoverOfCountedChatIsCapacityNeutral(t *testing.T) {
t.Parallel()
f := newAdmissionFixture(t)
ctx := testutil.Context(t, testutil.WaitLong)
a := testAdmission()
counted := f.occupiedRoot(t)
f.occupiedRoot(t)
chat, err := f.db.GetChatByID(ctx, counted.ID)
require.NoError(t, err)
admitted, err := a.Admit(ctx, f.db, chat)
require.NoError(t, err)
require.True(t, admitted, "an already-counted chat must re-admit for takeover at full capacity")
}
// The single transaction verifies that the admission lock covers the
// ownership write.
func TestAdmission_ConcurrentAdmitNeverOverAdmits(t *testing.T) {
t.Parallel()
f := newAdmissionFixture(t)
ctx := testutil.Context(t, testutil.WaitLong)
a := testAdmission()
const attempts = 8
chats := make([]database.Chat, attempts)
for i := range chats {
chats[i] = f.chat(t, database.Chat{})
}
errRefused := xerrors.New("refused")
var (
admitted atomic.Int64
unexpected atomic.Int64
wg sync.WaitGroup
)
for _, chat := range chats {
wg.Go(func() {
machine := chatstate.NewChatMachine(f.db, f.ps, chat.ID)
err := machine.Update(ctx, func(tx *chatstate.Tx, store database.Store) error {
ok, err := a.Admit(ctx, store, chat)
if err != nil {
return err
}
if !ok {
return errRefused
}
_, err = tx.Acquire(chatstate.AcquireInput{WorkerID: uuid.New(), RunnerID: uuid.New()})
return err
})
switch {
case err == nil:
admitted.Add(1)
case errors.Is(err, errRefused):
default:
unexpected.Add(1)
}
})
}
wg.Wait()
require.EqualValues(t, 0, unexpected.Load(), "admission attempts must not error")
require.EqualValues(t, 2, admitted.Load(), "exactly rootCapacity chats must admit")
}
func TestAdmission_StaleHeartbeatsFreeSlots(t *testing.T) {
t.Parallel()
f := newAdmissionFixture(t)
ctx := testutil.Context(t, testutil.WaitLong)
f.occupiedRoot(t)
f.occupiedRoot(t)
// A zero staleness window makes every heartbeat stale.
a := newAgentCapacityLimiter(nil, 0)
a.rootCapacity = 2
a.subagentCapacity = 2
admitted, err := a.Admit(ctx, f.db, f.chat(t, database.Chat{}))
require.NoError(t, err)
require.True(t, admitted)
}
type staticAgentCapacityUnlock bool
func (u staticAgentCapacityUnlock) Unlocked() bool {
return bool(u)
}
func TestAdmission_UnlockBypassesCaps(t *testing.T) {
t.Parallel()
a := newAgentCapacityLimiter(staticAgentCapacityUnlock(true), 30)
admitted, err := a.Admit(t.Context(), nil, database.Chat{Status: database.ChatStatusRunning})
require.NoError(t, err)
require.True(t, admitted)
_, capped := a.Limits()
require.False(t, capped)
}
func TestAdmission_LimitsReportsCaps(t *testing.T) {
t.Parallel()
a := newAgentCapacityLimiter(nil, 30)
limits, capped := a.Limits()
require.True(t, capped)
require.EqualValues(t, defaultMaxConcurrentRootAgents, limits.Root)
require.EqualValues(t, defaultMaxConcurrentSubagents, limits.Subagent)
}
+84
View File
@@ -0,0 +1,84 @@
package chatd
import (
"context"
"github.com/google/uuid"
"github.com/prometheus/client_golang/prometheus"
"github.com/coder/coder/v2/coderd/database"
)
type capacityMetrics struct {
active *prometheus.GaugeVec
queued *prometheus.GaugeVec
}
func newCapacityMetrics(registerer prometheus.Registerer) *capacityMetrics {
m := &capacityMetrics{
active: prometheus.NewGaugeVec(prometheus.GaugeOpts{
Namespace: "coderd",
Subsystem: "chatd",
Name: "agents_active",
Help: "Deployment-wide number of chats holding a concurrent-agent capacity slot. Every replica reports the same database-derived value; aggregate with max, not sum.",
}, []string{"pool"}),
queued: prometheus.NewGaugeVec(prometheus.GaugeOpts{
Namespace: "coderd",
Subsystem: "chatd",
Name: "agents_queued_for_capacity",
Help: "Deployment-wide number of chats waiting for a concurrent-agent capacity slot. Every replica reports the same database-derived value; aggregate with max, not sum.",
}, []string{"pool"}),
}
registerer.MustRegister(m.active, m.queued)
return m
}
func (w *chatWorker) capacityMetricsLoop(ctx context.Context) {
ticker := w.opts.Clock.NewTicker(w.opts.CapacityMetricsInterval, "chatworker", "capacity-metrics")
defer ticker.Stop()
for {
select {
case <-ticker.C:
case <-ctx.Done():
return
}
w.refreshCapacityMetrics(ctx)
}
}
func (w *chatWorker) refreshCapacityMetrics(ctx context.Context) {
active, err := w.opts.Store.CountChatCapacityActiveByPool(ctx, database.CountChatCapacityActiveByPoolParams{
ExcludeChatID: uuid.Nil,
StaleSeconds: w.opts.HeartbeatStaleSeconds,
})
if err != nil {
if ctx.Err() == nil {
w.opts.Logger.Warn(ctx, "chatworker count active capacity chats failed", slogError(err))
}
return
}
limits, capped := w.opts.AgentCapacityLimiter.Limits()
var queuedRoot, queuedSubagent int64
if capped && (active.ActiveRootCount >= limits.Root || active.ActiveSubagentCount >= limits.Subagent) {
queued, err := w.opts.Store.CountChatCapacityQueuedByPool(ctx, w.opts.HeartbeatStaleSeconds)
if err != nil {
if ctx.Err() == nil {
w.opts.Logger.Warn(ctx, "chatworker count queued capacity chats failed", slogError(err))
}
return
}
if active.ActiveRootCount >= limits.Root {
queuedRoot = queued.QueuedRootCount
}
if active.ActiveSubagentCount >= limits.Subagent {
queuedSubagent = queued.QueuedSubagentCount
}
}
metrics := w.opts.CapacityMetrics
metrics.active.WithLabelValues("root").Set(float64(active.ActiveRootCount))
metrics.active.WithLabelValues("subagent").Set(float64(active.ActiveSubagentCount))
metrics.queued.WithLabelValues("root").Set(float64(queuedRoot))
metrics.queued.WithLabelValues("subagent").Set(float64(queuedSubagent))
}
+42 -10
View File
@@ -189,13 +189,14 @@ type Server struct {
configCacheUnsubscribe func()
providerCacheUnsubscribe func()
usageTracker *workspacestats.UsageTracker
clock quartz.Clock
metrics *chatloop.Metrics
chatWorker *chatWorker
messagePartBuffer *messagepartbuffer.Buffer
streamSyncPoller *streamSyncPoller
recordingSem chan struct{}
usageTracker *workspacestats.UsageTracker
clock quartz.Clock
metrics *chatloop.Metrics
chatWorker *chatWorker
messagePartBuffer *messagepartbuffer.Buffer
streamSyncPoller *streamSyncPoller
recordingSem chan struct{}
agentCapacityLimiter AgentCapacityLimiter
aibridgeTransportFactory *atomic.Pointer[aibridge.TransportFactory]
experiments codersdk.Experiments
@@ -3050,6 +3051,8 @@ type Config struct {
PrometheusRegistry prometheus.Registerer
AgentCapacityUnlock AgentCapacityUnlock
// OIDCTokenSource resolves the calling user's OIDC access
// token for MCP servers configured with auth_type=user_oidc.
// May be nil if the deployment has no OIDC provider; servers
@@ -3181,6 +3184,15 @@ func New(ps pubsub.Pubsub, cfg Config) *Server {
p.streamPartsDialer = streamPartsDialerForServer(workerID, localStreamPartsDialer, cfg.StreamPartsDialer)
p.streamSyncPoller = newStreamSyncPoller(ctx, cfg.Database, clk, cfg.Logger.Named("chatstream"))
p.streamSyncPoller.Start()
agentCapacityLimiter := newAgentCapacityLimiter(
cfg.AgentCapacityUnlock,
int32(inFlightChatStaleAfter.Seconds()),
)
var agentCapacityMetrics *capacityMetrics
if cfg.PrometheusRegistry != nil {
agentCapacityMetrics = newCapacityMetrics(cfg.PrometheusRegistry)
}
p.agentCapacityLimiter = agentCapacityLimiter
chatWorker, err := newChatWorker(p, chatWorkerOptions{
WorkerID: workerID,
Store: cfg.Database,
@@ -3188,6 +3200,8 @@ func New(ps pubsub.Pubsub, cfg Config) *Server {
Logger: cfg.Logger.Named("chatworker"),
Clock: clk,
MessagePartBuffer: p.messagePartBuffer,
AgentCapacityLimiter: agentCapacityLimiter,
CapacityMetrics: agentCapacityMetrics,
AcquisitionInterval: pendingChatAcquireInterval,
AcquisitionBatchSize: maxChatsPerAcquire,
HeartbeatInterval: chatHeartbeatInterval,
@@ -3302,9 +3316,6 @@ func chatWatchEventSDKChat(chat database.Chat, diffStatus *codersdk.ChatDiffStat
// publishChatPubsubEvent broadcasts a chat lifecycle event via PostgreSQL
// pubsub so that all replicas can push updates to watching clients.
func (p *Server) publishChatPubsubEvent(chat database.Chat, kind codersdk.ChatWatchEventKind, diffStatus *codersdk.ChatDiffStatus) {
if p.pubsub == nil {
return
}
event := codersdk.ChatWatchEvent{
Kind: kind,
Chat: chatWatchEventSDKChat(chat, diffStatus),
@@ -3326,6 +3337,27 @@ func (p *Server) publishChatPubsubEvent(chat database.Chat, kind codersdk.ChatWa
}
}
// ChatQueuedForCapacity reports whether the chat is waiting for a
// concurrent-agent capacity slot. Uncapped deployments always return false.
func (p *Server) ChatQueuedForCapacity(ctx context.Context, chat database.Chat) (bool, error) {
limits, capped := p.agentCapacityLimiter.Limits()
if !capped {
return false, nil
}
if chat.Archived || chat.Status != database.ChatStatusRunning {
return false, nil
}
// The pool count spans other users' chats, which the requester cannot
// read directly.
//nolint:gocritic // Capacity accounting is chatd-internal state.
return p.db.GetChatQueuedForCapacity(dbauthz.AsChatd(ctx), database.GetChatQueuedForCapacityParams{
ChatID: chat.ID,
StaleSeconds: int32(p.inFlightChatStaleAfter.Seconds()),
RootCapacity: limits.Root,
SubagentCapacity: limits.Subagent,
})
}
// PublishDiffStatusChange broadcasts a diff_status_change event for
// the given chat so that watching clients know to re-fetch the diff
// status. This is called from the HTTP layer after the diff status
+4 -4
View File
@@ -187,7 +187,7 @@ func TestStoreSubagentReportSummary(t *testing.T) {
ctrl := gomock.NewController(t)
db := dbmock.NewMockStore(ctrl)
server := &Server{db: db}
server := &Server{db: db, pubsub: dbpubsub.NewInMemory()}
chat := database.Chat{
ID: uuid.New(),
OwnerID: uuid.New(),
@@ -216,7 +216,7 @@ func TestStoreSubagentReportSummary(t *testing.T) {
ctrl := gomock.NewController(t)
db := dbmock.NewMockStore(ctrl)
server := &Server{db: db}
server := &Server{db: db, pubsub: dbpubsub.NewInMemory()}
chat := database.Chat{
ID: uuid.New(),
OwnerID: uuid.New(),
@@ -237,7 +237,7 @@ func TestStoreSubagentReportSummary(t *testing.T) {
ctrl := gomock.NewController(t)
db := dbmock.NewMockStore(ctrl)
server := &Server{db: db}
server := &Server{db: db, pubsub: dbpubsub.NewInMemory()}
chat := database.Chat{
ID: uuid.New(),
OwnerID: uuid.New(),
@@ -268,7 +268,7 @@ func TestStoreSubagentReportSummary(t *testing.T) {
ctrl := gomock.NewController(t)
db := dbmock.NewMockStore(ctrl)
server := &Server{db: db}
server := &Server{db: db, pubsub: dbpubsub.NewInMemory()}
chat := database.Chat{
ID: uuid.New(),
OwnerID: uuid.New(),
+19
View File
@@ -148,6 +148,25 @@ func (f *workerTestFixture) createRunningChat(t *testing.T) database.Chat {
return res.Chat
}
func (f *workerTestFixture) createRunningSubagentChat(t *testing.T, parentID uuid.UUID) database.Chat {
t.Helper()
ctx := testutil.Context(t, testutil.WaitShort)
res, err := chatstate.CreateChat(ctx, f.db, f.pubsub, chatstate.CreateChatInput{
OrganizationID: f.org.ID,
OwnerID: f.user.ID,
LastModelConfigID: f.model.ID,
Title: "subagent",
ClientType: database.ChatClientTypeApi,
ParentChatID: uuid.NullUUID{UUID: parentID, Valid: true},
RootChatID: uuid.NullUUID{UUID: parentID, Valid: true},
InitialMessages: []chatstate.Message{
userTextMessage(t, "hello", f.user.ID, f.model.ID, f.apiKey.ID),
},
})
require.NoError(t, err)
return res.Chat
}
func (f *workerTestFixture) createRequiresActionChat(t *testing.T) database.Chat {
t.Helper()
ctx := testutil.Context(t, testutil.WaitShort)
+17 -6
View File
@@ -21,12 +21,13 @@ import (
)
const (
defaultAcquisitionInterval = 30 * time.Second
defaultAcquisitionBatchSize = int32(10)
defaultRunnerSyncInterval = 15 * time.Second
defaultHeartbeatInterval = 9 * time.Second
defaultHeartbeatCleanupEvery = 30 * time.Second
defaultHeartbeatStaleSeconds = int32(30)
defaultAcquisitionInterval = 30 * time.Second
defaultAcquisitionBatchSize = int32(10)
defaultCapacityMetricsInterval = 30 * time.Second
defaultRunnerSyncInterval = 15 * time.Second
defaultHeartbeatInterval = 9 * time.Second
defaultHeartbeatCleanupEvery = 30 * time.Second
defaultHeartbeatStaleSeconds = int32(30)
// The archive cutoff is based on UTC start-of-day and only moves
// once per day, so hourly runs are more than enough to keep up
// while still catching chats that cross the threshold shortly
@@ -192,7 +193,11 @@ type chatWorkerOptions struct {
Auditor *atomic.Pointer[audit.Auditor]
AutoArchiveRecords prometheus.Counter
AgentCapacityLimiter AgentCapacityLimiter
CapacityMetrics *capacityMetrics
AcquisitionInterval time.Duration
CapacityMetricsInterval time.Duration
AcquisitionBatchSize int32
ArchiveInterval time.Duration
ArchiveBatchSize int32
@@ -226,6 +231,9 @@ func (o chatWorkerOptions) withDefaults() (chatWorkerOptions, error) {
if o.AcquisitionInterval <= 0 {
o.AcquisitionInterval = defaultAcquisitionInterval
}
if o.CapacityMetricsInterval <= 0 {
o.CapacityMetricsInterval = defaultCapacityMetricsInterval
}
if o.AcquisitionBatchSize <= 0 {
o.AcquisitionBatchSize = defaultAcquisitionBatchSize
}
@@ -250,6 +258,9 @@ func (o chatWorkerOptions) withDefaults() (chatWorkerOptions, error) {
if o.HeartbeatStaleSeconds <= 0 {
o.HeartbeatStaleSeconds = defaultHeartbeatStaleSeconds
}
if o.AgentCapacityLimiter == nil {
o.AgentCapacityLimiter = newAgentCapacityLimiter(nil, o.HeartbeatStaleSeconds)
}
if o.StateChannelSize <= 0 {
o.StateChannelSize = defaultStateChannelSize
}
+2 -1
View File
@@ -24,6 +24,7 @@ import (
"github.com/coder/coder/v2/coderd/database/dbgen"
"github.com/coder/coder/v2/coderd/database/dbmock"
"github.com/coder/coder/v2/coderd/database/dbtestutil"
dbpubsub "github.com/coder/coder/v2/coderd/database/pubsub"
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
"github.com/coder/coder/v2/coderd/x/chatd/chatprovider"
"github.com/coder/coder/v2/coderd/x/chatd/chattest"
@@ -580,7 +581,7 @@ func TestMaybeGenerateChatTitlePreservesUpdatedAt(t *testing.T) {
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true})
generated := &generatedChatTitle{}
server := &Server{db: db}
server := &Server{db: db, pubsub: dbpubsub.NewInMemory()}
server.maybeGenerateChatTitle(
ctx,
chat,
+1
View File
@@ -293,6 +293,7 @@ func TestCreateChildSubagentChatDispatchesUserPromptSubmit(t *testing.T) {
t.Cleanup(consumer.Close)
server := &Server{
db: db,
pubsub: pubsub.NewInMemory(),
logger: slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}),
hooks: chathooks.NewTrigger(dispatch.New(
slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}),
@@ -24,6 +24,7 @@ import (
"github.com/coder/coder/v2/coderd/aibridge"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/database/dbmock"
dbpubsub "github.com/coder/coder/v2/coderd/database/pubsub"
"github.com/coder/coder/v2/coderd/util/ptr"
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
"github.com/coder/coder/v2/coderd/x/chatd/chatprovider"
@@ -709,6 +710,7 @@ func titleOverrideTestServer(db database.Store, logger slog.Logger) *Server {
})}
return &Server{
db: db,
pubsub: dbpubsub.NewInMemory(),
logger: logger,
configCache: newChatConfigCache(context.Background(), db, quartz.NewReal()),
aibridgeTransportFactory: aibridgeTestFactoryPointer(factory),
+69 -36
View File
@@ -97,6 +97,11 @@ func (w *chatWorker) Start(ctx context.Context) error {
w.wg.Go(func() {
w.archiveLoop(workerCtx)
})
if w.opts.CapacityMetrics != nil {
w.wg.Go(func() {
w.capacityMetricsLoop(workerCtx)
})
}
wake(wakeCh)
return nil
}
@@ -187,49 +192,65 @@ func (w *chatWorker) acquisitionLoop(
}
func (w *chatWorker) acquireOnce(ctx context.Context, workerID uuid.UUID, manager *runnerManager) {
attempted := make(map[uuid.UUID]struct{})
for {
rows, err := w.opts.Store.GetChatWorkerAcquisitionCandidates(ctx, database.GetChatWorkerAcquisitionCandidatesParams{
StaleSeconds: w.opts.HeartbeatStaleSeconds,
LimitCount: w.opts.AcquisitionBatchSize,
})
// Fetch twice the budget so one full pool cannot hide candidates in the other.
rows, err := w.opts.Store.GetChatWorkerAcquisitionCandidates(ctx, database.GetChatWorkerAcquisitionCandidatesParams{
StaleSeconds: w.opts.HeartbeatStaleSeconds,
LimitCount: w.opts.AcquisitionBatchSize * 2,
})
if err != nil {
if ctx.Err() == nil {
w.opts.Logger.Warn(ctx, "chatworker acquisition query failed", slogError(err))
}
return
}
acquired := int32(0)
rootPoolRefused := false
subagentPoolRefused := false
for _, row := range rows {
if acquired >= w.opts.AcquisitionBatchSize {
return
}
// Interrupting and requires-action chats bypass capacity so their runners
// can finish work or enforce the action deadline.
isSubagent := row.ParentChatID.Valid
if row.Status == database.ChatStatusRunning &&
((isSubagent && subagentPoolRefused) || (!isSubagent && rootPoolRefused)) {
continue
}
candidateAcquired, err := w.acquireCandidateSafely(ctx, workerID, manager, row.ID)
if errors.Is(err, errCapacityRefused) {
if isSubagent {
subagentPoolRefused = true
} else {
rootPoolRefused = true
}
continue
}
if err != nil {
if ctx.Err() == nil {
w.opts.Logger.Warn(ctx, "chatworker acquisition query failed", slogError(err))
if ctx.Err() != nil {
return
}
return
w.opts.Logger.Warn(ctx, "chatworker acquisition candidate failed", slogError(err))
continue
}
if len(rows) == 0 {
return
}
newRows := 0
for _, row := range rows {
if _, ok := attempted[row.ID]; ok {
continue
}
attempted[row.ID] = struct{}{}
newRows++
if err := w.acquireCandidateSafely(ctx, workerID, manager, row.ID); err != nil {
if ctx.Err() != nil {
return
}
w.opts.Logger.Warn(ctx, "chatworker acquisition candidate failed", slogError(err))
}
}
if len(rows) < int(w.opts.AcquisitionBatchSize) || newRows == 0 {
return
if candidateAcquired {
acquired++
}
}
}
var errSkipAcquire = xerrors.New("skip acquire")
var (
errSkipAcquire = xerrors.New("skip acquire")
errCapacityRefused = xerrors.New("capacity refused")
)
func (w *chatWorker) acquireCandidateSafely(
ctx context.Context,
workerID uuid.UUID,
manager *runnerManager,
chatID uuid.UUID,
) (err error) {
) (acquired bool, err error) {
defer func() {
if recovered := recover(); recovered != nil {
err = xerrors.Errorf("chatworker acquisition panic: %v", recovered)
@@ -243,7 +264,7 @@ func (w *chatWorker) acquireCandidate(
workerID uuid.UUID,
manager *runnerManager,
chatID uuid.UUID,
) error {
) (bool, error) {
runnerID := uuid.New()
machine := chatstate.NewChatMachine(w.opts.Store, w.opts.Pubsub, chatID)
err := machine.Update(ctx, func(tx *chatstate.Tx, store database.Store) error {
@@ -274,22 +295,34 @@ func (w *chatWorker) acquireCandidate(
return errSkipAcquire
}
}
admitted, err := w.opts.AgentCapacityLimiter.Admit(ctx, store, chat)
if err != nil {
return xerrors.Errorf("agent admission: %w", err)
}
if !admitted {
// Roll back to suppress the ownership hint, which would wake every
// worker into an immediate retry of this unowned chat.
return errCapacityRefused
}
_, err = tx.Acquire(chatstate.AcquireInput{WorkerID: workerID, RunnerID: runnerID})
return err
})
if errors.Is(err, errCapacityRefused) {
return false, errCapacityRefused
}
if errors.Is(err, errSkipAcquire) || errors.Is(err, chatstate.ErrChatNotFound) {
return nil
return false, nil
}
if err != nil {
return err
return false, err
}
if err := manager.Spawn(ctx, spawnRunnerRequest{ChatID: chatID, WorkerID: workerID, RunnerID: runnerID}); err != nil {
if errAbandon := w.abandonAcquiredChat(ctx, workerID, runnerID, chatID); errAbandon != nil {
return errors.Join(err, errAbandon)
return false, errors.Join(err, errAbandon)
}
return err
return false, err
}
return nil
return true, nil
}
func (w *chatWorker) abandonAcquiredChat(ctx context.Context, workerID uuid.UUID, runnerID uuid.UUID, chatID uuid.UUID) error {
+4 -2
View File
@@ -156,7 +156,7 @@ func TestWorker_TwoWorkersRaceSingleOwner(t *testing.T) {
require.Equal(t, call.input.RunnerID, latest.RunnerID.UUID)
}
func TestWorker_DrainsMultipleRunnableChatsOnWake(t *testing.T) {
func TestWorker_AcquisitionBatchSizeLimitsSuccessfulAcquisitions(t *testing.T) {
t.Parallel()
f := newWorkerTestFixture(t)
first := f.createRunningChat(t)
@@ -165,12 +165,14 @@ func TestWorker_DrainsMultipleRunnableChatsOnWake(t *testing.T) {
starter := newRecordingTaskStarter()
opts := testOptions(t, f, starter)
opts.AcquisitionBatchSize = 1
startWorker(t, opts)
worker := startWorker(t, opts)
want := map[uuid.UUID]bool{first.ID: true, second.ID: true, third.ID: true}
for range 3 {
call := starter.waitCall(t, taskKindGeneration, uuid.Nil)
delete(want, call.input.ChatID)
starter.assertNoCall(t)
worker.Wake()
}
require.Empty(t, want)
}