test_event_batch.py 35 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904905906907908909910911912913914915916917918919920921922923924925926927928929930931932933934935936937938939940941942943944945946947948949950951952953954955956957958959960961962963964965966967968969970971972973974975976977978979980981982983984985986987988989990991992993994995996997998999100010011002100310041005100610071008100910101011101210131014101510161017101810191020102110221023102410251026102710281029103010311032103310341035103610371038103910401041104210431044104510461047104810491050105110521053105410551056105710581059106010611062106310641065106610671068106910701071107210731074107510761077107810791080108110821083108410851086108710881089109010911092109310941095109610971098109911001101110211031104110511061107110811091110111111121113111411151116111711181119112011211122112311241125112611271128112911301131113211331134113511361137113811391140114111421143114411451146114711481149115011511152115311541155115611571158115911601161116211631164116511661167116811691170117111721173117411751176
  1. from __future__ import annotations
  2. import asyncio
  3. import json
  4. import threading
  5. from collections.abc import Callable
  6. from typing import Any
  7. import pytest
  8. import agent_lab.application.events as events_module
  9. from agent_lab.application.events import (
  10. EventDefinition,
  11. EventKernel,
  12. EventRegistry,
  13. EventRequest,
  14. EventSource,
  15. EventStatus,
  16. ResultPolicy,
  17. )
  18. from agent_lab.application.tools import ToolDefinition, ToolRegistry
  19. from agent_lab.domain.events import ToolCallEvent
  20. def _executor_type():
  21. executor_type = getattr(events_module, "EventBatchExecutor", None)
  22. assert executor_type is not None, "EventBatchExecutor is not implemented"
  23. return executor_type
  24. def _definition(
  25. name: str,
  26. handler: Callable[[EventRequest], Any],
  27. *,
  28. conflict_keys: tuple[str, ...] = (),
  29. timeout_seconds: float | None = None,
  30. terminal: bool = False,
  31. normalizer: Callable[[dict[str, Any]], dict[str, Any]] | None = None,
  32. ) -> EventDefinition:
  33. return EventDefinition(
  34. name=name,
  35. description=f"Execute {name}.",
  36. parameters={"type": "object"},
  37. handler=handler,
  38. conflict_keys=conflict_keys,
  39. timeout_seconds=timeout_seconds,
  40. terminal=terminal,
  41. normalizer=normalizer,
  42. )
  43. def _request(event_id: str, name: str, arguments: dict[str, Any] | None = None):
  44. return EventRequest(
  45. id=event_id,
  46. name=name,
  47. arguments=arguments or {},
  48. )
  49. def _executor(
  50. definitions: list[EventDefinition],
  51. *,
  52. max_parallel_events: int = 4,
  53. batch_timeout_seconds: float = 1.0,
  54. **options: Any,
  55. ):
  56. return _executor_type()(
  57. EventKernel(EventRegistry(definitions)),
  58. max_parallel_events=max_parallel_events,
  59. batch_timeout_seconds=batch_timeout_seconds,
  60. **options,
  61. )
  62. @pytest.mark.asyncio
  63. async def test_batch_executor_bounds_global_parallelism():
  64. active = 0
  65. max_active = 0
  66. async def handler(request: EventRequest) -> dict[str, Any]:
  67. nonlocal active, max_active
  68. active += 1
  69. max_active = max(max_active, active)
  70. await asyncio.sleep(0.02)
  71. active -= 1
  72. return {"event_id": request.id}
  73. executor = _executor(
  74. [_definition("example.work", handler)],
  75. max_parallel_events=2,
  76. )
  77. batch = await executor.execute(
  78. [
  79. _request(str(index), "example.work", {"index": index})
  80. for index in range(4)
  81. ],
  82. scope="session-1:turn-1",
  83. )
  84. assert max_active == 2
  85. assert [result.status for result in batch.results] == [
  86. EventStatus.SUCCESS,
  87. ] * 4
  88. @pytest.mark.asyncio
  89. async def test_batch_executor_preserves_input_order_when_completion_order_differs():
  90. release_first = asyncio.Event()
  91. second_finished = asyncio.Event()
  92. async def handler(request: EventRequest) -> dict[str, Any]:
  93. if request.id == "first":
  94. await release_first.wait()
  95. else:
  96. second_finished.set()
  97. return {"event_id": request.id}
  98. executor = _executor([_definition("example.work", handler)])
  99. task = asyncio.create_task(
  100. executor.execute(
  101. [
  102. _request("first", "example.work", {"order": 1}),
  103. _request("second", "example.work", {"order": 2}),
  104. ],
  105. scope="session-1:turn-1",
  106. )
  107. )
  108. await asyncio.wait_for(second_finished.wait(), timeout=0.2)
  109. release_first.set()
  110. batch = await asyncio.wait_for(task, timeout=0.2)
  111. assert [result.event_id for result in batch.results] == ["first", "second"]
  112. @pytest.mark.asyncio
  113. async def test_batch_executor_serializes_shared_conflict_keys():
  114. active = 0
  115. max_active = 0
  116. async def handler(request: EventRequest) -> dict[str, Any]:
  117. nonlocal active, max_active
  118. active += 1
  119. max_active = max(max_active, active)
  120. await asyncio.sleep(0.02)
  121. active -= 1
  122. return {"event_id": request.id}
  123. executor = _executor(
  124. [
  125. _definition(
  126. "example.write",
  127. handler,
  128. conflict_keys=("shared-resource",),
  129. )
  130. ]
  131. )
  132. await executor.execute(
  133. [
  134. _request("first", "example.write", {"order": 1}),
  135. _request("second", "example.write", {"order": 2}),
  136. ],
  137. scope="session-1:turn-1",
  138. )
  139. assert max_active == 1
  140. @pytest.mark.asyncio
  141. async def test_batch_executor_isolates_per_event_timeout_from_successful_sibling():
  142. async def slow_handler(request: EventRequest) -> dict[str, Any]:
  143. await asyncio.sleep(1)
  144. return {"event_id": request.id}
  145. executor = _executor(
  146. [
  147. _definition(
  148. "example.slow",
  149. slow_handler,
  150. timeout_seconds=0.01,
  151. ),
  152. _definition(
  153. "example.fast",
  154. lambda request: {"event_id": request.id},
  155. ),
  156. ],
  157. batch_timeout_seconds=0.2,
  158. )
  159. batch = await executor.execute(
  160. [
  161. _request("slow", "example.slow"),
  162. _request("fast", "example.fast"),
  163. ],
  164. scope="session-1:turn-1",
  165. )
  166. assert [result.status for result in batch.results] == [
  167. EventStatus.TIMEOUT,
  168. EventStatus.SUCCESS,
  169. ]
  170. assert batch.deadline_exceeded is False
  171. @pytest.mark.asyncio
  172. async def test_batch_executor_applies_overall_deadline_to_waiting_work():
  173. async def blocked_handler(request: EventRequest) -> dict[str, Any]:
  174. await asyncio.Event().wait()
  175. return {"event_id": request.id}
  176. executor = _executor(
  177. [_definition("example.blocked", blocked_handler)],
  178. max_parallel_events=1,
  179. batch_timeout_seconds=0.02,
  180. )
  181. batch = await executor.execute(
  182. [
  183. _request("first", "example.blocked"),
  184. _request("second", "example.blocked"),
  185. ],
  186. scope="session-1:turn-1",
  187. )
  188. assert [result.status for result in batch.results] == [
  189. EventStatus.TIMEOUT,
  190. EventStatus.TIMEOUT,
  191. ]
  192. assert batch.deadline_exceeded is True
  193. @pytest.mark.asyncio
  194. async def test_batch_executor_coalesces_exact_duplicates_inside_one_batch():
  195. calls = 0
  196. def handler(request: EventRequest) -> dict[str, Any]:
  197. nonlocal calls
  198. calls += 1
  199. return {"event_id": request.id, "arguments": request.arguments}
  200. executor = _executor([_definition("example.write", handler)])
  201. batch = await executor.execute(
  202. [
  203. _request("same", "example.write", {"a": 1, "b": 2}),
  204. _request("same", "example.write", {"b": 2, "a": 1}),
  205. ],
  206. scope="session-1:turn-1",
  207. )
  208. assert calls == 1
  209. assert batch.coalesced_count == 1
  210. assert batch.results[0].event_id == batch.results[1].event_id == "same"
  211. assert batch.results[0].payload == batch.results[1].payload
  212. assert batch.results[0].deduplicated_from is None
  213. assert batch.results[1].deduplicated_from == "same"
  214. @pytest.mark.asyncio
  215. async def test_batch_executor_replay_key_includes_scope_event_id_and_request():
  216. calls = 0
  217. def handler(request: EventRequest) -> dict[str, Any]:
  218. nonlocal calls
  219. calls += 1
  220. return {"call": calls, "arguments": request.arguments}
  221. executor = _executor([_definition("example.write", handler)])
  222. original = _request("same", "example.write", {"a": 1, "b": 2})
  223. canonical_replay = _request(
  224. "same",
  225. "example.write",
  226. {"b": 2, "a": 1},
  227. )
  228. first = await executor.execute([original], scope="session-1:turn-1")
  229. replay = await executor.execute(
  230. [canonical_replay],
  231. scope="session-1:turn-1",
  232. )
  233. changed_request = await executor.execute(
  234. [_request("same", "example.write", {"a": 2})],
  235. scope="session-1:turn-1",
  236. )
  237. changed_id = await executor.execute(
  238. [_request("other", "example.write", {"a": 1, "b": 2})],
  239. scope="session-1:turn-1",
  240. )
  241. changed_scope = await executor.execute(
  242. [original],
  243. scope="session-1:turn-2",
  244. )
  245. assert first.results[0] == replay.results[0]
  246. assert replay.replayed_count == 1
  247. assert changed_request.replayed_count == 0
  248. assert changed_id.replayed_count == 0
  249. assert changed_scope.replayed_count == 0
  250. assert calls == 4
  251. @pytest.mark.asyncio
  252. async def test_batch_executor_runs_terminal_events_after_all_nonterminal_results():
  253. observed: list[str] = []
  254. async def nonterminal(request: EventRequest) -> dict[str, Any]:
  255. observed.append(f"start:{request.id}")
  256. await asyncio.sleep(0.01)
  257. observed.append(f"finish:{request.id}")
  258. return {"event_id": request.id}
  259. def terminal(request: EventRequest) -> dict[str, Any]:
  260. observed.append(f"terminal:{request.id}")
  261. return {"event_id": request.id}
  262. executor = _executor(
  263. [
  264. _definition("example.work", nonterminal),
  265. _definition("example.terminate", terminal, terminal=True),
  266. ]
  267. )
  268. batch = await executor.execute(
  269. [
  270. _request("terminate", "example.terminate"),
  271. _request("first", "example.work", {"order": 1}),
  272. _request("second", "example.work", {"order": 2}),
  273. ],
  274. scope="session-1:turn-1",
  275. )
  276. terminal_index = observed.index("terminal:terminate")
  277. assert observed.index("finish:first") < terminal_index
  278. assert observed.index("finish:second") < terminal_index
  279. assert [result.event_id for result in batch.results] == [
  280. "terminate",
  281. "first",
  282. "second",
  283. ]
  284. @pytest.mark.asyncio
  285. async def test_batch_executor_does_not_swallow_external_cancellation():
  286. started = asyncio.Event()
  287. async def handler(request: EventRequest) -> dict[str, Any]:
  288. started.set()
  289. await asyncio.Event().wait()
  290. return {"event_id": request.id}
  291. executor = _executor([_definition("example.blocked", handler)])
  292. task = asyncio.create_task(
  293. executor.execute(
  294. [_request("blocked", "example.blocked")],
  295. scope="session-1:turn-1",
  296. )
  297. )
  298. await asyncio.wait_for(started.wait(), timeout=0.2)
  299. task.cancel()
  300. with pytest.raises(asyncio.CancelledError):
  301. await task
  302. @pytest.mark.asyncio
  303. async def test_batch_executor_coalesces_normalized_provider_arguments_across_ids():
  304. calls: list[EventRequest] = []
  305. def normalize(arguments: dict[str, Any]) -> dict[str, Any]:
  306. value = arguments["value"]
  307. return {"value": int(value) if isinstance(value, float) else value}
  308. executor = _executor(
  309. [
  310. _definition(
  311. "example.normalized",
  312. lambda request: calls.append(request) or {"value": request.arguments["value"]},
  313. normalizer=normalize,
  314. )
  315. ]
  316. )
  317. requests = [
  318. EventRequest(
  319. id="integer-id",
  320. name="example.normalized",
  321. arguments={"value": 30},
  322. source=EventSource.PROVIDER_RESOLVED,
  323. raw_arguments='{"value":30}',
  324. ),
  325. EventRequest(
  326. id="float-id",
  327. name="example.normalized",
  328. arguments={"value": 30.0},
  329. source=EventSource.PROVIDER_RESOLVED,
  330. raw_arguments='{"value":30.0}',
  331. ),
  332. ]
  333. batch = await executor.execute(requests, scope="normalized-scope")
  334. assert len(calls) == 1
  335. assert batch.coalesced_count == 1
  336. assert [result.event_id for result in batch.results] == [
  337. "integer-id",
  338. "float-id",
  339. ]
  340. assert batch.results[0].payload == {"value": 30}
  341. assert batch.results[1].payload == {"value": 30}
  342. assert getattr(batch.results[0], "deduplicated_from", None) is None
  343. assert getattr(batch.results[1], "deduplicated_from", None) == "integer-id"
  344. @pytest.mark.asyncio
  345. async def test_coalesced_clone_rewrites_only_correlated_payload_event_id():
  346. correlated = _executor(
  347. [
  348. _definition(
  349. "example.correlated",
  350. lambda request: {"event_id": request.id, "value": "same"},
  351. )
  352. ]
  353. )
  354. correlated_batch = await correlated.execute(
  355. [
  356. EventRequest(
  357. id="primary",
  358. name="example.correlated",
  359. arguments={},
  360. raw_arguments='{"call":"primary"}',
  361. ),
  362. EventRequest(
  363. id="duplicate",
  364. name="example.correlated",
  365. arguments={},
  366. raw_arguments='{"call":"duplicate"}',
  367. ),
  368. ],
  369. scope="correlated-payload",
  370. )
  371. primary, duplicate = correlated_batch.results
  372. assert primary.event_id == "primary"
  373. assert duplicate.event_id == "duplicate"
  374. assert duplicate.raw_arguments == '{"call":"duplicate"}'
  375. assert duplicate.payload == {"event_id": "duplicate", "value": "same"}
  376. assert getattr(duplicate, "deduplicated_from", None) == "primary"
  377. domain = _executor(
  378. [
  379. _definition(
  380. "example.domain",
  381. lambda request: {"event_id": "domain-object"},
  382. )
  383. ]
  384. )
  385. domain_batch = await domain.execute(
  386. [
  387. _request("primary", "example.domain"),
  388. _request("duplicate", "example.domain"),
  389. ],
  390. scope="domain-payload",
  391. )
  392. assert domain_batch.results[1].payload == {"event_id": "domain-object"}
  393. assert getattr(domain_batch.results[1], "deduplicated_from", None) == "primary"
  394. @pytest.mark.asyncio
  395. async def test_batch_executor_tool_replies_keep_each_coalesced_call_id():
  396. registry = ToolRegistry(
  397. [
  398. ToolDefinition(
  399. name="example.tool",
  400. description="Execute a coalesced tool.",
  401. parameters={"type": "object"},
  402. handler=lambda event: {"event_id": event.id},
  403. )
  404. ]
  405. )
  406. executor = _executor_type()(
  407. EventKernel(registry.event_registry),
  408. batch_timeout_seconds=1,
  409. )
  410. events = [
  411. ToolCallEvent(id=event_id, name="example.tool", arguments={}, raw_arguments="{}")
  412. for event_id in ("first-id", "second-id")
  413. ]
  414. batch = await executor.execute(
  415. [
  416. registry.event_request(event, source=EventSource.PROVIDER_RESOLVED)
  417. for event in events
  418. ],
  419. scope="tool-scope",
  420. )
  421. replies = registry.tool_replies(events, batch)
  422. assert [reply.tool_call_id for reply in replies] == ["first-id", "second-id"]
  423. assert [json.loads(reply.content) for reply in replies] == [
  424. {"event_id": "first-id"},
  425. {"event_id": "second-id"},
  426. ]
  427. assert [result.event_id for result in batch.results] == [
  428. "first-id",
  429. "second-id",
  430. ]
  431. @pytest.mark.asyncio
  432. async def test_batch_policy_and_compact_summary_skip_logical_duplicates():
  433. registry = ToolRegistry(
  434. [
  435. ToolDefinition(
  436. name="example.template",
  437. description="Render one template.",
  438. parameters={"type": "object"},
  439. handler=lambda event: {"event_id": event.id},
  440. result_policy=ResultPolicy.TEMPLATE_FOLLOW_UP,
  441. result_message_factory=lambda result: "Completed once.",
  442. )
  443. ]
  444. )
  445. executor = _executor_type()(EventKernel(registry.event_registry))
  446. events = [
  447. ToolCallEvent(
  448. id=event_id,
  449. name="example.template",
  450. arguments={},
  451. raw_arguments="{}",
  452. )
  453. for event_id in ("primary", "duplicate")
  454. ]
  455. batch = await executor.execute(
  456. [registry.event_request(event) for event in events],
  457. scope="policy-dedupe",
  458. )
  459. decision = registry.batch_decision(batch)
  460. summary = registry.compact_results_message(batch)
  461. assert decision.template_messages == ("Completed once.",)
  462. assert summary is not None
  463. assert len(summary.content.splitlines()) == 2
  464. @pytest.mark.asyncio
  465. async def test_batch_key_falls_back_to_original_arguments_when_normalizer_fails():
  466. calls = 0
  467. def normalize(arguments: dict[str, Any]) -> dict[str, Any]:
  468. if isinstance(arguments.get("value"), str):
  469. raise ValueError("unsupported value")
  470. return {"value": int(arguments["value"])}
  471. def handler(request: EventRequest) -> dict[str, Any]:
  472. nonlocal calls
  473. calls += 1
  474. return {"value": request.arguments["value"]}
  475. executor = _executor(
  476. [_definition("example.normalized", handler, normalizer=normalize)]
  477. )
  478. batch = await executor.execute(
  479. [
  480. EventRequest(
  481. id="valid",
  482. name="example.normalized",
  483. arguments={"value": 30},
  484. source=EventSource.PROVIDER_RESOLVED,
  485. ),
  486. EventRequest(
  487. id="invalid",
  488. name="example.normalized",
  489. arguments={"value": "30"},
  490. source=EventSource.PROVIDER_RESOLVED,
  491. ),
  492. ],
  493. scope="normalizer-failure",
  494. )
  495. assert batch.coalesced_count == 0
  496. assert [result.status for result in batch.results] == [
  497. EventStatus.SUCCESS,
  498. EventStatus.INVALID_ARGUMENTS,
  499. ]
  500. assert calls == 1
  501. @pytest.mark.asyncio
  502. @pytest.mark.parametrize(
  503. ("invalid_arguments", "raw_arguments"),
  504. [
  505. ({"value": float("nan")}, '{"value":NaN}'),
  506. ({"value": object()}, '{"value":"object"}'),
  507. ],
  508. )
  509. async def test_invalid_json_arguments_do_not_coalesce_with_valid_empty_object(
  510. invalid_arguments: dict[str, Any],
  511. raw_arguments: str,
  512. ):
  513. calls = 0
  514. def handler(request: EventRequest) -> dict[str, Any]:
  515. nonlocal calls
  516. calls += 1
  517. return {"event_id": request.id}
  518. executor = _executor([_definition("example.strict", handler)])
  519. batch = await executor.execute(
  520. [
  521. EventRequest(
  522. id="valid",
  523. name="example.strict",
  524. arguments={},
  525. source=EventSource.PROVIDER_RESOLVED,
  526. raw_arguments="{}",
  527. ),
  528. EventRequest(
  529. id="invalid",
  530. name="example.strict",
  531. arguments=invalid_arguments,
  532. source=EventSource.PROVIDER_RESOLVED,
  533. raw_arguments=raw_arguments,
  534. ),
  535. ],
  536. scope="strict-json-batch",
  537. )
  538. assert batch.coalesced_count == 0
  539. assert [result.status for result in batch.results] == [
  540. EventStatus.SUCCESS,
  541. EventStatus.INVALID_ARGUMENTS,
  542. ]
  543. assert calls == 1
  544. @pytest.mark.asyncio
  545. @pytest.mark.parametrize(
  546. ("invalid_arguments", "raw_arguments"),
  547. [
  548. ({"value": float("nan")}, '{"value":NaN}'),
  549. ({"value": object()}, '{"value":"object"}'),
  550. ],
  551. )
  552. async def test_invalid_json_arguments_do_not_replay_valid_success(
  553. invalid_arguments: dict[str, Any],
  554. raw_arguments: str,
  555. ):
  556. executor = _executor(
  557. [_definition("example.strict", lambda request: {"event_id": request.id})]
  558. )
  559. valid = EventRequest(
  560. id="same-id",
  561. name="example.strict",
  562. arguments={},
  563. source=EventSource.PROVIDER_RESOLVED,
  564. raw_arguments="{}",
  565. )
  566. invalid = EventRequest(
  567. id="same-id",
  568. name="example.strict",
  569. arguments=invalid_arguments,
  570. source=EventSource.PROVIDER_RESOLVED,
  571. raw_arguments=raw_arguments,
  572. )
  573. first = await executor.execute([valid], scope="strict-json-replay")
  574. second = await executor.execute([invalid], scope="strict-json-replay")
  575. assert first.results[0].status is EventStatus.SUCCESS
  576. assert second.replayed_count == 0
  577. assert second.results[0].status is EventStatus.INVALID_ARGUMENTS
  578. @pytest.mark.asyncio
  579. async def test_non_json_normalizer_output_has_distinct_canonical_key():
  580. calls = 0
  581. def normalize(arguments: dict[str, Any]) -> dict[str, Any]:
  582. if arguments.get("invalid"):
  583. return {"value": object()}
  584. return {}
  585. def handler(request: EventRequest) -> dict[str, Any]:
  586. nonlocal calls
  587. calls += 1
  588. return {"event_id": request.id}
  589. executor = _executor(
  590. [_definition("example.normalized", handler, normalizer=normalize)]
  591. )
  592. batch = await executor.execute(
  593. [
  594. EventRequest(
  595. id="valid",
  596. name="example.normalized",
  597. arguments={},
  598. source=EventSource.PROVIDER_RESOLVED,
  599. ),
  600. EventRequest(
  601. id="invalid",
  602. name="example.normalized",
  603. arguments={"invalid": True},
  604. source=EventSource.PROVIDER_RESOLVED,
  605. ),
  606. ],
  607. scope="normalizer-invalid-json",
  608. )
  609. assert batch.coalesced_count == 0
  610. assert [result.status for result in batch.results] == [
  611. EventStatus.SUCCESS,
  612. EventStatus.INVALID_ARGUMENTS,
  613. ]
  614. assert calls == 1
  615. @pytest.mark.asyncio
  616. async def test_concurrent_replay_waiters_share_execution_and_cancel_independently():
  617. calls = 0
  618. started = asyncio.Event()
  619. release = asyncio.Event()
  620. async def handler(request: EventRequest) -> dict[str, Any]:
  621. nonlocal calls
  622. calls += 1
  623. started.set()
  624. await release.wait()
  625. return {"event_id": request.id}
  626. executor = _executor([_definition("example.shared", handler)])
  627. request = _request("same-id", "example.shared")
  628. first = asyncio.create_task(executor.execute([request], scope="shared-scope"))
  629. second = asyncio.create_task(executor.execute([request], scope="shared-scope"))
  630. await asyncio.wait_for(started.wait(), timeout=0.2)
  631. first.cancel()
  632. with pytest.raises(asyncio.CancelledError):
  633. await first
  634. release.set()
  635. second_batch = await asyncio.wait_for(second, timeout=0.2)
  636. assert calls == 1
  637. assert second_batch.results[0].status is EventStatus.SUCCESS
  638. @pytest.mark.asyncio
  639. async def test_release_scope_removes_replay_entries_for_reuse():
  640. calls = 0
  641. def handler(request: EventRequest) -> dict[str, Any]:
  642. nonlocal calls
  643. calls += 1
  644. return {"call": calls}
  645. executor = _executor([_definition("example.cached", handler)])
  646. request = _request("same-id", "example.cached")
  647. await executor.execute([request], scope="released-scope")
  648. executor.release_scope("released-scope")
  649. batch = await executor.execute([request], scope="released-scope")
  650. assert calls == 2
  651. assert batch.replayed_count == 0
  652. @pytest.mark.asyncio
  653. async def test_release_scope_starts_new_generation_before_old_task_finishes():
  654. calls = 0
  655. started = [asyncio.Event(), asyncio.Event()]
  656. releases = [asyncio.Event(), asyncio.Event()]
  657. async def handler(request: EventRequest) -> dict[str, Any]:
  658. nonlocal calls
  659. call_index = calls
  660. calls += 1
  661. started[call_index].set()
  662. await releases[call_index].wait()
  663. return {"call": call_index + 1, "event_id": request.id}
  664. executor = _executor([_definition("example.generated", handler)])
  665. request = _request("same-id", "example.generated")
  666. first = asyncio.create_task(
  667. executor.execute([request], scope="generation-scope")
  668. )
  669. await asyncio.wait_for(started[0].wait(), timeout=0.2)
  670. executor.release_scope("generation-scope")
  671. second = asyncio.create_task(
  672. executor.execute([request], scope="generation-scope")
  673. )
  674. await asyncio.sleep(0.02)
  675. new_generation_started = started[1].is_set()
  676. releases[0].set()
  677. releases[1].set()
  678. first_batch, second_batch = await asyncio.gather(first, second)
  679. assert new_generation_started is True
  680. assert first_batch.results[0].payload["call"] == 1
  681. assert second_batch.results[0].payload["call"] == 2
  682. @pytest.mark.asyncio
  683. async def test_old_generation_completion_cannot_replace_new_inflight_or_cache():
  684. calls = 0
  685. started = [asyncio.Event(), asyncio.Event(), asyncio.Event()]
  686. releases = [asyncio.Event(), asyncio.Event(), asyncio.Event()]
  687. async def handler(request: EventRequest) -> dict[str, Any]:
  688. nonlocal calls
  689. call_index = calls
  690. calls += 1
  691. started[call_index].set()
  692. await releases[call_index].wait()
  693. return {"call": call_index + 1, "event_id": request.id}
  694. executor = _executor([_definition("example.generated", handler)])
  695. request = _request("same-id", "example.generated")
  696. old = asyncio.create_task(executor.execute([request], scope="reused-scope"))
  697. await asyncio.wait_for(started[0].wait(), timeout=0.2)
  698. executor.release_scope("reused-scope")
  699. current = asyncio.create_task(
  700. executor.execute([request], scope="reused-scope")
  701. )
  702. await asyncio.sleep(0.02)
  703. new_generation_started = started[1].is_set()
  704. if not new_generation_started:
  705. releases[0].set()
  706. await asyncio.gather(old, current)
  707. assert new_generation_started is True
  708. releases[0].set()
  709. old_batch = await asyncio.wait_for(old, timeout=0.2)
  710. waiter = asyncio.create_task(
  711. executor.execute([request], scope="reused-scope")
  712. )
  713. await asyncio.sleep(0.02)
  714. assert calls == 2
  715. releases[1].set()
  716. current_batch, waiter_batch = await asyncio.gather(current, waiter)
  717. replayed = await executor.execute([request], scope="reused-scope")
  718. assert old_batch.results[0].payload["call"] == 1
  719. assert current_batch.results[0].payload["call"] == 2
  720. assert waiter_batch.results[0].payload["call"] == 2
  721. assert replayed.replayed_count == 1
  722. assert replayed.results[0].payload["call"] == 2
  723. @pytest.mark.asyncio
  724. async def test_replay_cache_stays_bounded_across_250_scopes():
  725. executor = _executor(
  726. [_definition("example.cached", lambda request: {"id": request.id})],
  727. max_replay_entries=100,
  728. )
  729. request = _request("same-id", "example.cached")
  730. for index in range(250):
  731. await executor.execute([request], scope=f"scope-{index}")
  732. assert len(executor._replay_results) <= 100
  733. @pytest.mark.asyncio
  734. async def test_terminal_uses_independent_grace_after_sibling_batch_timeout():
  735. terminal_calls: list[str] = []
  736. async def blocked(request: EventRequest) -> dict[str, Any]:
  737. await asyncio.Event().wait()
  738. return {"event_id": request.id}
  739. def terminate(request: EventRequest) -> dict[str, Any]:
  740. terminal_calls.append(request.id)
  741. return {"terminated": True}
  742. executor = _executor(
  743. [
  744. _definition("example.blocked", blocked),
  745. _definition(
  746. "example.terminate",
  747. terminate,
  748. terminal=True,
  749. timeout_seconds=0.05,
  750. ),
  751. ],
  752. batch_timeout_seconds=0.01,
  753. terminal_grace_seconds=0.1,
  754. )
  755. batch = await executor.execute(
  756. [
  757. _request("sibling", "example.blocked"),
  758. _request("terminal", "example.terminate"),
  759. ],
  760. scope="terminal-grace",
  761. )
  762. assert [result.status for result in batch.results] == [
  763. EventStatus.TIMEOUT,
  764. EventStatus.SUCCESS,
  765. ]
  766. assert batch.deadline_exceeded is True
  767. assert terminal_calls == ["terminal"]
  768. @pytest.mark.asyncio
  769. async def test_terminal_grace_is_bounded_by_terminal_event_timeout():
  770. async def blocked_terminal(request: EventRequest) -> dict[str, Any]:
  771. await asyncio.Event().wait()
  772. return {"event_id": request.id}
  773. executor = _executor(
  774. [
  775. _definition(
  776. "example.terminate",
  777. blocked_terminal,
  778. terminal=True,
  779. timeout_seconds=0.01,
  780. )
  781. ],
  782. terminal_grace_seconds=0.2,
  783. )
  784. batch = await executor.execute(
  785. [_request("terminal", "example.terminate")],
  786. scope="terminal-timeout",
  787. )
  788. assert batch.results[0].status is EventStatus.TIMEOUT
  789. assert batch.results[0].error == "event timed out after 0.01 seconds"
  790. assert batch.deadline_exceeded is False
  791. @pytest.mark.asyncio
  792. async def test_replayed_terminal_grace_timeout_preserves_deadline_flag():
  793. async def blocked_terminal(request: EventRequest) -> dict[str, Any]:
  794. await asyncio.Event().wait()
  795. return {"event_id": request.id}
  796. executor = _executor(
  797. [_definition("example.terminate", blocked_terminal, terminal=True)],
  798. terminal_grace_seconds=0.01,
  799. )
  800. request = _request("terminal", "example.terminate")
  801. first = await executor.execute([request], scope="terminal-replay")
  802. replayed = await executor.execute([request], scope="terminal-replay")
  803. assert first.results[0].status is EventStatus.TIMEOUT
  804. assert first.results[0].error == "batch deadline exceeded"
  805. assert first.deadline_exceeded is False
  806. assert replayed.replayed_count == 1
  807. assert replayed.deadline_exceeded is False
  808. @pytest.mark.asyncio
  809. async def test_replayed_sibling_batch_timeout_preserves_deadline_flag():
  810. async def blocked(request: EventRequest) -> dict[str, Any]:
  811. await asyncio.Event().wait()
  812. return {"event_id": request.id}
  813. executor = _executor(
  814. [_definition("example.blocked", blocked)],
  815. batch_timeout_seconds=0.01,
  816. )
  817. request = _request("sibling", "example.blocked")
  818. first = await executor.execute([request], scope="sibling-replay")
  819. replayed = await executor.execute([request], scope="sibling-replay")
  820. assert first.results[0].status is EventStatus.TIMEOUT
  821. assert first.deadline_exceeded is True
  822. assert replayed.replayed_count == 1
  823. assert replayed.deadline_exceeded is True
  824. @pytest.mark.asyncio
  825. async def test_blocking_sync_handler_timeout_returns_before_thread_finishes():
  826. started = threading.Event()
  827. release = threading.Event()
  828. finished = threading.Event()
  829. def blocking_handler(request: EventRequest) -> dict[str, Any]:
  830. started.set()
  831. release.wait(timeout=1)
  832. finished.set()
  833. return {"event_id": request.id}
  834. executor = _executor(
  835. [
  836. _definition(
  837. "example.blocking",
  838. blocking_handler,
  839. timeout_seconds=0.02,
  840. )
  841. ],
  842. batch_timeout_seconds=0.2,
  843. )
  844. loop = asyncio.get_running_loop()
  845. started_at = loop.time()
  846. batch = await executor.execute(
  847. [_request("blocking", "example.blocking")],
  848. scope="blocking-handler",
  849. )
  850. elapsed = loop.time() - started_at
  851. release.set()
  852. await asyncio.to_thread(finished.wait, 0.2)
  853. assert started.is_set()
  854. assert batch.results[0].status is EventStatus.TIMEOUT
  855. assert elapsed < 0.1
  856. assert finished.is_set()
  857. @pytest.mark.asyncio
  858. async def test_timed_out_sync_handler_keeps_conflict_lock_until_thread_finishes():
  859. worker_started = threading.Event()
  860. worker_release = threading.Event()
  861. follower_started = asyncio.Event()
  862. def blocking_handler(request: EventRequest) -> dict[str, Any]:
  863. worker_started.set()
  864. worker_release.wait(timeout=1)
  865. return {"event_id": request.id}
  866. async def follower_handler(request: EventRequest) -> dict[str, Any]:
  867. follower_started.set()
  868. return {"event_id": request.id}
  869. executor = _executor(
  870. [
  871. _definition(
  872. "example.blocking",
  873. blocking_handler,
  874. conflict_keys=("shared",),
  875. timeout_seconds=0.02,
  876. ),
  877. _definition(
  878. "example.follower",
  879. follower_handler,
  880. conflict_keys=("shared",),
  881. timeout_seconds=0.2,
  882. ),
  883. ]
  884. )
  885. timed_out = await executor.execute(
  886. [_request("blocking", "example.blocking")],
  887. scope="sync-conflict-blocking",
  888. )
  889. follower = asyncio.create_task(
  890. executor.execute(
  891. [_request("follower", "example.follower")],
  892. scope="sync-conflict-follower",
  893. )
  894. )
  895. await asyncio.sleep(0.03)
  896. entered_before_worker_finished = follower_started.is_set()
  897. worker_release.set()
  898. follower_batch = await asyncio.wait_for(follower, timeout=0.3)
  899. assert worker_started.is_set()
  900. assert timed_out.results[0].status is EventStatus.TIMEOUT
  901. assert entered_before_worker_finished is False
  902. assert follower_batch.results[0].status is EventStatus.SUCCESS
  903. @pytest.mark.asyncio
  904. async def test_timed_out_sync_handler_keeps_parallel_slot_until_thread_finishes():
  905. worker_started = threading.Event()
  906. worker_release = threading.Event()
  907. follower_started = asyncio.Event()
  908. def blocking_handler(request: EventRequest) -> dict[str, Any]:
  909. worker_started.set()
  910. worker_release.wait(timeout=1)
  911. return {"event_id": request.id}
  912. async def follower_handler(request: EventRequest) -> dict[str, Any]:
  913. follower_started.set()
  914. return {"event_id": request.id}
  915. executor = _executor(
  916. [
  917. _definition(
  918. "example.blocking",
  919. blocking_handler,
  920. timeout_seconds=0.02,
  921. ),
  922. _definition(
  923. "example.follower",
  924. follower_handler,
  925. timeout_seconds=0.2,
  926. ),
  927. ],
  928. max_parallel_events=1,
  929. )
  930. timed_out = await executor.execute(
  931. [_request("blocking", "example.blocking")],
  932. scope="sync-slot-blocking",
  933. )
  934. follower = asyncio.create_task(
  935. executor.execute(
  936. [_request("follower", "example.follower")],
  937. scope="sync-slot-follower",
  938. )
  939. )
  940. await asyncio.sleep(0.03)
  941. entered_before_worker_finished = follower_started.is_set()
  942. worker_release.set()
  943. follower_batch = await asyncio.wait_for(follower, timeout=0.3)
  944. assert worker_started.is_set()
  945. assert timed_out.results[0].status is EventStatus.TIMEOUT
  946. assert entered_before_worker_finished is False
  947. assert follower_batch.results[0].status is EventStatus.SUCCESS
  948. @pytest.mark.asyncio
  949. async def test_timed_out_async_handler_releases_resources_after_cancellation():
  950. cancelled = asyncio.Event()
  951. follower_started = asyncio.Event()
  952. async def blocked_handler(request: EventRequest) -> dict[str, Any]:
  953. try:
  954. await asyncio.Event().wait()
  955. finally:
  956. cancelled.set()
  957. return {"event_id": request.id}
  958. async def follower_handler(request: EventRequest) -> dict[str, Any]:
  959. follower_started.set()
  960. return {"event_id": request.id}
  961. executor = _executor(
  962. [
  963. _definition(
  964. "example.blocked",
  965. blocked_handler,
  966. conflict_keys=("shared",),
  967. timeout_seconds=0.01,
  968. ),
  969. _definition(
  970. "example.follower",
  971. follower_handler,
  972. conflict_keys=("shared",),
  973. timeout_seconds=0.2,
  974. ),
  975. ],
  976. max_parallel_events=1,
  977. )
  978. timed_out = await executor.execute(
  979. [_request("blocked", "example.blocked")],
  980. scope="async-release-blocked",
  981. )
  982. follower = await executor.execute(
  983. [_request("follower", "example.follower")],
  984. scope="async-release-follower",
  985. )
  986. assert timed_out.results[0].status is EventStatus.TIMEOUT
  987. assert cancelled.is_set()
  988. assert follower_started.is_set()
  989. assert follower.results[0].status is EventStatus.SUCCESS