mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
Generated
+4
@@ -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"
|
||||
|
||||
Generated
+4
@@ -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"
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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] {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
Generated
+45
@@ -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()
|
||||
|
||||
Generated
+1
-1
@@ -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);
|
||||
|
||||
|
||||
@@ -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;
|
||||
Generated
+8
-11
@@ -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.
|
||||
|
||||
Generated
+183
-142
@@ -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)
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user