Skip to content

Commit 820128f

Browse files
committed
test(agenttask,assistant): pin agent_wait non-blocking and Await subscription semantics
1 parent 2520aa6 commit 820128f

3 files changed

Lines changed: 249 additions & 7 deletions

File tree

‎internal/agenttask/service.go‎

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -109,6 +109,7 @@ type Options struct {
109109
type Service struct {
110110
runner Runner
111111
getTaskFn func(context.Context, string) (*database.TaskEntity, bool, error)
112+
awaitGetFn func(context.Context, string) (*database.AgentTaskEntity, bool, error)
112113
renewLeaseFn func(context.Context, string, string, time.Time) (bool, error)
113114
active map[string]context.CancelFunc
114115
cancelSources map[string]string
@@ -126,6 +127,7 @@ type Service struct {
126127
mu sync.Mutex
127128
lifecycle sync.Mutex
128129
nextSubscriber uint64
130+
awaitPollEvery time.Duration
129131
timeout time.Duration
130132
leaseDuration time.Duration
131133
leaseHeartbeatInterval time.Duration
@@ -187,6 +189,8 @@ func NewStopped(ctx context.Context, options *Options) (*Service, error) {
187189
nextSubscriber: 0, wg: sync.WaitGroup{}, timeout: timeout,
188190
concurrency: concurrency, sessionConcurrency: sessionConcurrency, logger: logger, leaseOwner: leaseOwner,
189191
getTaskFn: options.Tasks.Get, renewLeaseFn: options.Tasks.RenewLease, leaseDuration: leaseDuration,
192+
awaitGetFn: options.AgentTasks.Get,
193+
awaitPollEvery: awaitPollInterval,
190194
leaseHeartbeatInterval: leaseHeartbeatInterval,
191195
leaseRenewalRetryInterval: leaseRenewalRetryInterval,
192196
leaseRenewalWindow: leaseRenewalWindow, mu: sync.Mutex{},
@@ -583,13 +587,13 @@ func (service *Service) Await(ctx context.Context, taskID string) (*database.Age
583587
subscription := service.Subscribe(taskID)
584588
defer subscription.Cancel()
585589

586-
ticker := time.NewTicker(awaitPollInterval)
590+
ticker := time.NewTicker(service.awaitPollEvery)
587591
defer ticker.Stop()
588592

589593
events := subscription.Events
590594

591595
for {
592-
task, found, err := service.agentTasks.Get(ctx, taskID)
596+
task, found, err := service.awaitGetFn(ctx, taskID)
593597
if err != nil {
594598
return nil, oops.In("agenttask").Code("await_task").Wrapf(err, "await agent task")
595599
}

‎internal/agenttask/service_internal_test.go‎

Lines changed: 169 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -21,8 +21,15 @@ import (
2121
)
2222

2323
const (
24-
childSessionName = "child"
25-
workerName = "worker"
24+
childSessionName = "child"
25+
workerName = "worker"
26+
taskSucceededKind = "task_succeeded"
27+
taskFailedKind = "task_failed"
28+
taskCanceledKind = "task_canceled"
29+
30+
// awaitTestTimeout bounds Await calls in tests so a failed terminal
31+
// transition fails fast instead of consuming the suite timeout.
32+
awaitTestTimeout = 5 * time.Second
2633
)
2734

2835
func emptyService() *Service {
@@ -32,6 +39,7 @@ func emptyService() *Service {
3239
cancel: nil, done: nil, sessionSlots: nil, tasks: nil, logger: nil, leaseOwner: "", wg: sync.WaitGroup{},
3340
nextSubscriber: 0, timeout: 0, sessionConcurrency: 0, leaseDuration: 0,
3441
leaseHeartbeatInterval: 0, leaseRenewalRetryInterval: 0, leaseRenewalWindow: 0,
42+
awaitGetFn: nil, awaitPollEvery: awaitPollInterval,
3543
mu: sync.Mutex{}, lifecycle: sync.Mutex{}, concurrency: 0, started: false, closed: false,
3644
}
3745
}
@@ -101,6 +109,7 @@ func serviceWithRepositories(tasks *database.TaskRepository, agentTasks *databas
101109
service.tasks = tasks
102110
service.getTaskFn = tasks.Get
103111
service.agentTasks = agentTasks
112+
service.awaitGetFn = agentTasks.Get
104113
service.active = make(map[string]context.CancelFunc)
105114
service.cancelSources = make(map[string]string)
106115
service.subscribers = make(map[string]map[uint64]chan database.TaskEventEntity)
@@ -438,7 +447,7 @@ func TestServiceInternalTerminalEventSurvivesFullSubscriberBuffer(t *testing.T)
438447
service.publish(&database.TaskEventEntity{
439448
TaskID: taskID, Sequence: eventBuffer + 1,
440449
Event: database.EventEntity{
441-
CreatedAt: time.Time{}, ID: "", Kind: "task_succeeded", PayloadJSON: "",
450+
CreatedAt: time.Time{}, ID: "", Kind: taskSucceededKind, PayloadJSON: "",
442451
},
443452
})
444453

@@ -448,7 +457,7 @@ func TestServiceInternalTerminalEventSurvivesFullSubscriberBuffer(t *testing.T)
448457
}
449458

450459
assert.Equal(t, int64(2), received[0].Sequence, "oldest stream event should be evicted")
451-
assert.Equal(t, "task_succeeded", received[len(received)-1].Event.Kind)
460+
assert.Equal(t, taskSucceededKind, received[len(received)-1].Event.Kind)
452461
}
453462

454463
func TestServiceInternalClosesAllSubscriptions(t *testing.T) {
@@ -602,7 +611,7 @@ func TestServiceInternalCancelRecordsProvenance(t *testing.T) {
602611
latest, found, err := fixture.tasks.LatestEvent(t.Context(), created.Task.ID)
603612
require.NoError(t, err)
604613
require.True(t, found)
605-
assert.Equal(t, "task_canceled", latest.Event.Kind)
614+
assert.Equal(t, taskCanceledKind, latest.Event.Kind)
606615
assert.JSONEq(t, `{"canceled_by":"`+test.source+`"}`, latest.Event.PayloadJSON)
607616
})
608617
}
@@ -1170,3 +1179,158 @@ func TestServiceInternalLeaseRenewsThroughoutLongRun(t *testing.T) {
11701179
receive(t, done, &struct{}{}, "timed out waiting for lease renewal to stop")
11711180
assert.GreaterOrEqual(t, renewals.Load(), int32(wantedRenewals))
11721181
}
1182+
1183+
// gateAwaitOnQueuedRead parks Await in its wait loop by signaling when its
1184+
// first repository read observes the queued task. The returned channel
1185+
// closes once; terminal transitions sent after it can only be observed
1186+
// through the subscription or the safety-net poll.
1187+
func gateAwaitOnQueuedRead(t *testing.T, service *Service) <-chan struct{} {
1188+
t.Helper()
1189+
1190+
readQueued := make(chan struct{})
1191+
1192+
var once sync.Once
1193+
1194+
inner := service.awaitGetFn
1195+
service.awaitGetFn = func(ctx context.Context, id string) (*database.AgentTaskEntity, bool, error) {
1196+
entity, found, err := inner(ctx, id)
1197+
if err == nil && found && entity.Task.State == database.TaskQueued {
1198+
once.Do(func() { close(readQueued) })
1199+
}
1200+
1201+
return entity, found, err
1202+
}
1203+
1204+
return readQueued
1205+
}
1206+
1207+
// TestServiceAwaitIsWokenBySubscriptionPublish pins the fire-once wait
1208+
// contract: the terminal event published through the subscription wakes
1209+
// Await without waiting for the safety-net poll.
1210+
func TestServiceAwaitIsWokenBySubscriptionPublish(t *testing.T) {
1211+
t.Parallel()
1212+
1213+
fixture := newServiceRepositoryFixture(t)
1214+
task := fixture.createQueuedAgentTask(t)
1215+
service := serviceWithRepositories(fixture.tasks, fixture.agentTasks)
1216+
// Stretch the safety-net poll beyond the test deadline so only the
1217+
// published subscription event can wake Await.
1218+
service.awaitPollEvery = 10 * time.Minute
1219+
readQueued := gateAwaitOnQueuedRead(t, service)
1220+
transitionErr := make(chan error, 1)
1221+
1222+
go func() {
1223+
<-readQueued
1224+
1225+
_, err := service.tasks.Finish(t.Context(), &database.TaskFinish{
1226+
TaskID: task.Task.ID, EventKind: taskSucceededKind, Result: "done",
1227+
ErrorCode: "", ErrorMessage: "", PayloadJSON: `{}`, LeaseOwner: "",
1228+
TargetState: database.TaskSucceeded, From: []database.TaskState{database.TaskQueued},
1229+
})
1230+
if err != nil {
1231+
transitionErr <- err
1232+
1233+
return
1234+
}
1235+
1236+
service.publishLatest(t.Context(), task.Task.ID)
1237+
}()
1238+
1239+
awaitCtx, awaitCancel := context.WithTimeout(t.Context(), awaitTestTimeout)
1240+
defer awaitCancel()
1241+
1242+
awaited, err := service.Await(awaitCtx, task.Task.ID)
1243+
require.NoError(t, err)
1244+
1245+
select {
1246+
case err := <-transitionErr:
1247+
t.Fatalf("finish task for Await wake: %v", err)
1248+
default:
1249+
}
1250+
1251+
assert.Equal(t, database.TaskSucceeded, awaited.Task.State)
1252+
assert.Nil(t, service.subscribers[task.Task.ID], "Await should cancel its subscription before returning")
1253+
}
1254+
1255+
// TestServiceAwaitSafetyNetObservesCrossProcessTerminalState pins the poll
1256+
// fallback: when no in-process publisher exists, the bounded poll still
1257+
// observes a terminal state written by another process.
1258+
func TestServiceAwaitSafetyNetObservesCrossProcessTerminalState(t *testing.T) {
1259+
t.Parallel()
1260+
1261+
fixture := newServiceRepositoryFixture(t)
1262+
task := fixture.createQueuedAgentTask(t)
1263+
service := serviceWithRepositories(fixture.tasks, fixture.agentTasks)
1264+
readQueued := gateAwaitOnQueuedRead(t, service)
1265+
transitionErr := make(chan error, 1)
1266+
1267+
go func() {
1268+
<-readQueued
1269+
1270+
_, err := service.tasks.Transition(
1271+
t.Context(), task.Task.ID,
1272+
[]database.TaskState{database.TaskQueued}, database.TaskCanceled,
1273+
taskCanceledKind, database.CancelEventPayload(database.CancelSourceParent),
1274+
)
1275+
if err != nil {
1276+
transitionErr <- err
1277+
}
1278+
}()
1279+
1280+
awaitCtx, awaitCancel := context.WithTimeout(t.Context(), awaitTestTimeout)
1281+
defer awaitCancel()
1282+
1283+
awaited, err := service.Await(awaitCtx, task.Task.ID)
1284+
require.NoError(t, err)
1285+
1286+
select {
1287+
case err := <-transitionErr:
1288+
t.Fatalf("transition task for Await poll: %v", err)
1289+
default:
1290+
}
1291+
1292+
assert.Equal(t, database.TaskCanceled, awaited.Task.State)
1293+
}
1294+
1295+
// TestServiceAwaitReturnsPromptlyOnAlreadyTerminalTask pins that a task
1296+
// which is already terminal resolves on the first repository read without
1297+
// waiting for a subscription event or the poll ticker.
1298+
func TestServiceAwaitReturnsPromptlyOnAlreadyTerminalTask(t *testing.T) {
1299+
t.Parallel()
1300+
1301+
cases := []struct {
1302+
name string
1303+
state database.TaskState
1304+
kind string
1305+
}{
1306+
{name: "succeeded", state: database.TaskSucceeded, kind: taskSucceededKind},
1307+
{name: "failed", state: database.TaskFailed, kind: taskFailedKind},
1308+
{name: taskCanceledKind, state: database.TaskCanceled, kind: taskCanceledKind},
1309+
{name: "interrupted", state: database.TaskInterrupted, kind: taskInterruptedEvent},
1310+
}
1311+
1312+
for _, testCase := range cases {
1313+
t.Run(testCase.name, func(t *testing.T) {
1314+
t.Parallel()
1315+
1316+
fixture := newServiceRepositoryFixture(t)
1317+
task := fixture.createQueuedAgentTask(t)
1318+
service := serviceWithRepositories(fixture.tasks, fixture.agentTasks)
1319+
1320+
changed, err := service.tasks.Transition(
1321+
t.Context(), task.Task.ID,
1322+
[]database.TaskState{database.TaskQueued}, testCase.state,
1323+
testCase.kind, `{}`,
1324+
)
1325+
require.NoError(t, err)
1326+
require.True(t, changed)
1327+
1328+
awaitCtx, awaitCancel := context.WithTimeout(t.Context(), awaitTestTimeout)
1329+
defer awaitCancel()
1330+
1331+
awaited, err := service.Await(awaitCtx, task.Task.ID)
1332+
require.NoError(t, err)
1333+
assert.Equal(t, testCase.state, awaited.Task.State)
1334+
})
1335+
}
1336+
}

‎internal/assistant/agent_tool_internal_test.go‎

Lines changed: 74 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,7 @@ type agentControllerStub struct {
3232
listed []database.AgentTaskEntity
3333
lastLimit int
3434
getCalls int
35+
awaitCalls int
3536
subscriptions int
3637
found bool
3738
}
@@ -73,6 +74,8 @@ func (stub *agentControllerStub) Cancel(context.Context, string, string, string)
7374
return &stub.task.Task, found, stub.cancelErr
7475
}
7576
func (stub *agentControllerStub) Await(context.Context, string) (*database.AgentTaskEntity, error) {
77+
stub.awaitCalls++
78+
7679
return stub.task, nil
7780
}
7881
func (stub *agentControllerStub) SubscribeAgentTask(
@@ -422,3 +425,74 @@ func agentToolTask(id, owner string, state database.TaskState) *database.AgentTa
422425
Model: "", Provider: "", PolicyJSON: `{}`, UsageJSON: `{}`, Depth: 1,
423426
}
424427
}
428+
429+
// TestAgentWaitToolReturnsImmediatelyWithoutBlocking pins the non-blocking
430+
// contract of agent_wait: one ownership Get per invocation, no Await
431+
// subscription, and the current state returned as-is so the parent can keep
432+
// working while the agent runs.
433+
func TestAgentWaitToolReturnsImmediatelyWithoutBlocking(t *testing.T) {
434+
t.Parallel()
435+
436+
cases := []struct {
437+
name string
438+
state database.TaskState
439+
resultText string
440+
wantContain string
441+
}{
442+
{
443+
name: "running task reports state",
444+
state: database.TaskRunning,
445+
resultText: "",
446+
wantContain: "is running",
447+
},
448+
{
449+
name: "queued task reports state",
450+
state: database.TaskQueued,
451+
resultText: "",
452+
wantContain: "is queued",
453+
},
454+
{
455+
name: "terminal task reports result",
456+
state: database.TaskSucceeded,
457+
resultText: "summary of findings",
458+
wantContain: "summary of findings",
459+
},
460+
}
461+
462+
for _, testCase := range cases {
463+
t.Run(testCase.name, func(t *testing.T) {
464+
t.Parallel()
465+
466+
task := agentToolTask("task", "owner", testCase.state)
467+
task.Task.Result = testCase.resultText
468+
stub := newAgentControllerStub(task, nil, true)
469+
executor := newAgentToolExecutor(
470+
stub, nil, isolatedAgentCatalog(t), agentWaitToolName, "owner", "",
471+
)
472+
473+
result, err := executor.Execute(t.Context(), agentArguments(t, `{"task_id":" task "}`))
474+
require.NoError(t, err)
475+
assert.Contains(t, result.Text(), testCase.wantContain)
476+
assert.Equal(t, 1, stub.getCalls, "agent_wait must perform exactly one Get")
477+
assert.Equal(t, 0, stub.subscriptions, "agent_wait must not subscribe")
478+
})
479+
}
480+
}
481+
482+
// TestAgentWaitToolNeverCallsAwait pins that agent_wait stays non-blocking at
483+
// the controller boundary: a blocking Await call would freeze the parent tool
484+
// loop and reintroduce the polling pressure the design removed.
485+
func TestAgentWaitToolNeverCallsAwait(t *testing.T) {
486+
t.Parallel()
487+
488+
stub := newAgentControllerStub(
489+
agentToolTask("task", "owner", database.TaskRunning), nil, true,
490+
)
491+
executor := newAgentToolExecutor(
492+
stub, nil, isolatedAgentCatalog(t), agentWaitToolName, "owner", "",
493+
)
494+
495+
_, err := executor.Execute(t.Context(), agentArguments(t, `{"task_id":"task"}`))
496+
require.NoError(t, err)
497+
assert.Zero(t, stub.awaitCalls, "agent_wait must never call the blocking Await")
498+
}

0 commit comments

Comments
 (0)