@@ -21,8 +21,15 @@ import (
2121)
2222
2323const (
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
2835func 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
454463func 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+ }
0 commit comments