From e16c07ce3b8f87e0f2a32c6410f41d3f4ce1c6bb Mon Sep 17 00:00:00 2001 From: Remy Tuyeras Date: Thu, 2 Apr 2026 16:17:22 -0400 Subject: [PATCH 1/4] new feature: add in Event ( in decorators) --- .gitignore | 1 + .../question_spell_1_dna.json | 3 + .../question_spell_2_dna.json | 3 + .../question_spell_dna.json | 6 + .../question_spell_merge_1_dna.json | 6 + .../question_spell_merge_2_dna.json | 18 +- .../question_spell_merge_3_dna.json | 6 + simulations/simul_flow_2.py | 4 +- summoner/_version.py | 2 +- summoner/client/client.py | 93 +++-- summoner/client/merger.py | 96 +++-- summoner/protocol/payload.py | 17 +- summoner/protocol/process.py | 22 +- summoner/protocol/triggers.py | 10 +- summoner/protocol/validation.py | 11 +- tests/helpers.py | 10 + tests/test_client_config.py | 96 +++++ tests/test_client_context.py | 105 ++++++ tests/test_client_integration.py | 210 +++++++++++ tests/test_client_runtime.py | 146 ++++++++ tests/test_client_send_data.py | 345 ++++++++++++++++++ tests/test_merger_context.py | 110 ++++++ tests/test_payload.py | 56 +++ tests/test_process.py | 13 +- tests/test_triggers.py | 15 +- 25 files changed, 1326 insertions(+), 78 deletions(-) create mode 100644 tests/helpers.py create mode 100644 tests/test_client_config.py create mode 100644 tests/test_client_context.py create mode 100644 tests/test_client_integration.py create mode 100644 tests/test_client_runtime.py create mode 100644 tests/test_client_send_data.py create mode 100644 tests/test_merger_context.py create mode 100644 tests/test_payload.py diff --git a/.gitignore b/.gitignore index 282b4d1..68c3eb2 100644 --- a/.gitignore +++ b/.gitignore @@ -6,6 +6,7 @@ experiments/ *.mov *.db *.old +*.backup.* logs/ img/assets/ simulations/summoner_mock.py diff --git a/examples/4_spell_dna_agents/question_spell_1_dna.json b/examples/4_spell_dna_agents/question_spell_1_dna.json index 0e6b778..cb5dacd 100644 --- a/examples/4_spell_dna_agents/question_spell_1_dna.json +++ b/examples/4_spell_dna_agents/question_spell_1_dna.json @@ -33,6 +33,7 @@ { "type": "receive", "route": "spell --> effect", + "route_key": "spell--[]-->effect", "priority": [], "source": "@agent.receive(route=\"spell --> effect\")\nasync def custom_receive(msg: Union[dict,str]) -> None:\n global is_back\n msg = (msg[\"content\"] if isinstance(msg, dict) and \"content\" in msg else msg) \n tag = (\"\\r[From server]\" if isinstance(msg, str) and msg[:len(\"Warning:\")] == \"Warning:\" else \"\\r[Received]\")\n print(tag, msg, flush=True)\n if msg == \"/travel\":\n await agent.travel_to(host = \"testnet.summoner.org\", port = 8888)\n is_back = True\n return Move(Trigger.ok)\n elif msg == \"/quit\":\n await agent.quit()\n return None\n elif msg == \"/go_home\":\n await agent.travel_to(host = agent.default_host, port = agent.default_port)\n is_back = True\n return None\n print(\"waiting for instructions... \", flush=True)\n", "module": "__main__", @@ -41,7 +42,9 @@ { "type": "send", "route": "spell", + "route_key": "spell", "multi": false, + "use_data": false, "on_triggers": [], "on_actions": [], "source": "@agent.send(route=\"spell\")\nasync def send_question() -> str:\n global is_back\n if is_back and agent.host == agent.default_host and agent.port == agent.default_port:\n is_back = False\n return \"I am back\"\n else:\n await asyncio.sleep(0.1)\n", diff --git a/examples/4_spell_dna_agents/question_spell_2_dna.json b/examples/4_spell_dna_agents/question_spell_2_dna.json index 6b0d2ce..b7bc4b8 100644 --- a/examples/4_spell_dna_agents/question_spell_2_dna.json +++ b/examples/4_spell_dna_agents/question_spell_2_dna.json @@ -41,6 +41,7 @@ { "type": "receive", "route": "effect --> spell", + "route_key": "effect--[]-->spell", "priority": [], "source": "@agent.receive(route=\"effect --> spell\")\nasync def receive_response(msg: Union[str, dict]) -> None:\n global tracker\n content = msg[\"content\"] if isinstance(msg, dict) else msg\n print(f\"Received[{tracker}]: {content}\")\n if content != \"waiting\":\n async with tracker_lock:\n if tracker >= 6:\n await agent.travel_to(host = agent.default_host, port = agent.default_port)\n tracker = 0\n return Move(Trigger.ok)\n else:\n tracker += 1\n", "module": "__main__", @@ -49,7 +50,9 @@ { "type": "send", "route": "effect", + "route_key": "effect", "multi": false, + "use_data": false, "on_triggers": [], "on_actions": [], "source": "@agent.send(route=\"effect\")\nasync def send_question() -> str:\n global tracker\n await asyncio.sleep(2)\n async with tracker_lock:\n return QUESTIONS[tracker % len(QUESTIONS)]\n", diff --git a/examples/4_spell_dna_agents/question_spell_dna.json b/examples/4_spell_dna_agents/question_spell_dna.json index 523940d..ef91c9e 100644 --- a/examples/4_spell_dna_agents/question_spell_dna.json +++ b/examples/4_spell_dna_agents/question_spell_dna.json @@ -42,6 +42,7 @@ { "type": "receive", "route": "spell --> effect", + "route_key": "spell--[]-->effect", "priority": [], "source": "@agent.receive(route=\"spell --> effect\")\nasync def custom_receive(msg: Union[dict,str]) -> None:\n msg = (msg[\"content\"] if isinstance(msg, dict) and \"content\" in msg else msg) \n tag = (\"\\r[From server]\" if isinstance(msg, str) and msg[:len(\"Warning:\")] == \"Warning:\" else \"\\r[Received]\")\n print(tag, msg, flush=True)\n if msg == \"/travel\":\n await agent.travel_to(host = \"testnet.summoner.org\", port = 8888)\n return Move(Trigger.ok)\n elif msg == \"/quit\":\n await agent.quit()\n return None\n elif msg == \"/go_home\":\n await agent.travel_to(host = agent.default_host, port = agent.default_port)\n return None\n print(\"waiting for instructions... \", flush=True)\n", "module": "__main__", @@ -50,6 +51,7 @@ { "type": "receive", "route": "effect --> spell", + "route_key": "effect--[]-->spell", "priority": [], "source": "@agent.receive(route=\"effect --> spell\")\nasync def receive_response(msg: Union[str, dict]) -> None:\n global tracker, is_back\n content = msg[\"content\"] if isinstance(msg, dict) else msg\n print(f\"Received[{tracker}]: {content}\")\n if content != \"waiting\":\n async with tracker_lock:\n if tracker >= 10:\n await agent.travel_to(host = agent.default_host, port = agent.default_port)\n tracker = 0\n is_back = True\n return Move(Trigger.ok)\n else:\n tracker += 1\n", "module": "__main__", @@ -58,7 +60,9 @@ { "type": "send", "route": "effect", + "route_key": "effect", "multi": false, + "use_data": false, "on_triggers": [], "on_actions": [], "source": "@agent.send(route=\"effect\")\nasync def send_question() -> str:\n global tracker\n await asyncio.sleep(2)\n async with tracker_lock:\n return QUESTIONS[tracker % len(QUESTIONS)]\n", @@ -68,7 +72,9 @@ { "type": "send", "route": "spell", + "route_key": "spell", "multi": false, + "use_data": false, "on_triggers": [], "on_actions": [], "source": "@agent.send(route=\"spell\")\nasync def send_question() -> str:\n global is_back\n if is_back:\n is_back = False\n return \"I am back\"\n else:\n await asyncio.sleep(0.1)\n", diff --git a/examples/4_spell_dna_agents/question_spell_merge_1_dna.json b/examples/4_spell_dna_agents/question_spell_merge_1_dna.json index 80e6c4e..138659c 100644 --- a/examples/4_spell_dna_agents/question_spell_merge_1_dna.json +++ b/examples/4_spell_dna_agents/question_spell_merge_1_dna.json @@ -42,6 +42,7 @@ { "type": "receive", "route": "spell --> effect", + "route_key": "spell--[]-->effect", "priority": [], "source": "@agent.receive(route=\"spell --> effect\")\nasync def custom_receive(msg: Union[dict,str]) -> None:\n global is_back\n msg = (msg[\"content\"] if isinstance(msg, dict) and \"content\" in msg else msg) \n tag = (\"\\r[From server]\" if isinstance(msg, str) and msg[:len(\"Warning:\")] == \"Warning:\" else \"\\r[Received]\")\n print(tag, msg, flush=True)\n if msg == \"/travel\":\n await agent.travel_to(host = \"testnet.summoner.org\", port = 8888)\n is_back = True\n return Move(Trigger.ok)\n elif msg == \"/quit\":\n await agent.quit()\n return None\n elif msg == \"/go_home\":\n await agent.travel_to(host = agent.default_host, port = agent.default_port)\n is_back = True\n return None\n print(\"waiting for instructions... \", flush=True)\n", "module": "question_spell_1", @@ -50,6 +51,7 @@ { "type": "receive", "route": "effect --> spell", + "route_key": "effect--[]-->spell", "priority": [], "source": "@agent.receive(route=\"effect --> spell\")\nasync def receive_response(msg: Union[str, dict]) -> None:\n global tracker\n content = msg[\"content\"] if isinstance(msg, dict) else msg\n print(f\"Received[{tracker}]: {content}\")\n if content != \"waiting\":\n async with tracker_lock:\n if tracker >= 6:\n await agent.travel_to(host = agent.default_host, port = agent.default_port)\n tracker = 0\n return Move(Trigger.ok)\n else:\n tracker += 1\n", "module": "question_spell_2", @@ -58,7 +60,9 @@ { "type": "send", "route": "spell", + "route_key": "spell", "multi": false, + "use_data": false, "on_triggers": [], "on_actions": [], "source": "@agent.send(route=\"spell\")\nasync def send_question() -> str:\n global is_back\n if is_back and agent.host == agent.default_host and agent.port == agent.default_port:\n is_back = False\n return \"I am back\"\n else:\n await asyncio.sleep(0.1)\n", @@ -68,7 +72,9 @@ { "type": "send", "route": "effect", + "route_key": "effect", "multi": false, + "use_data": false, "on_triggers": [], "on_actions": [], "source": "@agent.send(route=\"effect\")\nasync def send_question() -> str:\n global tracker\n await asyncio.sleep(2)\n async with tracker_lock:\n return QUESTIONS[tracker % len(QUESTIONS)]\n", diff --git a/examples/4_spell_dna_agents/question_spell_merge_2_dna.json b/examples/4_spell_dna_agents/question_spell_merge_2_dna.json index 3104f56..03bcb26 100644 --- a/examples/4_spell_dna_agents/question_spell_merge_2_dna.json +++ b/examples/4_spell_dna_agents/question_spell_merge_2_dna.json @@ -27,49 +27,55 @@ { "type": "upload_states", "source": "@agent.upload_states()\nasync def upload(msg):\n global state\n print(\"[upload]\", state)\n return state\n", - "module": "summoner_merge_3bf060a8e6894bc8adb7c0a8fe929ee2", + "module": "summoner_merge_bd13e49211184e05bcdae8db224429e1", "fn_name": "upload" }, { "type": "download_states", "source": "@agent.download_states()\nasync def download(possible_state):\n global state\n print(\"[download]\", possible_state, state)\n if set(possible_state) == {Node(\"spell\")}:\n state = \"spell\"\n if set(possible_state) == {Node(\"effect\")}:\n state = \"effect\"\n print(\"[download] -->\", state)\n", - "module": "summoner_merge_3bf060a8e6894bc8adb7c0a8fe929ee2", + "module": "summoner_merge_bd13e49211184e05bcdae8db224429e1", "fn_name": "download" }, { "type": "receive", "route": "spell --> effect", + "route_key": "spell--[]-->effect", "priority": [], "source": "@agent.receive(route=\"spell --> effect\")\nasync def custom_receive(msg: Union[dict,str]) -> None:\n global is_back\n msg = (msg[\"content\"] if isinstance(msg, dict) and \"content\" in msg else msg) \n tag = (\"\\r[From server]\" if isinstance(msg, str) and msg[:len(\"Warning:\")] == \"Warning:\" else \"\\r[Received]\")\n print(tag, msg, flush=True)\n if msg == \"/travel\":\n await agent.travel_to(host = \"testnet.summoner.org\", port = 8888)\n is_back = True\n return Move(Trigger.ok)\n elif msg == \"/quit\":\n await agent.quit()\n return None\n elif msg == \"/go_home\":\n await agent.travel_to(host = agent.default_host, port = agent.default_port)\n is_back = True\n return None\n print(\"waiting for instructions... \", flush=True)\n", - "module": "summoner_merge_3e720ddabff84bec8b0af2d00526d59d", + "module": "summoner_merge_193541db02964a57969ae895d1a1c9e2", "fn_name": "custom_receive" }, { "type": "receive", "route": "effect --> spell", + "route_key": "effect--[]-->spell", "priority": [], "source": "@agent.receive(route=\"effect --> spell\")\nasync def receive_response(msg: Union[str, dict]) -> None:\n global tracker\n content = msg[\"content\"] if isinstance(msg, dict) else msg\n print(f\"Received[{tracker}]: {content}\")\n if content != \"waiting\":\n async with tracker_lock:\n if tracker >= 6:\n await agent.travel_to(host = agent.default_host, port = agent.default_port)\n tracker = 0\n return Move(Trigger.ok)\n else:\n tracker += 1\n", - "module": "summoner_merge_3bf060a8e6894bc8adb7c0a8fe929ee2", + "module": "summoner_merge_bd13e49211184e05bcdae8db224429e1", "fn_name": "receive_response" }, { "type": "send", "route": "spell", + "route_key": "spell", "multi": false, + "use_data": false, "on_triggers": [], "on_actions": [], "source": "@agent.send(route=\"spell\")\nasync def send_question() -> str:\n global is_back\n if is_back and agent.host == agent.default_host and agent.port == agent.default_port:\n is_back = False\n return \"I am back\"\n else:\n await asyncio.sleep(0.1)\n", - "module": "summoner_merge_3e720ddabff84bec8b0af2d00526d59d", + "module": "summoner_merge_193541db02964a57969ae895d1a1c9e2", "fn_name": "send_question" }, { "type": "send", "route": "effect", + "route_key": "effect", "multi": false, + "use_data": false, "on_triggers": [], "on_actions": [], "source": "@agent.send(route=\"effect\")\nasync def send_question() -> str:\n global tracker\n await asyncio.sleep(2)\n async with tracker_lock:\n return QUESTIONS[tracker % len(QUESTIONS)]\n", - "module": "summoner_merge_3bf060a8e6894bc8adb7c0a8fe929ee2", + "module": "summoner_merge_bd13e49211184e05bcdae8db224429e1", "fn_name": "send_question" } ] \ No newline at end of file diff --git a/examples/4_spell_dna_agents/question_spell_merge_3_dna.json b/examples/4_spell_dna_agents/question_spell_merge_3_dna.json index 80e6c4e..138659c 100644 --- a/examples/4_spell_dna_agents/question_spell_merge_3_dna.json +++ b/examples/4_spell_dna_agents/question_spell_merge_3_dna.json @@ -42,6 +42,7 @@ { "type": "receive", "route": "spell --> effect", + "route_key": "spell--[]-->effect", "priority": [], "source": "@agent.receive(route=\"spell --> effect\")\nasync def custom_receive(msg: Union[dict,str]) -> None:\n global is_back\n msg = (msg[\"content\"] if isinstance(msg, dict) and \"content\" in msg else msg) \n tag = (\"\\r[From server]\" if isinstance(msg, str) and msg[:len(\"Warning:\")] == \"Warning:\" else \"\\r[Received]\")\n print(tag, msg, flush=True)\n if msg == \"/travel\":\n await agent.travel_to(host = \"testnet.summoner.org\", port = 8888)\n is_back = True\n return Move(Trigger.ok)\n elif msg == \"/quit\":\n await agent.quit()\n return None\n elif msg == \"/go_home\":\n await agent.travel_to(host = agent.default_host, port = agent.default_port)\n is_back = True\n return None\n print(\"waiting for instructions... \", flush=True)\n", "module": "question_spell_1", @@ -50,6 +51,7 @@ { "type": "receive", "route": "effect --> spell", + "route_key": "effect--[]-->spell", "priority": [], "source": "@agent.receive(route=\"effect --> spell\")\nasync def receive_response(msg: Union[str, dict]) -> None:\n global tracker\n content = msg[\"content\"] if isinstance(msg, dict) else msg\n print(f\"Received[{tracker}]: {content}\")\n if content != \"waiting\":\n async with tracker_lock:\n if tracker >= 6:\n await agent.travel_to(host = agent.default_host, port = agent.default_port)\n tracker = 0\n return Move(Trigger.ok)\n else:\n tracker += 1\n", "module": "question_spell_2", @@ -58,7 +60,9 @@ { "type": "send", "route": "spell", + "route_key": "spell", "multi": false, + "use_data": false, "on_triggers": [], "on_actions": [], "source": "@agent.send(route=\"spell\")\nasync def send_question() -> str:\n global is_back\n if is_back and agent.host == agent.default_host and agent.port == agent.default_port:\n is_back = False\n return \"I am back\"\n else:\n await asyncio.sleep(0.1)\n", @@ -68,7 +72,9 @@ { "type": "send", "route": "effect", + "route_key": "effect", "multi": false, + "use_data": false, "on_triggers": [], "on_actions": [], "source": "@agent.send(route=\"effect\")\nasync def send_question() -> str:\n global tracker\n await asyncio.sleep(2)\n async with tracker_lock:\n return QUESTIONS[tracker % len(QUESTIONS)]\n", diff --git a/simulations/simul_flow_2.py b/simulations/simul_flow_2.py index 5da4bb3..1a42ebf 100644 --- a/simulations/simul_flow_2.py +++ b/simulations/simul_flow_2.py @@ -238,7 +238,8 @@ async def fn(): fn=generate_sender_fn(route_str, actions, triggers), multi=False, actions=actions, - triggers=triggers + triggers=triggers, + use_data=False, ) sender_index.setdefault(route_str, []) sender_index[route_str].append(sender) @@ -342,4 +343,3 @@ async def sender_test(): # run the simulation asyncio.run(sender_test()) - diff --git a/summoner/_version.py b/summoner/_version.py index a82b376..c68196d 100644 --- a/summoner/_version.py +++ b/summoner/_version.py @@ -1 +1 @@ -__version__ = "1.1.1" +__version__ = "1.2.0" diff --git a/summoner/client/client.py b/summoner/client/client.py index 6cfd4ea..b2fa656 100644 --- a/summoner/client/client.py +++ b/summoner/client/client.py @@ -387,6 +387,8 @@ def receive( route: str, priority: Union[int, tuple[int, ...]] = () ): + if not isinstance(route, str): + raise TypeError(f"Argument `route` must be string. Provided: {route}") route = route.strip() def decorator(fn: Callable[[Union[str, dict]], Awaitable[Optional[Event]]]): @@ -407,9 +409,6 @@ def decorator(fn: Callable[[Union[str, dict]], Awaitable[Optional[Event]]]): logger=self.logger, ) - if not isinstance(route, str): - raise TypeError(f"Argument `route` must be string. Provided: {route}") - if isinstance(priority, int): tuple_priority = (priority,) elif isinstance(priority, tuple) and all(isinstance(p, int) for p in priority): @@ -467,41 +466,53 @@ def send( multi: bool = False, on_triggers: Optional[set[Signal]] = None, on_actions: Optional[set[Type]] = None, + use_data: bool = False, ): + if not isinstance(route, str): + raise TypeError(f"Argument `route` must be string. Provided: {route}") route = route.strip() - def decorator(fn: Callable[[], Awaitable]): + def decorator(fn: Callable[..., Awaitable]): # ----[ Safety Checks ]---- if not inspect.iscoroutinefunction(fn): raise TypeError(f"@send sender '{fn.__name__}' must be async") - sig = inspect.signature(fn) - if len(sig.parameters) != 0: - raise TypeError(f"@send '{fn.__name__}' must accept no arguments") - + expected_params = 1 if use_data else 0 + decorator_name = "@send" + if multi and use_data: + decorator_name = "@send[multi=True,use_data=True]" + elif multi: + decorator_name = "@send[multi=True]" + elif use_data: + decorator_name = "@send[use_data=True]" + if not multi: _check_param_and_return( fn, - decorator_name="@send", - allow_param=(), # no args allowed + decorator_name=decorator_name, + allow_param=(Any,) if use_data else (), # one data arg when enabled allow_return=(type(None), Any, str, dict), logger=self.logger, + expected_params=expected_params, + skip_param_type_check=use_data, ) else: _check_param_and_return( fn, - decorator_name="@send[multi=True]", - allow_param=(), # no args allowed + decorator_name=decorator_name, + allow_param=(Any,) if use_data else (), # one data arg when enabled allow_return=(Any, list, list[str], list[dict], list[Union[str, dict]]), logger=self.logger, + expected_params=expected_params, + skip_param_type_check=use_data, ) - if not isinstance(route, str): - raise TypeError(f"Argument `route` must be string. Provided: {route}") - if not isinstance(multi, bool): raise TypeError(f"Argument `multi` must be Boolean. Provided: {multi}") + if not isinstance(use_data, bool): + raise TypeError(f"Argument `use_data` must be Boolean. Provided: {use_data}") + if on_triggers is not None and ( not isinstance(on_triggers, set) or not all(isinstance(sig, Signal) for sig in on_triggers) @@ -513,6 +524,17 @@ def decorator(fn: Callable[[], Awaitable]): not all(isinstance(act, type) and issubclass(act, Event) and act in {Action.MOVE, Action.STAY, Action.TEST} for act in on_actions) ): raise TypeError(f"Argument `on_actions` must be `None` or a set of Action event classes: {{Action.MOVE, Action.STAY, Action.TEST}}. Provided: {on_actions!r}") + + reactive_filters = ( + (isinstance(on_triggers, set) and bool(on_triggers)) or + (isinstance(on_actions, set) and bool(on_actions)) + ) + + if use_data and not reactive_filters: + raise ValueError("Argument `use_data=True` requires a non-empty `on_triggers` or `on_actions` set so the sender has queued event data to receive") + + if use_data and not self._flow.in_use: + raise RuntimeError("Argument `use_data=True` requires `client.flow().activate()` before sender registration") # ----[ DNA capture ]---- self._dna_senders.append({ @@ -521,13 +543,14 @@ def decorator(fn: Callable[[], Awaitable]): "multi": multi, "on_triggers": on_triggers, "on_actions": on_actions, + "use_data": use_data, "source": inspect.getsource(fn), }) # ----[ Registration Code ]---- async def register(): - sender = Sender(fn=fn, multi=multi, actions=on_actions, triggers=on_triggers) + sender = Sender(fn=fn, multi=multi, actions=on_actions, triggers=on_triggers, use_data=use_data) actions_exist = isinstance(on_actions, set) and bool(on_actions) triggers_exist = isinstance(on_triggers, set) and bool(on_triggers) @@ -759,18 +782,20 @@ def dna(self, include_context: bool = False) -> str: route_key = raw_route route_key = "".join(str(route_key).split()) - entries.append({ + entry = { "type": "send", "route": raw_route, # original route string "route_key": route_key, # stable route representative "multi": dna["multi"], + "use_data": dna["use_data"], # Serialize triggers/actions by name so they can be re-resolved later. "on_triggers": [t.name for t in (dna["on_triggers"] or [])], "on_actions": [a.__name__ for a in (dna["on_actions"] or [])], "source": get_callable_source(fn, dna.get("source")), "module": fn.__module__, "fn_name": fn.__name__, - }) + } + entries.append(entry) # All hooks for dna in self._dna_hooks: @@ -1128,14 +1153,17 @@ async def _send_worker( while True: - item: Optional[tuple[str, Sender]] = await self.send_queue.get() + item: Optional[tuple[str, Sender, Any]] = await self.send_queue.get() if item is None: self.send_queue.task_done() break - route, sender = item + route, sender, sender_data = item try: - result = await sender.fn() + if sender.use_data: + result = await sender.fn(sender_data) + else: + result = await sender.fn() # ----[ Urgent: Handle Aborts ]---- async with self.connection_lock: @@ -1263,9 +1291,10 @@ def _route_accepts( pending.sort(key=lambda it: hook_priority_order(it[0])) # ----[ Build Sender Batch ]---- - senders: list[tuple[str, Sender]] = [] + senders: list[tuple[str, Sender, Any]] = [] # De-dup set: at most one sender per (route, key-from-recv, recv-handler-name) this cycle. + # `use_data=True` senders intentionally bypass this so each queued event payload is delivered. emitted: set[tuple[str, Optional[str], str]] = set() for route, routed_senders in sender_index.items(): @@ -1273,7 +1302,13 @@ def _route_accepts( # Non-reactive (no actions/triggers): preserve current behavior if (not self._flow.in_use) or (sender.actions is None and sender.triggers is None): - senders.append((route, sender)) + if sender.use_data: + self.logger.warning( + f"Skipping sender '{sender.fn.__name__}' on route '{route}': " + "use_data=True requires a reactive sender running with flow enabled" + ) + continue + senders.append((route, sender, None)) # Reactive: require matching a pending activation (existential) elif self._flow.in_use and ((sender.actions and isinstance(sender.actions, set)) or @@ -1283,12 +1318,18 @@ def _route_accepts( if sender_parsed_route is None: continue - # Iterate pending in queue order; first match "wins" for this (route,key,fn_name) - for (priority, key, parsed_route, event) in pending: + # Iterate pending in queue order. + # `use_data=True` senders receive one call per matching queued event. + # Other reactive senders keep the existing first-match de-dup behavior. + for (_priority, key, parsed_route, event) in pending: if _route_accepts(sender_parsed_route, parsed_route) and sender.responds_to(event): + if sender.use_data: + senders.append((route, sender, event.data)) + continue + dedup_key = (route, key, sender.fn.__name__) # key scopes to the activation thread/peer if dedup_key not in emitted: - senders.append((route, sender)) + senders.append((route, sender, None)) emitted.add(dedup_key) break # do not enqueue multiple times for this sender this cycle diff --git a/summoner/client/merger.py b/summoner/client/merger.py index 3a3bad4..b03bd12 100644 --- a/summoner/client/merger.py +++ b/summoner/client/merger.py @@ -58,6 +58,7 @@ import re import json import uuid +import textwrap import os, sys target_path = os.path.abspath(os.path.join(os.path.dirname(os.path.abspath(__file__)), "..")) @@ -151,6 +152,20 @@ def _resolve_action(ActionCls, name: str): raise KeyError(f"Unknown action '{name}' for {ActionCls}") +def _read_required_mapping_field(entry: dict[str, Any], field: str): + """ + Read a required DNA field and fail clearly if it is absent. + + Replay code accepts missing defaultable fields via `.get(...)`, but route and + other structural identifiers should still be explicit. This helper keeps the + resulting failure mode intentional instead of surfacing a later KeyError or + attribute error from deeper registration code. + """ + if field not in entry: + raise KeyError(f"Missing required DNA field '{field}'") + return entry[field] + + class ClientMerger(SummonerClient): """ Merge multiple sources into one client. @@ -633,7 +648,7 @@ def _make_from_source(self, entry: dict[str, Any], g: dict, sandbox_name: str) - if def_idx is None: raise RuntimeError(f"Could not find def for '{fn_name}'") - func_body = "\n".join(lines[def_idx:]) + func_body = textwrap.dedent("\n".join(lines[def_idx:])) # --------------------------------------------------------------------- # 2) Ensure rebinding happens in the same globals dict used by exec(). @@ -762,7 +777,7 @@ def initiate_receivers(self): def initiate_senders(self): """ - Replay @send(route, multi, on_triggers, on_actions) from every source onto the merged client. + Replay @send(route, multi, on_triggers, on_actions, use_data) from every source onto the merged client. Imported-client sources: - carry actual trigger/action objects in _dna_senders. @@ -774,6 +789,24 @@ def initiate_senders(self): - else fall back to load_triggers() - actions are resolved from Action by name via _resolve_action. """ + if not self._flow.in_use: + for src in self.sources: + if src["kind"] == "client": + if any(dna.get("use_data", False) for dna in src["client"]._dna_senders): + raise RuntimeError( + "ClientMerger.flow().activate() must be called before initiate_senders() " + "when replaying senders with use_data=True" + ) + else: + if any( + entry.get("type") == "send" and entry.get("use_data", False) + for entry in src["dna_entries"] + ): + raise RuntimeError( + "ClientMerger.flow().activate() must be called before initiate_senders() " + "when replaying senders with use_data=True" + ) + for src in self.sources: if src["kind"] == "client": client: SummonerClient = src["client"] @@ -781,37 +814,44 @@ def initiate_senders(self): for dna in client._dna_senders: fn_clone = self._clone_handler(dna["fn"], var_name) try: + route = _read_required_mapping_field(dna, "route") self.send( - dna["route"], - multi=dna["multi"], - on_triggers=dna["on_triggers"], - on_actions=dna["on_actions"], + route, + multi=dna.get("multi", False), + on_triggers=dna.get("on_triggers"), + on_actions=dna.get("on_actions"), + use_data=dna.get("use_data", False), )(fn_clone) except Exception as e: self.logger.warning( - f"[{var_name}] Failed to replay sender '{dna['fn'].__name__}' on route '{dna['route']}': {e}" + f"[{var_name}] Failed to replay sender '{dna['fn'].__name__}' " + f"on route '{dna.get('route', '')}': {e}" ) else: g = src["globals"] sandbox = src["sandbox_name"] - # Triggers: prefer a Trigger class provided by sandbox context; otherwise load defaults. - TriggerCls = g.get("Trigger") - if TriggerCls is None: - TriggerCls = load_triggers() - for entry in src["dna_entries"]: if entry.get("type") != "send": continue fn = self._make_from_source(entry, g, sandbox) - on_triggers = {_resolve_trigger(TriggerCls, t) for t in entry.get("on_triggers", [])} or None + route = _read_required_mapping_field(entry, "route") + trigger_names = entry.get("on_triggers", []) + on_triggers = None + if trigger_names: + TriggerCls = g.get("Trigger") + if TriggerCls is None: + TriggerCls = load_triggers() + g["Trigger"] = TriggerCls + on_triggers = {_resolve_trigger(TriggerCls, t) for t in trigger_names} or None on_actions = {_resolve_action(Action, a) for a in entry.get("on_actions", [])} or None dec = self.send( - entry["route"], + route, multi=entry.get("multi", False), on_triggers=on_triggers, on_actions=on_actions, + use_data=entry.get("use_data", False), ) self._apply_with_source_patch(dec, fn, entry["source"]) @@ -1055,7 +1095,7 @@ def _make_from_source(self, entry: dict[str, Any]) -> types.FunctionType: for idx, line in enumerate(lines): pattern = rf"\s*(async\s+)?def\s+{re.escape(fn_name)}\b" if re.match(pattern, line): - func_body = "\n".join(lines[idx:]) + func_body = textwrap.dedent("\n".join(lines[idx:])) break else: raise RuntimeError(f"Could not find definition for '{fn_name}'") @@ -1130,27 +1170,41 @@ def initiate_senders(self): - Trigger is resolved using a Trigger class found in sandbox globals, else load_triggers() - Action is resolved from the Action container by name """ + if (not self._flow.in_use) and any( + entry.get("type") == "send" and entry.get("use_data", False) + for entry in self._dna_list + ): + raise RuntimeError( + "ClientTranslation.flow().activate() must be called before initiate_senders() " + "when replaying senders with use_data=True" + ) + g = self._sandbox_globals # Ensure rebind globals are visible before resolving triggers/actions. if self._rebind_globals: g.update(self._rebind_globals) - TriggerCls = g.get("Trigger") - if TriggerCls is None: - TriggerCls = load_triggers() - for entry in self._dna_list: if entry.get("type") != "send": continue fn = self._make_from_source(entry) - on_triggers = {_resolve_trigger(TriggerCls, t) for t in entry.get("on_triggers", [])} or None + route = _read_required_mapping_field(entry, "route") + trigger_names = entry.get("on_triggers", []) + on_triggers = None + if trigger_names: + TriggerCls = g.get("Trigger") + if TriggerCls is None: + TriggerCls = load_triggers() + g["Trigger"] = TriggerCls + on_triggers = {_resolve_trigger(TriggerCls, t) for t in trigger_names} or None on_actions = {_resolve_action(Action, a) for a in entry.get("on_actions", [])} or None dec = self.send( - entry["route"], + route, multi=entry.get("multi", False), on_triggers=on_triggers, on_actions=on_actions, + use_data=entry.get("use_data", False), ) self._apply_with_source_patch(dec, fn, entry["source"]) diff --git a/summoner/protocol/payload.py b/summoner/protocol/payload.py index c6bd21e..258e1d2 100644 --- a/summoner/protocol/payload.py +++ b/summoner/protocol/payload.py @@ -153,12 +153,13 @@ def cast_v0_0_1(val: Any, expected: Any) -> Any: return val -# Register version 0.0.1 +# Register all versions since 0.0.1 register_envelope_version("0.0.1", parse_v0_0_1, cast_v0_0_1) register_envelope_version("1.0.0", parse_v0_0_1, cast_v0_0_1) register_envelope_version("1.0.1", parse_v0_0_1, cast_v0_0_1) register_envelope_version("1.1.0", parse_v0_0_1, cast_v0_0_1) -# register_envelope_version("1.1.1", parse_v0_0_1, cast_v0_0_1) +register_envelope_version("1.1.1", parse_v0_0_1, cast_v0_0_1) +# register_envelope_version("1.2.0", parse_v0_0_1, cast_v0_0_1) register_envelope_version(core_version, parse_v0_0_1, cast_v0_0_1) @@ -227,9 +228,12 @@ def recover_with_types(text: str) -> RelayedMessage: 3. Returns a dict: {"remote_addr": , "content": }. Fallbacks (in order): - - Invalid JSON or non-JSON warning strings: strip trailing newline and return raw string. - - Missing outer keys ("remote_addr"/"content"): strip newline and return raw string. - - Missing envelope keys inside content: return the parsed `{"remote_addr":…, "content":…}` as-is. + - Invalid JSON or non-JSON warning strings: strip a trailing newline and return the raw string. + - Plain JSON strings that are not transport envelopes: strip a trailing newline and return the parsed string. + - Parsed JSON values that are not Summoner relay wrappers + (missing outer "remote_addr"/"content"): return the parsed value as-is. + - Summoner relay wrappers whose `content` is not a typed envelope: + return the parsed `{"remote_addr": ..., "content": ...}` object as-is. Args: text: Raw text received from the server (may include newline). @@ -248,6 +252,9 @@ def recover_with_types(text: str) -> RelayedMessage: # If that fails (e.g. pure warning string), strip newline and relay the raw text return remove_last_newline(text) + if isinstance(obj, str): + return remove_last_newline(obj) + # 2) Ensure we have the outer {"remote_addr":…, "content":…} wrapper if not (isinstance(obj, dict) and "remote_addr" in obj and "content" in obj): # Malformed protocol message; upstream code can catch this if needed diff --git a/summoner/protocol/process.py b/summoner/protocol/process.py index c432e23..bfaa4c9 100644 --- a/summoner/protocol/process.py +++ b/summoner/protocol/process.py @@ -304,13 +304,28 @@ def activated_nodes( # ======= PROTOCOL: SEND / RECEIVE ======= -@dataclass(frozen=True) +@dataclass(frozen=True, init=False) class Sender: - __slots__ = ('fn', 'multi', 'actions', 'triggers') - fn: Callable[[], Awaitable] + __slots__ = ('fn', 'multi', 'actions', 'triggers', 'use_data') + fn: Callable[..., Awaitable] multi: bool actions: Optional[set[Type]] triggers: Optional[set[Signal]] + use_data: bool + + def __init__( + self, + fn: Callable[..., Awaitable], + multi: bool, + actions: Optional[set[Type]], + triggers: Optional[set[Signal]], + use_data: bool = False, + ): + object.__setattr__(self, "fn", fn) + object.__setattr__(self, "multi", multi) + object.__setattr__(self, "actions", actions) + object.__setattr__(self, "triggers", triggers) + object.__setattr__(self, "use_data", use_data) def responds_to(self, event: Any) -> bool: action_check = True @@ -522,4 +537,3 @@ class ClientIntent(Enum): QUIT = auto() # brutal, immediate exit TRAVEL = auto() # switch to a new host/port ABORT = auto() # abort due to error - diff --git a/summoner/protocol/triggers.py b/summoner/protocol/triggers.py index d2f55f5..b31a726 100644 --- a/summoner/protocol/triggers.py +++ b/summoner/protocol/triggers.py @@ -222,11 +222,14 @@ def name_of(*args): class Event: - __slots__ = ("signal",) - def __init__(self, signal: Signal) -> None: + __slots__ = ("signal", "data") + def __init__(self, signal: Signal, data: Any = None) -> None: self.signal = signal + self.data = data def __repr__(self) -> str: - return f"{type(self).__name__}({self.signal!r})" + if self.data is None: + return f"{type(self).__name__}({self.signal!r})" + return f"{type(self).__name__}({self.signal!r}, data={self.data!r})" class Move(Event): pass @@ -287,4 +290,3 @@ def load_triggers( f"Could not find triggers file at {path if 'path' in locals() else ''}" ) from e return build_triggers(tree) - diff --git a/summoner/protocol/validation.py b/summoner/protocol/validation.py index f2659dc..5f1574e 100644 --- a/summoner/protocol/validation.py +++ b/summoner/protocol/validation.py @@ -51,20 +51,23 @@ def _check_param_and_return(fn, decorator_name: str, allow_param: tuple[type, ...], allow_return: tuple[type, ...], - logger: Logger): + logger: Logger, + expected_params: Union[int, None] = None, + skip_param_type_check: bool = False): sig = inspect.signature(fn) hints = get_type_hints(fn) # parameter check params = list(sig.parameters.values()) - expected_params = 1 if decorator_name in ("@hook", "@receive", "@upload_states", "@download_states") else 0 + if expected_params is None: + expected_params = 1 if decorator_name in ("@hook", "@receive", "@upload_states", "@download_states") else 0 if len(params) != expected_params: raise TypeError(f"{decorator_name} '{fn.__name__}' must have " f"{expected_params} parameter(s), not {len(params)}") - if expected_params == 1: + if expected_params == 1 and not skip_param_type_check: raw_param = params[0].annotation param_hint = _normalize_annotation(raw_param) or hints.get(params[0].name, None) if param_hint is None: @@ -87,4 +90,4 @@ def _check_param_and_return(fn, elif not _valid_type_hint(ret_hint, allow_return): raise TypeError( f"{decorator_name} '{fn.__name__}' must return one of {allow_return}, not {ret_hint!r}" - ) \ No newline at end of file + ) diff --git a/tests/helpers.py b/tests/helpers.py new file mode 100644 index 0000000..8680613 --- /dev/null +++ b/tests/helpers.py @@ -0,0 +1,10 @@ +class DummyWriter: + def __init__(self): + self.messages: list[bytes] = [] + self.drain_calls = 0 + + def write(self, data: bytes) -> None: + self.messages.append(data) + + async def drain(self) -> None: + self.drain_calls += 1 diff --git a/tests/test_client_config.py b/tests/test_client_config.py new file mode 100644 index 0000000..4538a99 --- /dev/null +++ b/tests/test_client_config.py @@ -0,0 +1,96 @@ +import pytest + +from summoner.client import SummonerClient + + +def test_apply_config_sets_values_and_defaults(): + client = SummonerClient("sdk-config") + + try: + client._apply_config( + { + "host": "127.0.0.1", + "port": 9000, + "logger": {}, + "hyper_parameters": { + "reconnection": { + "retry_delay_seconds": 7, + "primary_retry_limit": 4, + }, + "receiver": { + "max_bytes_per_line": 4096, + "read_timeout_seconds": 12, + }, + "sender": { + "concurrency_limit": 3, + "queue_maxsize": 5, + "event_bridge_maxsize": 21, + "max_worker_errors": 4, + "batch_drain": False, + }, + }, + } + ) + + assert client.host == "127.0.0.1" + assert client.port == 9000 + assert client.retry_delay_seconds == 7 + assert client.primary_retry_limit == 4 + assert client.default_host == "127.0.0.1" + assert client.default_port == 9000 + assert client.max_bytes_per_line == 4096 + assert client.read_timeout_seconds == 12 + assert client.max_concurrent_workers == 3 + assert client.send_queue_maxsize == 5 + assert client.event_bridge_maxsize == 21 + assert client.max_consecutive_worker_errors == 4 + assert client.batch_drain is False + finally: + client.loop.close() + + +def test_apply_config_rejects_invalid_sender_limits(): + client = SummonerClient("sdk-config-invalid") + + try: + with pytest.raises(ValueError): + client._apply_config( + { + "logger": {}, + "hyper_parameters": { + "sender": { + "concurrency_limit": 0, + "queue_maxsize": 1, + }, + }, + } + ) + + with pytest.raises(ValueError): + client._apply_config( + { + "logger": {}, + "hyper_parameters": { + "sender": { + "concurrency_limit": 1, + "queue_maxsize": 0, + }, + }, + } + ) + + with pytest.raises(ValueError): + client._apply_config( + { + "logger": {}, + "hyper_parameters": { + "sender": { + "concurrency_limit": 1, + "queue_maxsize": 1, + "max_worker_errors": 0, + }, + }, + } + ) + finally: + client.loop.close() diff --git a/tests/test_client_context.py b/tests/test_client_context.py new file mode 100644 index 0000000..e3f8819 --- /dev/null +++ b/tests/test_client_context.py @@ -0,0 +1,105 @@ +import json + +from summoner.client import SummonerClient +from summoner.protocol import Action +from summoner.protocol.process import Direction + + +def test_iter_registered_handler_functions_yields_all_registered_handlers(): + client = SummonerClient("client-context") + + try: + @client.upload_states() + async def upload_states(payload: dict) -> dict: + return {"payload": payload} + + @client.download_states() + async def download_states(payload: dict): + return payload + + @client.hook(Direction.SEND, priority=(1,)) + async def send_hook(payload: dict) -> dict: + return {"payload": payload} + + @client.receive(route="request", priority=(2,)) + async def receive_request(payload: dict): + return {"payload": payload} + + @client.send(route="request", on_actions={Action.STAY}) + async def send_request() -> dict: + return {"sent": True} + + client.loop.run_until_complete(client._wait_for_registration()) + + handlers = list(client._iter_registered_handler_functions()) + assert handlers == [ + upload_states, + download_states, + receive_request, + send_request, + send_hook, + ] + finally: + client.loop.close() + + +def test_infer_client_binding_name_is_exported_in_context_dna(): + client = SummonerClient("binding-name") + binding_name = "sdk_bound_client_name" + previous = globals().get(binding_name) + globals()[binding_name] = client + + try: + @client.send(route="request", on_actions={Action.STAY}) + async def send_request() -> dict: + return {"sent": True} + + client.loop.run_until_complete(client._wait_for_registration()) + + assert client._infer_client_binding_name() == binding_name + + dna_list = json.loads(client.dna(include_context=True)) + assert dna_list[0]["type"] == "__context__" + assert dna_list[0]["var_name"] == binding_name + finally: + if previous is None: + globals().pop(binding_name, None) + else: + globals()[binding_name] = previous + client.loop.close() + + +def test_reset_client_intent_clears_quit_and_travel_flags(): + client = SummonerClient("intent-reset") + + try: + client.loop.run_until_complete(client.travel_to("127.0.0.1", 9000)) + client.loop.run_until_complete(client.quit()) + + assert client._travel is True + assert client._quit is True + assert client.host == "127.0.0.1" + assert client.port == 9000 + + client.loop.run_until_complete(client._reset_client_intent()) + + assert client._travel is False + assert client._quit is False + assert client.host == "127.0.0.1" + assert client.port == 9000 + finally: + client.loop.close() + + +def test_initialize_compiles_arrow_patterns_when_flow_is_active(): + client = SummonerClient("flow-init") + + try: + client.flow().activate() + client.flow().add_arrow_style("-", ("[", "]"), ",", ">") + + assert client.flow()._regex_ready is False + client.initialize() + assert client.flow()._regex_ready is True + finally: + client.loop.close() diff --git a/tests/test_client_integration.py b/tests/test_client_integration.py new file mode 100644 index 0000000..86ca95d --- /dev/null +++ b/tests/test_client_integration.py @@ -0,0 +1,210 @@ +import json + +from summoner.client import ClientMerger, ClientTranslation, SummonerClient +from summoner.protocol import Action +from summoner.protocol.process import Direction + + +SDK_CONTEXT_VALUE = "sdk-context" + + +def test_client_registers_all_handler_types_and_exports_dna_schema(): + client = SummonerClient("sdk-setup") + + try: + @client.upload_states() + async def upload_states(payload: dict) -> dict: + return {"state": "ready", "context": SDK_CONTEXT_VALUE} + + @client.download_states() + async def download_states(payload: dict): + return {"downloaded": payload, "context": SDK_CONTEXT_VALUE} + + @client.hook(Direction.RECEIVE, priority=(1,)) + async def receive_hook(payload: dict) -> dict: + return {"hook": "receive", "payload": payload, "context": SDK_CONTEXT_VALUE} + + @client.hook(Direction.SEND, priority=(2,)) + async def send_hook(payload: dict) -> dict: + return {"hook": "send", "payload": payload, "context": SDK_CONTEXT_VALUE} + + @client.receive(route="request", priority=(3,)) + async def receive_request(payload: str): + return {"received": payload, "context": SDK_CONTEXT_VALUE} + + @client.send(route="request", multi=False, on_actions={Action.STAY}) + async def send_request() -> dict: + return {"sent": True, "context": SDK_CONTEXT_VALUE} + + client.loop.run_until_complete(client._wait_for_registration()) + + assert client._upload_states is upload_states + assert client._download_states is download_states + assert client.receiving_hooks[(1,)] is receive_hook + assert client.sending_hooks[(2,)] is send_hook + assert client.receiver_index["request"].fn is receive_request + assert client.sender_index["request"][0].fn is send_request + + dna_entries = json.loads(client.dna()) + assert [entry["type"] for entry in dna_entries] == [ + "upload_states", + "download_states", + "receive", + "send", + "hook", + "hook", + ] + + send_entry = dna_entries[3] + assert send_entry["route"] == "request" + assert send_entry["multi"] is False + assert send_entry["use_data"] is False + assert send_entry["on_actions"] == ["Stay"] + assert send_entry["on_triggers"] == [] + finally: + client.loop.close() + + +def test_client_dna_with_context_supports_translation_of_mixed_handlers(): + client = SummonerClient("sdk-context") + translated = None + + try: + @client.upload_states() + async def upload_states(payload: str) -> dict: + return {"state": payload, "context": SDK_CONTEXT_VALUE} + + @client.download_states() + async def download_states(payload: dict): + return {"downloaded": payload, "context": SDK_CONTEXT_VALUE} + + @client.hook(Direction.RECEIVE, priority=(1,)) + async def receive_hook(payload: str) -> dict: + return {"hook": "receive", "payload": payload, "context": SDK_CONTEXT_VALUE} + + @client.receive(route="request", priority=(2,)) + async def receive_request(payload: str): + return {"received": payload, "context": SDK_CONTEXT_VALUE} + + @client.send(route="request", on_actions={Action.STAY}) + async def send_request() -> dict: + return {"sent": True, "context": SDK_CONTEXT_VALUE} + + client.loop.run_until_complete(client._wait_for_registration()) + + dna_list = json.loads(client.dna(include_context=True)) + assert dna_list[0]["type"] == "__context__" + assert dna_list[0]["globals"]["SDK_CONTEXT_VALUE"] == SDK_CONTEXT_VALUE + + translated = ClientTranslation(dna_list, name="translated-context") + translated.initiate_all() + translated.loop.run_until_complete(translated._wait_for_registration()) + + assert translated.loop.run_until_complete(translated._upload_states("seed")) == { + "state": "seed", + "context": SDK_CONTEXT_VALUE, + } + assert translated.loop.run_until_complete(translated._download_states({"peer": "a"})) == { + "downloaded": {"peer": "a"}, + "context": SDK_CONTEXT_VALUE, + } + assert translated.loop.run_until_complete(translated.receiving_hooks[(1,)]("payload")) == { + "hook": "receive", + "payload": "payload", + "context": SDK_CONTEXT_VALUE, + } + assert translated.loop.run_until_complete(translated.receiver_index["request"].fn("hello")) == { + "received": "hello", + "context": SDK_CONTEXT_VALUE, + } + assert translated.loop.run_until_complete(translated.sender_index["request"][0].fn()) == { + "sent": True, + "context": SDK_CONTEXT_VALUE, + } + finally: + client.loop.close() + if translated is not None: + translated.loop.close() + + +def test_client_merger_initiate_all_combines_sources(): + client_a = SummonerClient("sdk-a") + client_b = SummonerClient("sdk-b") + merged = None + + try: + @client_a.upload_states() + async def upload_states(payload: str) -> dict: + return {"from": "a", "payload": payload} + + @client_a.receive(route="request_a", priority=(1,)) + async def receive_a(payload: str): + return {"route": "a", "payload": payload} + + @client_a.send(route="request_a", on_actions={Action.STAY}) + async def send_a() -> dict: + return {"sender": "a"} + + @client_b.download_states() + async def download_states(payload: dict): + return {"from": "b", "payload": payload} + + @client_b.hook(Direction.SEND, priority=(2,)) + async def send_hook(payload: str) -> dict: + return {"hook": "b", "payload": payload} + + @client_b.receive(route="request_b", priority=(3,)) + async def receive_b(payload: str): + return {"route": "b", "payload": payload} + + @client_b.send(route="request_b", on_actions={Action.TEST}) + async def send_b() -> dict: + return {"sender": "b"} + + client_a.loop.run_until_complete(client_a._wait_for_registration()) + client_b.loop.run_until_complete(client_b._wait_for_registration()) + + merged = ClientMerger([client_a, client_b], name="merged", close_subclients=False) + merged.initiate_all() + merged.loop.run_until_complete(merged._wait_for_registration()) + + assert merged._upload_states is not None + assert merged._download_states is not None + assert (2,) in merged.sending_hooks + assert "request_a" in merged.receiver_index + assert "request_b" in merged.receiver_index + assert "request_a" in merged.sender_index + assert "request_b" in merged.sender_index + + assert merged.loop.run_until_complete(merged._upload_states("seed")) == { + "from": "a", + "payload": "seed", + } + merged_payload = {"peer": ["request_b"]} + assert merged.loop.run_until_complete(merged._download_states(merged_payload)) == { + "from": "b", + "payload": merged_payload, + } + assert merged.loop.run_until_complete(merged.sending_hooks[(2,)]("seed")) == { + "hook": "b", + "payload": "seed", + } + assert merged.loop.run_until_complete(merged.receiver_index["request_a"].fn("hello")) == { + "route": "a", + "payload": "hello", + } + assert merged.loop.run_until_complete(merged.receiver_index["request_b"].fn("world")) == { + "route": "b", + "payload": "world", + } + assert merged.loop.run_until_complete(merged.sender_index["request_a"][0].fn()) == { + "sender": "a", + } + assert merged.loop.run_until_complete(merged.sender_index["request_b"][0].fn()) == { + "sender": "b", + } + finally: + client_a.loop.close() + client_b.loop.close() + if merged is not None: + merged.loop.close() diff --git a/tests/test_client_runtime.py b/tests/test_client_runtime.py new file mode 100644 index 0000000..d133004 --- /dev/null +++ b/tests/test_client_runtime.py @@ -0,0 +1,146 @@ +import asyncio +import json +from typing import Any + +from summoner.client import SummonerClient +from summoner.protocol import Action +from summoner.protocol.payload import wrap_with_types +from summoner.protocol.process import Direction, Node +from summoner.protocol.triggers import load_triggers +from tests.helpers import DummyWriter + + +def test_message_receiver_loop_applies_receive_hooks_updates_state_and_bridges_events(): + client = SummonerClient("runtime-receiver") + client.flow().activate() + Trigger = load_triggers(json_dict={"ok": None}) + + downloads: list[dict[str, list[Node]]] = [] + reader = asyncio.StreamReader() + stop_event = asyncio.Event() + + try: + @client.hook(Direction.RECEIVE, priority=(1,)) + async def receive_hook_one(payload: dict) -> dict: + updated = dict(payload) + updated["trace"] = [1] + return updated + + @client.hook(Direction.RECEIVE, priority=(2,)) + async def receive_hook_two(payload: dict) -> dict: + updated = dict(payload) + updated["trace"] = list(payload["trace"]) + [2] + return updated + + @client.upload_states() + async def upload_states(payload: dict) -> dict[str, str]: + return {payload["remote_addr"]: "request"} + + @client.download_states() + async def download_states(payload: dict) -> None: + downloads.append(payload) + return None + + @client.receive(route="request", priority=(5,)) + async def receive_request(payload: dict) -> Any: + await client.quit() + return Action.STAY( + Trigger.ok, + data={ + "trace": payload["trace"], + "content": payload["content"], + }, + ) + + client.loop.run_until_complete(client._wait_for_registration()) + + client.max_bytes_per_line = 4096 + client.read_timeout_seconds = 0.1 + client.event_bridge = asyncio.Queue(maxsize=8) + + raw_message = json.dumps( + { + "remote_addr": "peer-a", + "content": json.loads( + wrap_with_types({"body": "hello"}, version=client.core_version) + ), + } + ) + "\n" + + reader.feed_data(raw_message.encode()) + reader.feed_eof() + + client.loop.run_until_complete(client.message_receiver_loop(reader, stop_event)) + + assert stop_event.is_set() + assert downloads == [{"peer-a": [Node("request")]}] + + priority, key, route, event = client.event_bridge.get_nowait() + assert priority == (5,) + assert key == "tape:peer-a" + assert str(route) == "request" + assert isinstance(event, Action.STAY) + assert event.data == { + "trace": [1, 2], + "content": {"body": "hello"}, + } + finally: + client.loop.close() + + +def test_message_sender_loop_applies_send_hooks_and_writes_wrapped_payload(): + client = SummonerClient("runtime-sender") + client.flow().activate() + Trigger = load_triggers(json_dict={"go": None}) + writer = DummyWriter() + stop_event = asyncio.Event() + + try: + @client.hook(Direction.SEND, priority=(1,)) + async def send_hook_one(payload: Any) -> dict: + updated = dict(payload) + updated["trace"] = ["hook-1"] + return updated + + @client.hook(Direction.SEND, priority=(2,)) + async def send_hook_two(payload: Any) -> dict: + updated = dict(payload) + updated["trace"] = list(payload["trace"]) + ["hook-2"] + await client.quit() + return updated + + @client.send(route="request", on_actions={Action.STAY}) + async def send_request() -> dict: + return {"body": "hello"} + + client.loop.run_until_complete(client._wait_for_registration()) + + client.max_concurrent_workers = 1 + client.send_queue_maxsize = 8 + client.event_bridge_maxsize = 8 + client.max_consecutive_worker_errors = 3 + client.batch_drain = True + client.send_queue = asyncio.Queue(maxsize=client.send_queue_maxsize) + client.event_bridge = asyncio.Queue(maxsize=client.event_bridge_maxsize) + + parsed_route = client.flow().parse_route("request") + client.event_bridge.put_nowait( + ((3,), "tape:peer-a", parsed_route, Action.STAY(Trigger.go)) + ) + + client._start_send_workers(writer, stop_event) + client.loop.run_until_complete(client.message_sender_loop(writer, stop_event)) + client.loop.run_until_complete(client._cleanup_workers()) + + assert stop_event.is_set() + assert writer.drain_calls == 1 + assert len(writer.messages) == 1 + + envelope = json.loads(writer.messages[0].decode()) + assert envelope["_version"] == client.core_version + assert envelope["_payload"] == { + "body": "hello", + "trace": ["hook-1", "hook-2"], + } + finally: + client.loop.close() diff --git a/tests/test_client_send_data.py b/tests/test_client_send_data.py new file mode 100644 index 0000000..79a3684 --- /dev/null +++ b/tests/test_client_send_data.py @@ -0,0 +1,345 @@ +import asyncio +import json + +import pytest + +from summoner.client import ClientMerger, ClientTranslation +from summoner.client.client import SummonerClient +from summoner.protocol import Action +from summoner.protocol.triggers import load_triggers +from tests.helpers import DummyWriter + + +def test_send_use_data_requires_one_argument(): + client = SummonerClient("send-data") + client.flow().activate() + + try: + with pytest.raises(TypeError): + @client.send(route="request", on_actions={Action.STAY}, use_data=True) + async def bad_sender() -> None: + return None + finally: + client.loop.close() + + +def test_send_use_data_accepts_annotated_parameter_types(): + client = SummonerClient("send-data") + client.flow().activate() + + try: + @client.send(route="request", on_actions={Action.STAY}, use_data=True) + async def good_sender(data: int) -> None: + return None + + client.loop.run_until_complete(client._wait_for_registration()) + assert client.sender_index["request"][0].fn is good_sender + finally: + client.loop.close() + + +def test_send_default_dna_writes_use_data_false(): + client = SummonerClient("send-data") + + try: + @client.send(route="request") + async def plain_sender() -> None: + return None + + client.loop.run_until_complete(client._wait_for_registration()) + assert client._dna_senders[0]["use_data"] is False + dna_entries = json.loads(client.dna()) + assert dna_entries[0]["use_data"] is False + finally: + client.loop.close() + + +def test_client_translation_replays_use_data_sender(): + source = SummonerClient("source") + source.flow().activate() + + translated = None + + try: + @source.send(route="request", on_actions={Action.STAY}, use_data=True) + async def source_sender(data: dict) -> None: + return None + + source.loop.run_until_complete(source._wait_for_registration()) + + translated = ClientTranslation(json.loads(source.dna()), name="translated") + translated.flow().activate() + translated.initiate_senders() + translated.loop.run_until_complete(translated._wait_for_registration()) + + sender = translated.sender_index["request"][0] + assert sender.use_data is True + finally: + source.loop.close() + if translated is not None: + translated.loop.close() + + +def test_client_translation_tolerates_legacy_sender_dna_without_use_data_field(): + source = SummonerClient("source") + translated = None + + try: + @source.send(route="request") + async def source_sender() -> None: + return None + + source.loop.run_until_complete(source._wait_for_registration()) + + dna_entries = json.loads(source.dna()) + assert dna_entries[0]["use_data"] is False + del dna_entries[0]["use_data"] + + translated = ClientTranslation(dna_entries, name="translated") + translated.initiate_senders() + translated.loop.run_until_complete(translated._wait_for_registration()) + + sender = translated.sender_index["request"][0] + assert sender.use_data is False + finally: + source.loop.close() + if translated is not None: + translated.loop.close() + + +def test_client_merger_requires_flow_activation_for_use_data_sender(): + source = SummonerClient("source") + source.flow().activate() + + merged = None + + try: + @source.send(route="request", on_actions={Action.STAY}, use_data=True) + async def source_sender(data: dict) -> None: + return None + + source.loop.run_until_complete(source._wait_for_registration()) + + merged = ClientMerger([source], name="merged", close_subclients=False) + + with pytest.raises(RuntimeError): + merged.initiate_senders() + finally: + source.loop.close() + if merged is not None: + merged.loop.close() + + +def test_client_merger_tolerates_legacy_imported_client_sender_defaults(): + source = SummonerClient("source") + merged = None + + try: + @source.send(route="request") + async def source_sender() -> None: + return None + + source.loop.run_until_complete(source._wait_for_registration()) + + del source._dna_senders[0]["multi"] + del source._dna_senders[0]["on_triggers"] + del source._dna_senders[0]["on_actions"] + del source._dna_senders[0]["use_data"] + + merged = ClientMerger([source], name="merged", close_subclients=False) + merged.initiate_senders() + merged.loop.run_until_complete(merged._wait_for_registration()) + + sender = merged.sender_index["request"][0] + assert sender.multi is False + assert sender.triggers is None + assert sender.actions is None + assert sender.use_data is False + finally: + source.loop.close() + if merged is not None: + merged.loop.close() + + +def test_receive_validation_is_not_globally_weakened(): + client = SummonerClient("send-data") + + try: + with pytest.raises(TypeError): + @client.receive(route="request") + async def bad_receiver(payload: int) -> None: + return None + finally: + client.loop.close() + + +def test_send_use_data_requires_reactive_sender(): + client = SummonerClient("send-data") + client.flow().activate() + + try: + with pytest.raises(ValueError): + @client.send(route="request", use_data=True) + async def bad_sender(data: dict) -> None: + return None + finally: + client.loop.close() + + +def test_send_rejects_non_string_route_cleanly(): + client = SummonerClient("send-data") + + try: + with pytest.raises(TypeError): + client.send(route=None) + finally: + client.loop.close() + + +def test_receive_rejects_non_string_route_cleanly(): + client = SummonerClient("send-data") + + try: + with pytest.raises(TypeError): + client.receive(route=None) + finally: + client.loop.close() + + +def test_send_use_data_requires_non_empty_reactive_filters(): + client = SummonerClient("send-data") + client.flow().activate() + + try: + with pytest.raises(ValueError): + @client.send(route="request", on_actions=set(), use_data=True) + async def bad_sender(data: dict) -> None: + return None + finally: + client.loop.close() + + +def test_send_use_data_passes_queued_event_data_to_sender(): + client = SummonerClient("send-data") + client.flow().activate() + + seen: list[dict] = [] + + try: + @client.send(route="request", on_actions={Action.STAY}, use_data=True) + async def good_sender(data: dict) -> dict: + seen.append(data) + return {"echo": data} + + client.loop.run_until_complete(client._wait_for_registration()) + + writer = DummyWriter() + stop_event = asyncio.Event() + + client.send_queue = asyncio.Queue() + client.batch_drain = True + client.max_consecutive_worker_errors = 3 + + sender = client.sender_index["request"][0] + payload = {"turn": 3, "text": "hello"} + + client.send_queue.put_nowait(("request", sender, payload)) + client.send_queue.put_nowait(None) + + client.loop.run_until_complete(client._send_worker(writer, stop_event)) + + assert seen == [payload] + assert client._dna_senders[0]["use_data"] is True + + dna_entries = json.loads(client.dna()) + assert dna_entries[0]["use_data"] is True + + assert writer.messages + assert b'"echo"' in writer.messages[0] + assert b'"hello"' in writer.messages[0] + finally: + client.loop.close() + + +def test_send_use_data_preserves_one_call_per_queued_event(): + client = SummonerClient("send-data") + client.flow().activate() + + Trigger = load_triggers(json_dict={"go": None}) + seen: list[dict] = [] + + try: + @client.send(route="request", on_actions={Action.STAY}, use_data=True) + async def good_sender(data: dict) -> None: + seen.append(data) + if len(seen) == 2: + async with client.connection_lock: + client._quit = True + return None + + client.loop.run_until_complete(client._wait_for_registration()) + + writer = DummyWriter() + stop_event = asyncio.Event() + + client.send_queue = asyncio.Queue() + client.event_bridge = asyncio.Queue() + client.batch_drain = True + client.max_concurrent_workers = 1 + client.max_consecutive_worker_errors = 3 + client.send_queue_maxsize = 8 + + parsed_route = client.flow().parse_route("request") + first = {"turn": 1, "text": "hello"} + second = {"turn": 2, "text": "again"} + + client.event_bridge.put_nowait(((), "peer-a", parsed_route, Action.STAY(Trigger.go, data=first))) + client.event_bridge.put_nowait(((), "peer-a", parsed_route, Action.STAY(Trigger.go, data=second))) + + worker = client.loop.create_task(client._send_worker(writer, stop_event)) + client.loop.run_until_complete(client.message_sender_loop(writer, stop_event)) + client.loop.run_until_complete(worker) + + assert seen == [first, second] + finally: + client.loop.close() + + +def test_send_without_use_data_keeps_legacy_dedup_behavior(): + client = SummonerClient("send-data") + client.flow().activate() + + Trigger = load_triggers(json_dict={"go": None}) + seen: list[str] = [] + + try: + @client.send(route="request", on_actions={Action.STAY}) + async def good_sender() -> None: + seen.append("sent") + async with client.connection_lock: + client._quit = True + return None + + client.loop.run_until_complete(client._wait_for_registration()) + + writer = DummyWriter() + stop_event = asyncio.Event() + + client.send_queue = asyncio.Queue() + client.event_bridge = asyncio.Queue() + client.batch_drain = True + client.max_concurrent_workers = 1 + client.max_consecutive_worker_errors = 3 + client.send_queue_maxsize = 8 + + parsed_route = client.flow().parse_route("request") + client.event_bridge.put_nowait(((), "peer-a", parsed_route, Action.STAY(Trigger.go))) + client.event_bridge.put_nowait(((), "peer-a", parsed_route, Action.STAY(Trigger.go))) + + worker = client.loop.create_task(client._send_worker(writer, stop_event)) + client.loop.run_until_complete(client.message_sender_loop(writer, stop_event)) + client.loop.run_until_complete(worker) + + assert seen == ["sent"] + finally: + client.loop.close() diff --git a/tests/test_merger_context.py b/tests/test_merger_context.py new file mode 100644 index 0000000..c29fe71 --- /dev/null +++ b/tests/test_merger_context.py @@ -0,0 +1,110 @@ +from summoner.client import ClientMerger, ClientTranslation, SummonerClient +from summoner.protocol import Action + + +def test_client_translation_applies_context_globals_recipes_and_optional_imports(): + dna_list = [ + { + "type": "__context__", + "var_name": "sdk_agent", + "imports": ["from math import sqrt"], + "globals": {"SDK_FLAG": "ready"}, + "recipes": {"SDK_NUMBERS": "[1, 2, 3]"}, + } + ] + + translated = ClientTranslation(dna_list, name="translated", allow_context_imports=False) + + try: + assert translated._var_name == "sdk_agent" + assert translated._sandbox_globals["sdk_agent"] is translated + assert translated._sandbox_globals["SDK_FLAG"] == "ready" + assert translated._sandbox_globals["SDK_NUMBERS"] == [1, 2, 3] + assert "sqrt" not in translated._sandbox_globals + finally: + translated.loop.close() + + +def test_client_translation_can_execute_context_imports_when_allowed(): + dna_list = [ + { + "type": "__context__", + "var_name": "sdk_agent", + "imports": ["from math import sqrt"], + } + ] + + translated = ClientTranslation(dna_list, name="translated-imports", allow_context_imports=True) + + try: + assert translated._sandbox_globals["sqrt"](25) == 5 + finally: + translated.loop.close() + + +def test_client_merger_normalizes_dna_source_context_and_reports_skipped_imports(): + dna_list = [ + { + "type": "__context__", + "var_name": "sdk_agent", + "imports": ["from math import sqrt"], + "globals": {"SDK_FLAG": "ready"}, + "recipes": {"SDK_NUMBERS": "[1, 2]"}, + } + ] + + merged = ClientMerger([dna_list], name="merged-dna", allow_context_imports=False, close_subclients=False) + + try: + source = merged.sources[0] + assert source["kind"] == "dna" + assert source["var_name"] == "sdk_agent" + assert source["globals"]["sdk_agent"] is merged + assert source["globals"]["SDK_FLAG"] == "ready" + assert source["globals"]["SDK_NUMBERS"] == [1, 2] + assert "sqrt" not in source["globals"] + assert source["import_report"]["skipped"] == ["from math import sqrt"] + finally: + merged.loop.close() + + +def test_client_merger_infers_imported_client_var_name_and_rebinds_handler_globals(): + source = SummonerClient("source-client") + merged = None + binding_name = "sdk_merge_agent" + previous = globals().get(binding_name) + globals()[binding_name] = source + + try: + @source.receive(route="request", priority=(1,)) + async def receive_request(payload: str): + return {"payload": payload, "client_name": sdk_merge_agent.name} + + @source.send(route="request", on_actions={Action.STAY}) + async def send_request() -> dict: + return {"client_name": sdk_merge_agent.name} + + source.loop.run_until_complete(source._wait_for_registration()) + + merged = ClientMerger([source], name="merged-client", close_subclients=False) + assert merged.sources[0]["var_name"] == binding_name + + merged.initiate_receivers() + merged.initiate_senders() + merged.loop.run_until_complete(merged._wait_for_registration()) + + assert merged.loop.run_until_complete(merged.receiver_index["request"].fn("hello")) == { + "payload": "hello", + "client_name": "merged-client", + } + assert merged.loop.run_until_complete(merged.sender_index["request"][0].fn()) == { + "client_name": "merged-client", + } + finally: + if previous is None: + globals().pop(binding_name, None) + else: + globals()[binding_name] = previous + source.loop.close() + if merged is not None: + merged.loop.close() diff --git a/tests/test_payload.py b/tests/test_payload.py new file mode 100644 index 0000000..c5db77a --- /dev/null +++ b/tests/test_payload.py @@ -0,0 +1,56 @@ +import json + +import pytest + +from summoner.protocol.payload import recover_with_types, wrap_with_types + + +def test_wrap_and_recover_typed_payload_round_trip(): + payload = { + "text": "hello", + "flag": True, + "count": 3, + "ratio": 1.5, + "items": ["a", 2, None, {"nested": False}], + } + + wrapped = wrap_with_types(payload, version="1.0.1") + assert wrapped.endswith("\n") + + relayed = json.dumps({ + "remote_addr": "peer-a", + "content": json.loads(wrapped), + }) + + assert recover_with_types(relayed) == { + "remote_addr": "peer-a", + "content": payload, + } + + +def test_recover_with_types_returns_plain_content_for_non_enveloped_message(): + message = json.dumps({ + "remote_addr": "peer-a", + "content": {"plain": "message"}, + }) + + assert recover_with_types(message) == { + "remote_addr": "peer-a", + "content": {"plain": "message"}, + } + + +def test_recover_with_types_returns_raw_text_for_non_json_warning(): + assert recover_with_types("warning: disconnected\n") == "warning: disconnected" + + +def test_recover_with_types_returns_plain_json_without_outer_wrapper(): + assert recover_with_types('{"status": "ok", "count": 2}\n') == { + "status": "ok", + "count": 2, + } + + +def test_wrap_with_types_rejects_unknown_version(): + with pytest.raises(ValueError): + wrap_with_types({"hello": "world"}, version="9.9.9") diff --git a/tests/test_process.py b/tests/test_process.py index 5478c14..e85c5b9 100644 --- a/tests/test_process.py +++ b/tests/test_process.py @@ -144,23 +144,28 @@ def test_sender_responds_to_filters(): test_evt = TestEvt(sigX) # 1) No filters → always responds - sender1 = Sender(fn=lambda: None, multi=False, actions=None, triggers=None) + sender1 = Sender(fn=lambda: None, multi=False, actions=None, triggers=None, use_data=False) assert sender1.responds_to(move_evt) assert sender1.responds_to(stay_evt) # 2) on_actions only - sender_actions = Sender(fn=lambda: None, multi=False, actions={Action.MOVE}, triggers=None) + sender_actions = Sender(fn=lambda: None, multi=False, actions={Action.MOVE}, triggers=None, use_data=False) assert sender_actions.responds_to(move_evt) assert not sender_actions.responds_to(stay_evt) # 3) on_triggers only — any event carrying sigX is accepted - sender_triggers = Sender(fn=lambda: None, multi=False, actions=None, triggers={sigX}) + sender_triggers = Sender(fn=lambda: None, multi=False, actions=None, triggers={sigX}, use_data=False) assert sender_triggers.responds_to(move_evt) assert sender_triggers.responds_to(test_evt) # ← change to True assert not sender_triggers.responds_to(stay_evt) # different signal Y # 4) both filters - sender_both = Sender(fn=lambda: None, multi=False, actions={Action.STAY}, triggers={sigY}) + sender_both = Sender(fn=lambda: None, multi=False, actions={Action.STAY}, triggers={sigY}, use_data=False) assert sender_both.responds_to(stay_evt) assert not sender_both.responds_to(StayEvt(sigX)) assert not sender_both.responds_to(move_evt) + + +def test_sender_use_data_defaults_to_false(): + sender = Sender(fn=lambda: None, multi=False, actions=None, triggers=None) + assert sender.use_data is False diff --git a/tests/test_triggers.py b/tests/test_triggers.py index 95c38cb..87af0c7 100644 --- a/tests/test_triggers.py +++ b/tests/test_triggers.py @@ -120,4 +120,17 @@ def test_event_and_action_classes_and_extract_signal(): assert extract_signal(sigX) is sigX assert extract_signal(None) is None with pytest.raises(TypeError): - extract_signal(123) \ No newline at end of file + extract_signal(123) + + +def test_event_can_store_optional_data(): + Trigger = load_triggers(json_dict={"X": None}) + sigX = Trigger.X + + move_evt = Move(sigX, data={"payload": 7}) + stay_evt = Stay(sigX) + + assert move_evt.data == {"payload": 7} + assert stay_evt.data is None + assert extract_signal(move_evt) is sigX + assert "data={'payload': 7}" in repr(move_evt) From 405dd7d95618ca3cffcab9c348d1a47a321142bc Mon Sep 17 00:00:00 2001 From: Remy Tuyeras Date: Fri, 3 Apr 2026 18:06:52 -0400 Subject: [PATCH 2/4] new feature: add every and run_while for send orachestration --- summoner/client/client.py | 765 ++++++++++++++++++++++++----- summoner/client/merger.py | 159 +++++- summoner/protocol/process.py | 24 +- tests/test_client_send_data.py | 91 ++++ tests/test_client_send_stress.py | 681 +++++++++++++++++++++++++ tests/test_client_timed_senders.py | 370 ++++++++++++++ tests/test_process.py | 4 + 7 files changed, 1976 insertions(+), 118 deletions(-) create mode 100644 tests/test_client_send_stress.py create mode 100644 tests/test_client_timed_senders.py diff --git a/summoner/client/client.py b/summoner/client/client.py index b2fa656..cda443a 100644 --- a/summoner/client/client.py +++ b/summoner/client/client.py @@ -1,6 +1,7 @@ import os import sys import json +import copy from typing import ( Optional, Callable, @@ -12,7 +13,8 @@ import asyncio import signal import inspect -from collections import defaultdict +from collections import defaultdict, deque +from dataclasses import dataclass, field import platform target_path = os.path.abspath(os.path.join(os.path.dirname(os.path.abspath(__file__)), "..")) @@ -63,6 +65,32 @@ class ServerDisconnected(Exception): """Raised when the server closes the connection.""" pass + +# ==== CLIENT-INTERNAL SENDER RUNTIME TYPES ==== +# +# These types stay in `client.py` on purpose: they model the local scheduling +# and worker runtime of one Summoner client instance, not protocol-level send +# semantics shared across modules. + +@dataclass +class SendInvocation: + """Internal queue item for one sender invocation.""" + route: str + sender: Sender + data: Any = None + done: Optional[asyncio.Future] = None + + +@dataclass +class TimedSenderRuntime: + """Internal per-sender scheduler state for timed senders.""" + armed: bool = False + running: bool = False + pending_payloads: deque[Any] = field(default_factory=deque) + next_run_at: Optional[float] = None + in_flight_done: Optional[asyncio.Future] = None + run_while_task: Optional[asyncio.Task] = None + class SummonerClient: DEFAULT_MAX_BYTES_PER_LINE = 64 * 1024 # 64 KiB @@ -133,7 +161,15 @@ def __init__(self, name: Optional[str] = None): # Pass Event information from the receiving end to the sending end self.event_bridge: Optional[asyncio.Queue[tuple[tuple[int, ...], Optional[str], ParsedRoute, Event]]] = None + # Sender-side orchestration runtime. These structures belong to the + # client implementation rather than the protocol layer because they + # track local batching, timed admission, and worker handoff state. self.send_queue: Optional[asyncio.Queue] = None + self.timed_event_buffer: deque[tuple[tuple[int, ...], Optional[str], ParsedRoute, Event]] = deque() + self.timed_sender_state: dict[str, TimedSenderRuntime] = {} + self.timed_sender_lock = asyncio.Lock() + self.timed_sender_wakeup = asyncio.Event() + self._next_sender_registration_id = 0 self.send_workers_started = False # To avoid double-starting workers self.worker_tasks: list[asyncio.Task] = [] self.writer_lock = asyncio.Lock() @@ -460,6 +496,54 @@ async def register(): # ==== SENDER REGISTRATION ==== + def _allocate_sender_registration_id(self) -> str: + registration_id = f"sender:{self._next_sender_registration_id}" + self._next_sender_registration_id += 1 + return registration_id + + def _normalize_data_mode(self, use_data: bool, data_mode: Optional[str]) -> Optional[str]: + if not use_data: + return None + if data_mode is None: + return "live" + if data_mode not in {"live", "snapshot"}: + raise ValueError( + "Argument `data_mode` must be `None`, 'live', or 'snapshot'. " + f"Provided: {data_mode!r}" + ) + return data_mode + + def _serialize_run_while_spec( + self, + run_while: Any, + ) -> tuple[str, Optional[bool], Optional[str], Optional[str]]: + if run_while is None: + return ("none", None, None, None) + if isinstance(run_while, bool): + return ("bool", run_while, None, None) + if callable(run_while): + module_name = getattr(run_while, "__module__", None) + qualname = getattr(run_while, "__qualname__", None) + serialized_name = None + source = None + if isinstance(module_name, str) and module_name and isinstance(qualname, str) and qualname: + serialized_name = f"{module_name}:{qualname}" + else: + fallback_name = getattr(run_while, "__name__", None) + if isinstance(fallback_name, str) and fallback_name: + serialized_name = fallback_name + try: + source = inspect.getsource(run_while) + except Exception: + source = getattr(run_while, "__dna_source__", None) + if not (isinstance(source, str) and source.strip()): + source = None + return ("callable", None, serialized_name, source) + raise TypeError( + "Argument `run_while` must be `None`, a bool, or a callable returning " + f"bool/awaitable bool. Provided: {run_while!r}" + ) + def send( self, route: str, @@ -467,6 +551,9 @@ def send( on_triggers: Optional[set[Signal]] = None, on_actions: Optional[set[Type]] = None, use_data: bool = False, + data_mode: Optional[str] = None, + every: Optional[float] = None, + run_while: Any = None, ): if not isinstance(route, str): raise TypeError(f"Argument `route` must be string. Provided: {route}") @@ -513,6 +600,12 @@ def decorator(fn: Callable[..., Awaitable]): if not isinstance(use_data, bool): raise TypeError(f"Argument `use_data` must be Boolean. Provided: {use_data}") + if every is not None: + if isinstance(every, bool) or not isinstance(every, (int, float)): + raise TypeError(f"Argument `every` must be `None` or a positive number. Provided: {every!r}") + if every <= 0: + raise ValueError(f"Argument `every` must be positive. Provided: {every!r}") + if on_triggers is not None and ( not isinstance(on_triggers, set) or not all(isinstance(sig, Signal) for sig in on_triggers) @@ -535,6 +628,21 @@ def decorator(fn: Callable[..., Awaitable]): if use_data and not self._flow.in_use: raise RuntimeError("Argument `use_data=True` requires `client.flow().activate()` before sender registration") + + if every is not None and reactive_filters and not self._flow.in_use: + raise RuntimeError( + "Timed reactive senders require `client.flow().activate()` before sender registration" + ) + + if run_while is not None and every is None: + raise ValueError("Argument `run_while` requires `every`") + + if data_mode is not None and not use_data: + raise ValueError("Argument `data_mode` requires `use_data=True`") + + normalized_data_mode = self._normalize_data_mode(use_data, data_mode) + run_while_kind, run_while_value, run_while_name, run_while_source = self._serialize_run_while_spec(run_while) + registration_id = self._allocate_sender_registration_id() # ----[ DNA capture ]---- self._dna_senders.append({ @@ -544,13 +652,30 @@ def decorator(fn: Callable[..., Awaitable]): "on_triggers": on_triggers, "on_actions": on_actions, "use_data": use_data, + "data_mode": normalized_data_mode, + "every": every, + "run_while": run_while, + "run_while_kind": run_while_kind, + "run_while_value": run_while_value, + "run_while_name": run_while_name, + "run_while_source": run_while_source, "source": inspect.getsource(fn), }) # ----[ Registration Code ]---- async def register(): - sender = Sender(fn=fn, multi=multi, actions=on_actions, triggers=on_triggers, use_data=use_data) + sender = Sender( + fn=fn, + multi=multi, + actions=on_actions, + triggers=on_triggers, + use_data=use_data, + data_mode=normalized_data_mode, + every=every, + run_while=run_while, + registration_id=registration_id, + ) actions_exist = isinstance(on_actions, set) and bool(on_actions) triggers_exist = isinstance(on_triggers, set) and bool(on_triggers) @@ -619,6 +744,9 @@ def _iter_registered_handler_functions(self): fn = d.get("fn") if fn is not None: yield fn + run_while = d.get("run_while") + if callable(run_while): + yield run_while for d in self._dna_hooks: fn = d.get("fn") @@ -788,6 +916,12 @@ def dna(self, include_context: bool = False) -> str: "route_key": route_key, # stable route representative "multi": dna["multi"], "use_data": dna["use_data"], + "data_mode": dna["data_mode"], + "every": dna["every"], + "run_while_kind": dna["run_while_kind"], + "run_while_value": dna["run_while_value"], + "run_while_name": dna["run_while_name"], + "run_while_source": dna.get("run_while_source", None), # Serialize triggers/actions by name so they can be re-resolved later. "on_triggers": [t.name for t in (dna["on_triggers"] or [])], "on_actions": [a.__name__ for a in (dna["on_actions"] or [])], @@ -975,8 +1109,8 @@ async def _read_line_safe( async def message_receiver_loop( - self, - reader: asyncio.StreamReader, + self, + reader: asyncio.StreamReader, stop_event: asyncio.Event ): @@ -1114,7 +1248,7 @@ async def _safe_call(fn: Callable[[Any], Awaitable], payload: Any) -> Any: for priority, event_list in sorted(event_buffer.items(), key=lambda kv: kv[0]): for event_data in event_list: # this will block if the bridge is full, slowing down readers - await self.event_bridge.put((priority,) + event_data) + await self._enqueue_sender_event((priority,) + event_data) event_buffer = {} @@ -1131,6 +1265,204 @@ async def _safe_call(fn: Callable[[Any], Awaitable], payload: Any) -> Any: stop_event.set() raise + # ==== SEND DATA HELPERS ==== + + def _capture_send_data(self, value: Any, data_mode: Optional[str]) -> Any: + if data_mode in (None, "live"): + return value + if data_mode == "snapshot": + return copy.deepcopy(value) + raise TypeError(f"Unknown data_mode: {data_mode!r}") + + def _materialize_send_data(self, value: Any, data_mode: Optional[str]) -> Any: + if data_mode in (None, "live"): + return value + if data_mode == "snapshot": + return copy.deepcopy(value) + raise TypeError(f"Unknown data_mode: {data_mode!r}") + + # ==== TIMED SENDER GUARD HELPERS ==== + + async def _await_run_while_value(self, awaitable: Any) -> bool: + value = await awaitable + return bool(value) + + def _wake_timed_scheduler(self, _task: asyncio.Task) -> None: + self.timed_sender_wakeup.set() + + def _poll_run_while(self, sender: Sender, runtime: TimedSenderRuntime) -> Optional[bool]: + spec = sender.run_while + + if spec is None: + return True + if isinstance(spec, bool): + return spec + if not callable(spec): + raise TypeError(f"Invalid run_while specification: {spec!r}") + + task = runtime.run_while_task + if task is not None: + if not task.done(): + return None + runtime.run_while_task = None + try: + return bool(task.result()) + except Exception as e: + self.logger.warning(f"run_while predicate failed: {type(e).__name__}: {e}") + return False + + try: + value = spec() + if inspect.isawaitable(value): + # Async guards may legitimately take time to confirm whether the + # sender may proceed. Keep that wait local to this sender by + # tracking a dedicated task instead of stalling the whole timed + # scheduler. + task = self.loop.create_task(self._await_run_while_value(value)) + task.add_done_callback(self._wake_timed_scheduler) + runtime.run_while_task = task + return None + return bool(value) + except Exception as e: + self.logger.warning(f"run_while predicate failed: {type(e).__name__}: {e}") + return False + + # ==== SEND ROUTING / ORCHESTRATION HELPERS ==== + + @staticmethod + def _route_accepts(sender_pr: ParsedRoute, receiver_pr: ParsedRoute) -> bool: + source_ok = all(any(n.accepts(m) for m in receiver_pr.source) for n in sender_pr.source) + label_ok = all(any(n.accepts(m) for m in receiver_pr.label) for n in sender_pr.label) + target_ok = all(any(n.accepts(m) for m in receiver_pr.target) for n in sender_pr.target) + return source_ok and label_ok and target_ok + + @staticmethod + def _has_reactive_filters(sender: Sender) -> bool: + return bool( + (sender.actions and isinstance(sender.actions, set)) or + (sender.triggers and isinstance(sender.triggers, set)) + ) + + async def _snapshot_sender_registry(self) -> tuple[dict[str, list[Sender]], dict[str, ParsedRoute]]: + async with self.routes_lock: + sender_index = { + route: list(routed_senders) + for route, routed_senders in self.sender_index.items() + } + sender_parsed_routes = self.sender_parsed_routes.copy() + return sender_index, sender_parsed_routes + + def _drain_pending_events(self) -> list[tuple[tuple[int, ...], Optional[str], ParsedRoute, Event]]: + pending: list[tuple[tuple[int, ...], Optional[str], ParsedRoute, Event]] = [] + if not self._flow.in_use: + return pending + + try: + while True: + pending.append(self.event_bridge.get_nowait()) + except asyncio.QueueEmpty: + pass + + pending.sort(key=lambda it: hook_priority_order(it[0])) + return pending + + async def _enqueue_sender_event( + self, + item: tuple[tuple[int, ...], Optional[str], ParsedRoute, Event], + ) -> None: + await self.event_bridge.put(item) + + if self._flow.in_use: + # Timed senders keep their own admission feed so they can be armed + # promptly without changing the untimed batch-loop contract. + async with self.timed_sender_lock: + self.timed_event_buffer.append(item) + self.timed_sender_wakeup.set() + + def _make_send_invocation( + self, + route: str, + sender: Sender, + *, + data: Any = None, + track_completion: bool = False, + ) -> SendInvocation: + done = self.loop.create_future() if track_completion else None + return SendInvocation(route=route, sender=sender, data=data, done=done) + + async def _wait_for_send_invocations(self, invocations: list[SendInvocation]) -> None: + tracked = [ + invocation.done + for invocation in invocations + if invocation.done is not None + ] + if tracked: + await asyncio.gather(*tracked, return_exceptions=True) + + def _ensure_timed_runtime(self, sender: Sender) -> TimedSenderRuntime: + registration_id = sender.registration_id or sender.fn.__name__ + runtime = self.timed_sender_state.get(registration_id) + if runtime is None: + runtime = TimedSenderRuntime() + self.timed_sender_state[registration_id] = runtime + return runtime + + def _arm_timed_senders_from_pending_locked( + self, + sender_index: dict[str, list[Sender]], + sender_parsed_routes: dict[str, ParsedRoute], + pending: list[tuple[tuple[int, ...], Optional[str], ParsedRoute, Event]], + now: float, + ) -> None: + if not pending: + return + + for route, routed_senders in sender_index.items(): + sender_parsed_route = sender_parsed_routes.get(route) + if sender_parsed_route is None: + continue + + for sender in routed_senders: + if sender.every is None or (not self._has_reactive_filters(sender)): + continue + + runtime = self._ensure_timed_runtime(sender) + + for (_priority, _key, parsed_route, event) in pending: + if not ( + self._route_accepts(sender_parsed_route, parsed_route) + and sender.responds_to(event) + ): + continue + + if not runtime.armed: + runtime.armed = True + runtime.running = True + runtime.next_run_at = now + + if sender.use_data: + try: + runtime.pending_payloads.append( + self._capture_send_data(event.data, sender.data_mode) + ) + except Exception as e: + self.logger.warning( + f"Failed to capture timed sender data for '{sender.fn.__name__}' " + f"on route '{route}': {type(e).__name__}: {e}" + ) + else: + break + + async def _wait_for_timed_wakeup(self, timeout: float) -> None: + if timeout <= 0: + return + try: + await asyncio.wait_for(self.timed_sender_wakeup.wait(), timeout=timeout) + except asyncio.TimeoutError: + return + finally: + self.timed_sender_wakeup.clear() + # ==== SENDER EXECUTION ==== def _start_send_workers( @@ -1153,12 +1485,20 @@ async def _send_worker( while True: - item: Optional[tuple[str, Sender, Any]] = await self.send_queue.get() + item = await self.send_queue.get() if item is None: self.send_queue.task_done() break - - route, sender, sender_data = item + + invocation_done = None + if isinstance(item, SendInvocation): + route = item.route + sender = item.sender + sender_data = item.data + invocation_done = item.done + else: + route, sender, sender_data = item + try: if sender.use_data: result = await sender.fn(sender_data) @@ -1168,6 +1508,8 @@ async def _send_worker( # ----[ Urgent: Handle Aborts ]---- async with self.connection_lock: if self._quit: + if invocation_done is not None and not invocation_done.done(): + invocation_done.set_result(None) stop_event.set() break @@ -1219,18 +1561,23 @@ async def _send_worker( # ----[ Unpack: Post Messages ]---- async with self.writer_lock: writer.write(message) - + # No concurrency on batch_drain (initialized in run()) if not self.batch_drain: async with self.writer_lock: await writer.drain() + if invocation_done is not None and not invocation_done.done(): + invocation_done.set_result(None) + except Exception as e: consecutive_errors += 1 self.logger.error( f"Worker for {sender.fn.__name__} crashed ({consecutive_errors} in a row): {e}", exc_info=True ) + if invocation_done is not None and not invocation_done.done(): + invocation_done.set_result(e) # if 3 workers in a row have crashed, abort the session if consecutive_errors >= self.max_consecutive_worker_errors: self.logger.critical(f"{self.max_consecutive_worker_errors} consecutive worker failures; shutting down sender loop") @@ -1248,131 +1595,327 @@ async def _cleanup_workers(self): self.worker_tasks.clear() self.send_workers_started = False - async def message_sender_loop( - self, - writer: asyncio.StreamWriter, - stop_event: asyncio.Event - ): + async def _message_sender_batch_loop( + self, + writer: asyncio.StreamWriter, + stop_event: asyncio.Event, + ) -> None: + """ + Preserve the legacy sender contract for untimed senders. - # ----[ Helper: Matches Routes Between Senders and Receivers to Trigger Send ]---- - def _route_accepts( - sender_pr: ParsedRoute, - receiver_pr: ParsedRoute - ) -> bool: - source_ok = all(any(n.accepts(m) for m in receiver_pr.source) for n in sender_pr.source) - label_ok = all(any(n.accepts(m) for m in receiver_pr.label) for n in sender_pr.label) - target_ok = all(any(n.accepts(m) for m in receiver_pr.target) for n in sender_pr.target) - return source_ok and label_ok and target_ok + Untimed senders remain round-based: each loop pass collects the current + work, dispatches a batch, and waits for that batch to finish before + starting the next untimed round. This keeps pre-`every` behavior stable. + """ + while not stop_event.is_set(): + sender_index, sender_parsed_routes = await self._snapshot_sender_registry() + pending = self._drain_pending_events() - cancelled = False - try: + invocations: list[SendInvocation] = [] + emitted: set[tuple[str, Optional[str], str]] = set() - # ----[ Keep Sending While Actively Listening (No Travel) ]---- - while not stop_event.is_set(): - - # ----[ Prepare Sender Batch ]---- - - async with self.routes_lock: - sender_index: dict[str, list[Sender]] = self.sender_index.copy() + for route, routed_senders in sender_index.items(): + for sender in routed_senders: + if sender.every is not None: + continue - # ----[ Fast upload of pending event data ]---- - if self._flow.in_use: - - async with self.routes_lock: - sender_parsed_routes: dict[str, ParsedRoute] = self.sender_parsed_routes.copy() - - pending: list[tuple[tuple[int, ...], Optional[str], ParsedRoute, Event]] = [] - try: - while True: - pending.append(self.event_bridge.get_nowait()) - except asyncio.QueueEmpty: - pass - - pending.sort(key=lambda it: hook_priority_order(it[0])) + sender_is_reactive = self._has_reactive_filters(sender) - # ----[ Build Sender Batch ]---- - senders: list[tuple[str, Sender, Any]] = [] + if (not self._flow.in_use) or (not sender_is_reactive): + if sender.use_data: + self.logger.warning( + f"Skipping sender '{sender.fn.__name__}' on route '{route}': " + "use_data=True requires a reactive sender running with flow enabled" + ) + continue + invocations.append( + self._make_send_invocation( + route, + sender, + track_completion=True, + ) + ) + continue - # De-dup set: at most one sender per (route, key-from-recv, recv-handler-name) this cycle. - # `use_data=True` senders intentionally bypass this so each queued event payload is delivered. - emitted: set[tuple[str, Optional[str], str]] = set() + sender_parsed_route = sender_parsed_routes.get(route) + if sender_parsed_route is None: + continue - for route, routed_senders in sender_index.items(): - for sender in routed_senders: - - # Non-reactive (no actions/triggers): preserve current behavior - if (not self._flow.in_use) or (sender.actions is None and sender.triggers is None): - if sender.use_data: + for (_priority, key, parsed_route, event) in pending: + if not ( + self._route_accepts(sender_parsed_route, parsed_route) + and sender.responds_to(event) + ): + continue + + if sender.use_data: + try: + captured_data = self._capture_send_data(event.data, sender.data_mode) + invocations.append( + self._make_send_invocation( + route, + sender, + data=self._materialize_send_data(captured_data, sender.data_mode), + track_completion=True, + ) + ) + except Exception as e: self.logger.warning( - f"Skipping sender '{sender.fn.__name__}' on route '{route}': " - "use_data=True requires a reactive sender running with flow enabled" + f"Failed to prepare sender data for '{sender.fn.__name__}' " + f"on route '{route}': {type(e).__name__}: {e}" ) + continue + + dedup_key = (route, key, sender.fn.__name__) + if dedup_key not in emitted: + invocations.append( + self._make_send_invocation( + route, + sender, + track_completion=True, + ) + ) + emitted.add(dedup_key) + break + + if not invocations: + await asyncio.sleep(0.1) + continue + + queue_size = self.send_queue.qsize() + expected_queue_size = queue_size + len(invocations) + if expected_queue_size > self.send_queue_maxsize * 0.8: + self.logger.warning( + f"Queue is about to exceed 80% its capacity; " + f"Attempted load size: {expected_queue_size} out of {self.send_queue_maxsize}" + ) + + try: + for invocation in invocations: + await self.send_queue.put(invocation) + except asyncio.CancelledError: + self.logger.info("Sender enqueue loop cancelled mid-batch.") + raise + + await self._wait_for_send_invocations(invocations) + + if self.batch_drain: + async with self.writer_lock: + await writer.drain() + + async with self.connection_lock: + if self._travel or self._quit: + stop_event.set() + + async def _message_sender_timed_loop( + self, + writer: asyncio.StreamWriter, + stop_event: asyncio.Event, + ) -> None: + """ + Timed senders become scheduler-owned obligations after admission. + + Philosophy: + - The initial event is admission. It arms a timed sender when the system + can accept the work. + - Once armed, `every` creates an obligation: the scheduler should keep + servicing that sender on cadence without being blocked by unrelated + untimed batches. + + This loop intentionally owns only timed senders so legacy untimed + behavior stays round-based and unchanged. + """ + while not stop_event.is_set(): + sender_index, sender_parsed_routes = await self._snapshot_sender_registry() + invocations: list[SendInvocation] = [] + next_due_times: list[float] = [] + drain_needed = False + admission_now = self.loop.time() + + async with self.timed_sender_lock: + buffered_events = list(self.timed_event_buffer) + self.timed_event_buffer.clear() + if buffered_events: + self._arm_timed_senders_from_pending_locked( + sender_index, + sender_parsed_routes, + buffered_events, + admission_now, + ) + + for route, routed_senders in sender_index.items(): + for sender in routed_senders: + if sender.every is None: + continue + + sender_is_reactive = self._has_reactive_filters(sender) + runtime = self._ensure_timed_runtime(sender) + + if runtime.in_flight_done is not None: + if runtime.in_flight_done.done(): + runtime.in_flight_done = None + drain_needed = True + else: + if runtime.next_run_at is not None: + next_due_times.append(runtime.next_run_at) continue - senders.append((route, sender, None)) - - # Reactive: require matching a pending activation (existential) - elif self._flow.in_use and ((sender.actions and isinstance(sender.actions, set)) or - (sender.triggers and isinstance(sender.triggers, set))): - - sender_parsed_route = sender_parsed_routes.get(route) - if sender_parsed_route is None: - continue - - # Iterate pending in queue order. - # `use_data=True` senders receive one call per matching queued event. - # Other reactive senders keep the existing first-match de-dup behavior. - for (_priority, key, parsed_route, event) in pending: - if _route_accepts(sender_parsed_route, parsed_route) and sender.responds_to(event): - if sender.use_data: - senders.append((route, sender, event.data)) - continue - - dedup_key = (route, key, sender.fn.__name__) # key scopes to the activation thread/peer - if dedup_key not in emitted: - senders.append((route, sender, None)) - emitted.add(dedup_key) - break # do not enqueue multiple times for this sender this cycle - - # ----[ Empty: Skip and Prevent Client Overwhelming | Almost full: warning ]---- - if not senders: - await asyncio.sleep(0.1) # Time - continue - else: - queue_size = self.send_queue.qsize() - expected_queue_size = queue_size + len(senders) - if expected_queue_size > self.send_queue_maxsize * 0.8: # 80% full - self.logger.warning(f"Queue is about to exceed 80% its capacity; Attempted load size: {expected_queue_size} out of {self.send_queue_maxsize}") - # ----[ Enqueue Sender Batch | Senders Are Run in Background ]---- + if sender_is_reactive and not runtime.armed: + continue + run_while_allowed = self._poll_run_while(sender, runtime) + if run_while_allowed is None: + continue + + if not run_while_allowed: + if sender_is_reactive: + runtime.armed = False + runtime.running = False + runtime.pending_payloads.clear() + runtime.next_run_at = None + runtime.in_flight_done = None + if runtime.run_while_task is not None: + runtime.run_while_task.cancel() + runtime.run_while_task = None + else: + runtime.running = False + continue + + now = self.loop.time() + runtime.running = True + + if runtime.next_run_at is None: + runtime.next_run_at = now + + if now < runtime.next_run_at: + next_due_times.append(runtime.next_run_at) + continue + + timed_batch: list[SendInvocation] = [] + + if sender.use_data: + buffered_payloads = list(runtime.pending_payloads) + runtime.pending_payloads.clear() + + for buffered_payload in buffered_payloads: + try: + payload = self._materialize_send_data(buffered_payload, sender.data_mode) + except Exception as e: + self.logger.warning( + f"Failed to materialize timed sender data for '{sender.fn.__name__}' " + f"on route '{route}': {type(e).__name__}: {e}" + ) + continue + + timed_batch.append( + self._make_send_invocation( + route, + sender, + data=payload, + track_completion=True, + ) + ) + else: + timed_batch.append( + self._make_send_invocation( + route, + sender, + track_completion=True, + ) + ) + + if timed_batch: + runtime.in_flight_done = asyncio.gather( + *[ + invocation.done + for invocation in timed_batch + if invocation.done is not None + ], + return_exceptions=True, + ) + invocations.extend(timed_batch) + else: + runtime.in_flight_done = None + + runtime.next_run_at = now + float(sender.every) + next_due_times.append(runtime.next_run_at) + + if drain_needed and self.batch_drain: + async with self.writer_lock: + await writer.drain() + + if invocations: + queue_size = self.send_queue.qsize() + expected_queue_size = queue_size + len(invocations) + if expected_queue_size > self.send_queue_maxsize * 0.8: + self.logger.warning( + f"Queue is about to exceed 80% its capacity; " + f"Attempted load size: {expected_queue_size} out of {self.send_queue_maxsize}" + ) + try: - for sender in senders: - await self.send_queue.put(sender) # Will block if full (i.e., back-pressure) + for invocation in invocations: + await self.send_queue.put(invocation) except asyncio.CancelledError: - self.logger.info("Sender enqueue loop cancelled mid-batch.") + self.logger.info("Timed sender enqueue loop cancelled mid-batch.") raise - # ----[ Wait for Sender Batch to Finish]---- - await self.send_queue.join() + sleep_for = 0.1 + if next_due_times: + next_due_at = min(next_due_times) + sleep_for = max(0.0, min(0.1, next_due_at - self.loop.time())) + await self._wait_for_timed_wakeup(sleep_for) - if self.batch_drain: - async with self.writer_lock: - await writer.drain() + async with self.connection_lock: + if self._travel or self._quit: + stop_event.set() - # ----[ Quit or Travel ]---- - async with self.connection_lock: - if self._travel or self._quit: - stop_event.set() + async def message_sender_loop( + self, + writer: asyncio.StreamWriter, + stop_event: asyncio.Event + ): + cancelled = False + batch_task: Optional[asyncio.Task] = None + timed_task: Optional[asyncio.Task] = None + + try: + self.timed_sender_wakeup.clear() + batch_task = self.loop.create_task(self._message_sender_batch_loop(writer, stop_event)) + timed_task = self.loop.create_task(self._message_sender_timed_loop(writer, stop_event)) + await asyncio.gather(batch_task, timed_task) except asyncio.CancelledError: self.logger.info("Client about to disconnect...") cancelled = True - # do NOT re-raise yet; let finally run first + raise finally: - # Best-effort signal to workers; never block on shutdown + stop_event.set() + self.timed_sender_wakeup.set() + + tasks_to_finish = [ + task for task in (batch_task, timed_task) + if task is not None + ] + for task in tasks_to_finish: + if not task.done(): + task.cancel() + if tasks_to_finish: + await asyncio.gather(*tasks_to_finish, return_exceptions=True) + + async with self.timed_sender_lock: + for runtime in self.timed_sender_state.values(): + if runtime.run_while_task is not None and not runtime.run_while_task.done(): + runtime.run_while_task.cancel() + pending_run_while = [ + runtime.run_while_task + for runtime in self.timed_sender_state.values() + if runtime.run_while_task is not None + ] + if pending_run_while: + await asyncio.gather(*pending_run_while, return_exceptions=True) + if self.send_queue is not None: - # This may result in redundant cancellation if shutdown() is also called, - # but guarantees all workers get signaled even in abrupt exits. for _ in range(self.max_concurrent_workers): try: if cancelled: @@ -1401,6 +1944,7 @@ async def handle_session(self, host: str = '127.0.0.1', port: int = 8888): await self._cleanup_workers() self.send_queue = asyncio.Queue(maxsize=self.send_queue_maxsize) self.event_bridge = asyncio.Queue(maxsize = self.event_bridge_maxsize) + self.timed_sender_state.clear() # reset any previous travel/quit intent so each session starts fresh; # travel is only honored if set after this point, quit likewise @@ -1472,6 +2016,7 @@ async def handle_session(self, host: str = '127.0.0.1', port: int = 8888): # Clean up worker used in the sender loop await self._cleanup_workers() + self.timed_sender_state.clear() # Deregister this session and its children from active tasks async with self.tasks_lock: diff --git a/summoner/client/merger.py b/summoner/client/merger.py index b03bd12..d46428e 100644 --- a/summoner/client/merger.py +++ b/summoner/client/merger.py @@ -49,7 +49,7 @@ """ from importlib import import_module -from typing import Optional, Any +from typing import Optional, Any, Callable from contextlib import suppress from pathlib import Path import inspect @@ -166,6 +166,97 @@ def _read_required_mapping_field(entry: dict[str, Any], field: str): return entry[field] +def _resolve_callable_reference(globals_dict: dict[str, Any], ref: Optional[str]) -> Optional[Callable[..., Any]]: + if not isinstance(ref, str) or not ref: + return None + + if ":" in ref: + module_name, qualname = ref.split(":", 1) + try: + obj = import_module(module_name) + except Exception: + obj = None + + if obj is not None: + try: + for part in qualname.split("."): + if part == "": + return None + obj = getattr(obj, part) + if callable(obj): + return obj + except Exception: + pass + + fallback_name = qualname.split(".")[-1] + fallback = globals_dict.get(fallback_name) + if callable(fallback): + return fallback + return None + + candidate = globals_dict.get(ref) + if callable(candidate): + return candidate + return None + + +def _resolve_callable_reference_from_source( + globals_dict: dict[str, Any], + ref: Optional[str], + source: Optional[str], + ) -> Optional[Callable[..., Any]]: + if not isinstance(source, str) or not source.strip(): + return None + + expected_name = None + if isinstance(ref, str) and ref: + if ":" in ref: + _, qualname = ref.split(":", 1) + if "" not in qualname: + expected_name = qualname.split(".")[-1] + else: + expected_name = ref + + if not isinstance(expected_name, str) or not expected_name or expected_name == "": + return None + + try: + if "__builtins__" not in globals_dict: + globals_dict["__builtins__"] = __builtins__ + exec(compile(textwrap.dedent(source), filename="", mode="exec"), globals_dict) + except Exception: + return None + + candidate = globals_dict.get(expected_name) + if callable(candidate): + return candidate + return None + + +def _resolve_run_while_spec( + globals_dict: dict[str, Any], + kind: str, + value: Any, + name: Optional[str], + source: Optional[str] = None, + ) -> Any: + if kind == "none": + return None + if kind == "bool": + return bool(value) + if kind == "callable": + resolved = _resolve_callable_reference(globals_dict, name) + if not callable(resolved): + resolved = _resolve_callable_reference_from_source(globals_dict, name, source) + if callable(resolved): + return resolved + raise ValueError( + "Could not resolve serialized run_while callable " + f"{name!r} from available replay context" + ) + raise ValueError(f"Unknown run_while kind {kind!r}") + + class ClientMerger(SummonerClient): """ Merge multiple sources into one client. @@ -792,19 +883,34 @@ def initiate_senders(self): if not self._flow.in_use: for src in self.sources: if src["kind"] == "client": - if any(dna.get("use_data", False) for dna in src["client"]._dna_senders): + if any( + dna.get("use_data", False) or ( + dna.get("every") is not None and ( + (dna.get("on_triggers") or []) or + (dna.get("on_actions") or []) + ) + ) + for dna in src["client"]._dna_senders + ): raise RuntimeError( "ClientMerger.flow().activate() must be called before initiate_senders() " - "when replaying senders with use_data=True" + "when replaying reactive timed senders or senders with use_data=True" ) else: if any( - entry.get("type") == "send" and entry.get("use_data", False) + entry.get("type") == "send" and ( + entry.get("use_data", False) or ( + entry.get("every") is not None and ( + (entry.get("on_triggers") or []) or + (entry.get("on_actions") or []) + ) + ) + ) for entry in src["dna_entries"] ): raise RuntimeError( "ClientMerger.flow().activate() must be called before initiate_senders() " - "when replaying senders with use_data=True" + "when replaying reactive timed senders or senders with use_data=True" ) for src in self.sources: @@ -815,12 +921,24 @@ def initiate_senders(self): fn_clone = self._clone_handler(dna["fn"], var_name) try: route = _read_required_mapping_field(dna, "route") + run_while = dna.get("run_while") + if run_while is None: + run_while = _resolve_run_while_spec( + fn_clone.__globals__, + dna.get("run_while_kind", "none"), + dna.get("run_while_value", None), + dna.get("run_while_name", None), + dna.get("run_while_source", None), + ) self.send( route, multi=dna.get("multi", False), on_triggers=dna.get("on_triggers"), on_actions=dna.get("on_actions"), use_data=dna.get("use_data", False), + data_mode=dna.get("data_mode", None), + every=dna.get("every", None), + run_while=run_while, )(fn_clone) except Exception as e: self.logger.warning( @@ -846,12 +964,22 @@ def initiate_senders(self): g["Trigger"] = TriggerCls on_triggers = {_resolve_trigger(TriggerCls, t) for t in trigger_names} or None on_actions = {_resolve_action(Action, a) for a in entry.get("on_actions", [])} or None + run_while = _resolve_run_while_spec( + g, + entry.get("run_while_kind", "none"), + entry.get("run_while_value", None), + entry.get("run_while_name", None), + entry.get("run_while_source", None), + ) dec = self.send( route, multi=entry.get("multi", False), on_triggers=on_triggers, on_actions=on_actions, use_data=entry.get("use_data", False), + data_mode=entry.get("data_mode", None), + every=entry.get("every", None), + run_while=run_while, ) self._apply_with_source_patch(dec, fn, entry["source"]) @@ -1171,12 +1299,19 @@ def initiate_senders(self): - Action is resolved from the Action container by name """ if (not self._flow.in_use) and any( - entry.get("type") == "send" and entry.get("use_data", False) + entry.get("type") == "send" and ( + entry.get("use_data", False) or ( + entry.get("every") is not None and ( + (entry.get("on_triggers") or []) or + (entry.get("on_actions") or []) + ) + ) + ) for entry in self._dna_list ): raise RuntimeError( "ClientTranslation.flow().activate() must be called before initiate_senders() " - "when replaying senders with use_data=True" + "when replaying reactive timed senders or senders with use_data=True" ) g = self._sandbox_globals @@ -1199,12 +1334,22 @@ def initiate_senders(self): g["Trigger"] = TriggerCls on_triggers = {_resolve_trigger(TriggerCls, t) for t in trigger_names} or None on_actions = {_resolve_action(Action, a) for a in entry.get("on_actions", [])} or None + run_while = _resolve_run_while_spec( + g, + entry.get("run_while_kind", "none"), + entry.get("run_while_value", None), + entry.get("run_while_name", None), + entry.get("run_while_source", None), + ) dec = self.send( route, multi=entry.get("multi", False), on_triggers=on_triggers, on_actions=on_actions, use_data=entry.get("use_data", False), + data_mode=entry.get("data_mode", None), + every=entry.get("every", None), + run_while=run_while, ) self._apply_with_source_patch(dec, fn, entry["source"]) diff --git a/summoner/protocol/process.py b/summoner/protocol/process.py index bfaa4c9..cf1f86c 100644 --- a/summoner/protocol/process.py +++ b/summoner/protocol/process.py @@ -306,12 +306,26 @@ def activated_nodes( @dataclass(frozen=True, init=False) class Sender: - __slots__ = ('fn', 'multi', 'actions', 'triggers', 'use_data') + __slots__ = ( + 'fn', + 'multi', + 'actions', + 'triggers', + 'use_data', + 'data_mode', + 'every', + 'run_while', + 'registration_id', + ) fn: Callable[..., Awaitable] multi: bool actions: Optional[set[Type]] triggers: Optional[set[Signal]] use_data: bool + data_mode: Optional[str] + every: Optional[float] + run_while: Any + registration_id: Optional[str] def __init__( self, @@ -320,12 +334,20 @@ def __init__( actions: Optional[set[Type]], triggers: Optional[set[Signal]], use_data: bool = False, + data_mode: Optional[str] = None, + every: Optional[float] = None, + run_while: Any = None, + registration_id: Optional[str] = None, ): object.__setattr__(self, "fn", fn) object.__setattr__(self, "multi", multi) object.__setattr__(self, "actions", actions) object.__setattr__(self, "triggers", triggers) object.__setattr__(self, "use_data", use_data) + object.__setattr__(self, "data_mode", data_mode) + object.__setattr__(self, "every", every) + object.__setattr__(self, "run_while", run_while) + object.__setattr__(self, "registration_id", registration_id) def responds_to(self, event: Any) -> bool: action_check = True diff --git a/tests/test_client_send_data.py b/tests/test_client_send_data.py index 79a3684..05cc8de 100644 --- a/tests/test_client_send_data.py +++ b/tests/test_client_send_data.py @@ -6,6 +6,7 @@ from summoner.client import ClientMerger, ClientTranslation from summoner.client.client import SummonerClient from summoner.protocol import Action +from summoner.protocol.process import Direction from summoner.protocol.triggers import load_triggers from tests.helpers import DummyWriter @@ -48,8 +49,16 @@ async def plain_sender() -> None: client.loop.run_until_complete(client._wait_for_registration()) assert client._dna_senders[0]["use_data"] is False + assert client._dna_senders[0]["data_mode"] is None + assert client._dna_senders[0]["every"] is None + assert client._dna_senders[0]["run_while_kind"] == "none" + assert client._dna_senders[0]["run_while_source"] is None dna_entries = json.loads(client.dna()) assert dna_entries[0]["use_data"] is False + assert dna_entries[0]["data_mode"] is None + assert dna_entries[0]["every"] is None + assert dna_entries[0]["run_while_kind"] == "none" + assert dna_entries[0]["run_while_source"] is None finally: client.loop.close() @@ -261,6 +270,37 @@ async def good_sender(data: dict) -> dict: client.loop.close() +def test_send_multi_worker_writes_all_payloads(): + client = SummonerClient("send-data") + + try: + @client.send(route="request", multi=True) + async def good_sender() -> list[dict]: + return [{"part": 1}, {"part": 2}] + + client.loop.run_until_complete(client._wait_for_registration()) + + writer = DummyWriter() + stop_event = asyncio.Event() + + client.send_queue = asyncio.Queue() + client.batch_drain = True + client.max_consecutive_worker_errors = 3 + + sender = client.sender_index["request"][0] + + client.send_queue.put_nowait(("request", sender, None)) + client.send_queue.put_nowait(None) + + client.loop.run_until_complete(client._send_worker(writer, stop_event)) + + assert len(writer.messages) == 2 + assert b'"part": 1' in writer.messages[0] + assert b'"part": 2' in writer.messages[1] + finally: + client.loop.close() + + def test_send_use_data_preserves_one_call_per_queued_event(): client = SummonerClient("send-data") client.flow().activate() @@ -305,6 +345,57 @@ async def good_sender(data: dict) -> None: client.loop.close() +def test_send_multi_use_data_on_triggers_only_emits_all_payloads(): + client = SummonerClient("send-data") + client.flow().activate() + + Trigger = load_triggers(json_dict={"go": None}) + seen: list[int] = [] + + try: + @client.hook(Direction.SEND, priority=(1,)) + async def stop_after_last(payload: dict) -> dict: + if payload["turn"] == 2 and payload["part"] == 2: + await client.quit() + return payload + + @client.send(route="request", on_triggers={Trigger.go}, use_data=True, multi=True) + async def good_sender(data: dict) -> list[dict]: + seen.append(data["turn"]) + return [ + {"turn": data["turn"], "part": 1}, + {"turn": data["turn"], "part": 2}, + ] + + client.loop.run_until_complete(client._wait_for_registration()) + + writer = DummyWriter() + stop_event = asyncio.Event() + + client.send_queue = asyncio.Queue() + client.event_bridge = asyncio.Queue() + client.batch_drain = True + client.max_concurrent_workers = 1 + client.max_consecutive_worker_errors = 3 + client.send_queue_maxsize = 8 + + parsed_route = client.flow().parse_route("request") + client.event_bridge.put_nowait(((), "peer-a", parsed_route, Action.TEST(Trigger.go, data={"turn": 1}))) + client.event_bridge.put_nowait(((), "peer-a", parsed_route, Action.TEST(Trigger.go, data={"turn": 2}))) + + worker = client.loop.create_task(client._send_worker(writer, stop_event)) + client.loop.run_until_complete(client.message_sender_loop(writer, stop_event)) + client.loop.run_until_complete(worker) + + assert seen == [1, 2] + assert len(writer.messages) == 4 + assert b'"turn": 1' in writer.messages[0] + assert b'"part": 2' in writer.messages[1] + assert b'"turn": 2' in writer.messages[2] + finally: + client.loop.close() + + def test_send_without_use_data_keeps_legacy_dedup_behavior(): client = SummonerClient("send-data") client.flow().activate() diff --git a/tests/test_client_send_stress.py b/tests/test_client_send_stress.py new file mode 100644 index 0000000..c386f8f --- /dev/null +++ b/tests/test_client_send_stress.py @@ -0,0 +1,681 @@ +import asyncio +import json +import os +import subprocess +import sys +import textwrap +from pathlib import Path +from typing import Any, Optional + +from summoner.client import ClientMerger +from summoner.client.client import SummonerClient +from summoner.protocol import Action +from summoner.protocol.process import Direction +from summoner.protocol.triggers import load_triggers +from tests.helpers import DummyWriter + + +REPO_ROOT = Path(__file__).resolve().parents[1] + + +async def _wait_for(predicate, timeout: float = 0.5, interval: float = 0.002) -> None: + deadline = asyncio.get_running_loop().time() + timeout + while True: + if predicate(): + return + if asyncio.get_running_loop().time() >= deadline: + raise TimeoutError("condition was not met before timeout") + await asyncio.sleep(interval) + + +async def _run_sender_loop(client: SummonerClient, writer: DummyWriter, stop_event: asyncio.Event, timeout: float = 1.0) -> None: + await asyncio.wait_for(client.message_sender_loop(writer, stop_event), timeout=timeout) + + +def _configure_runtime(client: SummonerClient, *, workers: int = 1, queue_size: int = 8) -> None: + client.max_concurrent_workers = workers + client.send_queue_maxsize = queue_size + client.event_bridge_maxsize = max(queue_size, 8) + client.max_consecutive_worker_errors = 3 + client.batch_drain = True + client.send_queue = asyncio.Queue(maxsize=client.send_queue_maxsize) + client.event_bridge = asyncio.Queue(maxsize=client.event_bridge_maxsize) + + +def _event_for(client: SummonerClient, route: str, trigger: Any, *, data: Any = None): + parsed_route = client.flow().parse_route(route) + return ((1,), "tape:peer-a", parsed_route, Action.STAY(trigger, data=data)) + + +def _pythonpath_env(extra_path: Optional[Path] = None) -> dict[str, str]: + paths = [] + if extra_path is not None: + paths.append(str(extra_path)) + existing = os.environ.get("PYTHONPATH") + if existing: + paths.append(existing) + paths.append(str(REPO_ROOT)) + env = os.environ.copy() + env["PYTHONPATH"] = os.pathsep.join(paths) + return env + + +def _run_subprocess(script: Path, *, env: dict[str, str]) -> subprocess.CompletedProcess: + return subprocess.run( + [sys.executable, str(script)], + cwd=REPO_ROOT, + env=env, + capture_output=True, + text=True, + check=False, + ) + + +def test_non_reactive_timed_sender_pauses_and_resumes_without_catch_up_burst(): + client = SummonerClient("timed-pause-resume") + writer = DummyWriter() + stop_event = asyncio.Event() + allowed = True + call_times: list[float] = [] + orchestrator = None + + try: + @client.hook(Direction.SEND, priority=(1,)) + async def stop_on_second(payload: dict) -> dict: + if payload["count"] == 2: + await client.quit() + return payload + + @client.send(route="tick", every=0.01, run_while=lambda: allowed) + async def tick_sender() -> dict: + call_times.append(client.loop.time()) + return {"count": len(call_times)} + + async def orchestrate() -> None: + nonlocal allowed + await _wait_for(lambda: len(call_times) >= 1) + allowed = False + await asyncio.sleep(0.05) + allowed = True + + client.loop.run_until_complete(client._wait_for_registration()) + _configure_runtime(client, workers=1, queue_size=8) + + client._start_send_workers(writer, stop_event) + orchestrator = client.loop.create_task(orchestrate()) + client.loop.run_until_complete(_run_sender_loop(client, writer, stop_event)) + client.loop.run_until_complete(orchestrator) + client.loop.run_until_complete(client._cleanup_workers()) + + assert stop_event.is_set() + assert len(call_times) == 2 + assert len(writer.messages) == 2 + assert call_times[1] - call_times[0] >= 0.045 + finally: + if orchestrator is not None and not orchestrator.done(): + orchestrator.cancel() + client.loop.run_until_complete(asyncio.gather(orchestrator, return_exceptions=True)) + client.loop.close() + + + +def test_reactive_timed_sender_disarms_and_requires_new_event_to_restart(): + client = SummonerClient("timed-reactive-rearm") + client.flow().activate() + Trigger = load_triggers(json_dict={"go": None}) + writer = DummyWriter() + stop_event = asyncio.Event() + enabled = True + seen: list[int] = [] + orchestrator = None + + try: + @client.send( + route="request", + on_actions={Action.STAY}, + on_triggers={Trigger.go}, + every=0.01, + run_while=lambda: enabled, + ) + async def timed_sender() -> dict: + nonlocal enabled + seen.append(len(seen) + 1) + enabled = False + if len(seen) == 2: + await client.quit() + return {"count": len(seen)} + + async def orchestrate() -> None: + nonlocal enabled + await client._enqueue_sender_event(_event_for(client, "request", Trigger.go)) + await _wait_for(lambda: len(seen) >= 1) + enabled = True + await asyncio.sleep(0.04) + await client._enqueue_sender_event(_event_for(client, "request", Trigger.go)) + + client.loop.run_until_complete(client._wait_for_registration()) + _configure_runtime(client, workers=1, queue_size=8) + + client._start_send_workers(writer, stop_event) + orchestrator = client.loop.create_task(orchestrate()) + client.loop.run_until_complete(_run_sender_loop(client, writer, stop_event)) + client.loop.run_until_complete(orchestrator) + client.loop.run_until_complete(client._cleanup_workers()) + + assert stop_event.is_set() + assert seen == [1, 2] + assert len(writer.messages) >= 1 + finally: + if orchestrator is not None and not orchestrator.done(): + orchestrator.cancel() + client.loop.run_until_complete(asyncio.gather(orchestrator, return_exceptions=True)) + client.loop.close() + + + +def test_reactive_timed_duplicate_arm_does_not_reset_cadence(): + client = SummonerClient("timed-reactive-cadence") + client.flow().activate() + Trigger = load_triggers(json_dict={"go": None}) + writer = DummyWriter() + stop_event = asyncio.Event() + call_times: list[float] = [] + orchestrator = None + + try: + @client.send( + route="request", + on_actions={Action.STAY}, + on_triggers={Trigger.go}, + every=0.1, + run_while=True, + ) + async def timed_sender() -> dict: + call_times.append(client.loop.time()) + if len(call_times) == 2: + await client.quit() + return {"count": len(call_times)} + + async def orchestrate() -> None: + await client._enqueue_sender_event(_event_for(client, "request", Trigger.go)) + await _wait_for(lambda: len(call_times) >= 1) + await asyncio.sleep(0.015) + await client._enqueue_sender_event(_event_for(client, "request", Trigger.go)) + + client.loop.run_until_complete(client._wait_for_registration()) + _configure_runtime(client, workers=1, queue_size=8) + + client._start_send_workers(writer, stop_event) + orchestrator = client.loop.create_task(orchestrate()) + client.loop.run_until_complete(_run_sender_loop(client, writer, stop_event)) + client.loop.run_until_complete(orchestrator) + client.loop.run_until_complete(client._cleanup_workers()) + + assert stop_event.is_set() + assert len(call_times) == 2 + assert 0.09 <= (call_times[1] - call_times[0]) < 0.14 + finally: + if orchestrator is not None and not orchestrator.done(): + orchestrator.cancel() + client.loop.run_until_complete(asyncio.gather(orchestrator, return_exceptions=True)) + client.loop.close() + + + +def test_reactive_timed_use_data_snapshot_freezes_mutations_after_buffering(): + client = SummonerClient("timed-snapshot") + client.flow().activate() + Trigger = load_triggers(json_dict={"go": None}) + writer = DummyWriter() + stop_event = asyncio.Event() + seen: list[dict[str, Any]] = [] + payload = {"id": 2, "items": []} + orchestrator = None + + try: + @client.send(route="slow") + async def slow_sender() -> None: + await asyncio.sleep(0.05) + return None + + @client.send( + route="request", + on_actions={Action.STAY}, + on_triggers={Trigger.go}, + use_data=True, + data_mode="snapshot", + every=0.1, + run_while=True, + ) + async def timed_sender(data: dict) -> dict: + seen.append({"id": data["id"], "items": list(data["items"])}) + if len(seen) == 2: + await client.quit() + return {"id": data["id"], "items": list(data["items"])} + + async def orchestrate() -> None: + await client._enqueue_sender_event(_event_for(client, "request", Trigger.go, data={"id": 1, "items": ["seed"]})) + await _wait_for(lambda: len(seen) >= 1) + await client._enqueue_sender_event(_event_for(client, "request", Trigger.go, data=payload)) + await _wait_for( + lambda: any(len(runtime.pending_payloads) >= 1 for runtime in client.timed_sender_state.values()) + ) + payload["items"].append("mutated") + + client.loop.run_until_complete(client._wait_for_registration()) + _configure_runtime(client, workers=2, queue_size=8) + + client._start_send_workers(writer, stop_event) + orchestrator = client.loop.create_task(orchestrate()) + client.loop.run_until_complete(_run_sender_loop(client, writer, stop_event)) + client.loop.run_until_complete(orchestrator) + client.loop.run_until_complete(client._cleanup_workers()) + + assert stop_event.is_set() + assert seen == [ + {"id": 1, "items": ["seed"]}, + {"id": 2, "items": []}, + ] + finally: + if orchestrator is not None and not orchestrator.done(): + orchestrator.cancel() + client.loop.run_until_complete(asyncio.gather(orchestrator, return_exceptions=True)) + client.loop.close() + + + +def test_reactive_timed_use_data_live_observes_mutations_after_buffering(): + client = SummonerClient("timed-live") + client.flow().activate() + Trigger = load_triggers(json_dict={"go": None}) + writer = DummyWriter() + stop_event = asyncio.Event() + seen: list[dict[str, Any]] = [] + payload = {"id": 2, "items": []} + orchestrator = None + + try: + @client.send(route="slow") + async def slow_sender() -> None: + await asyncio.sleep(0.05) + return None + + @client.send( + route="request", + on_actions={Action.STAY}, + on_triggers={Trigger.go}, + use_data=True, + data_mode="live", + every=0.1, + run_while=True, + ) + async def timed_sender(data: dict) -> dict: + seen.append({"id": data["id"], "items": list(data["items"])}) + if len(seen) == 2: + await client.quit() + return {"id": data["id"], "items": list(data["items"])} + + async def orchestrate() -> None: + await client._enqueue_sender_event(_event_for(client, "request", Trigger.go, data={"id": 1, "items": ["seed"]})) + await _wait_for(lambda: len(seen) >= 1) + await client._enqueue_sender_event(_event_for(client, "request", Trigger.go, data=payload)) + await _wait_for( + lambda: any(len(runtime.pending_payloads) >= 1 for runtime in client.timed_sender_state.values()) + ) + payload["items"].append("mutated") + + client.loop.run_until_complete(client._wait_for_registration()) + _configure_runtime(client, workers=2, queue_size=8) + + client._start_send_workers(writer, stop_event) + orchestrator = client.loop.create_task(orchestrate()) + client.loop.run_until_complete(_run_sender_loop(client, writer, stop_event)) + client.loop.run_until_complete(orchestrator) + client.loop.run_until_complete(client._cleanup_workers()) + + assert stop_event.is_set() + assert seen == [ + {"id": 1, "items": ["seed"]}, + {"id": 2, "items": ["mutated"]}, + ] + finally: + if orchestrator is not None and not orchestrator.done(): + orchestrator.cancel() + client.loop.run_until_complete(asyncio.gather(orchestrator, return_exceptions=True)) + client.loop.close() + + + +def test_reactive_timed_use_data_high_volume_backpressure_preserves_all_payloads(): + client = SummonerClient("timed-stress-backpressure") + client.flow().activate() + Trigger = load_triggers(json_dict={"go": None}) + writer = DummyWriter() + stop_event = asyncio.Event() + seen: list[int] = [] + total_payloads = 20 + + try: + @client.hook(Direction.SEND, priority=(1,)) + async def stop_after_last(payload: dict) -> dict: + if payload["id"] == total_payloads - 1: + await client.quit() + return payload + + @client.send( + route="request", + on_actions={Action.STAY}, + on_triggers={Trigger.go}, + use_data=True, + data_mode="snapshot", + every=0.01, + run_while=True, + ) + async def timed_sender(data: dict) -> dict: + seen.append(data["id"]) + return {"id": data["id"]} + + client.loop.run_until_complete(client._wait_for_registration()) + _configure_runtime(client, workers=1, queue_size=2) + client.event_bridge = asyncio.Queue(maxsize=total_payloads + 1) + + async def enqueue_all() -> None: + for idx in range(total_payloads): + await client._enqueue_sender_event( + _event_for(client, "request", Trigger.go, data={"id": idx}) + ) + + client.loop.run_until_complete(enqueue_all()) + + client._start_send_workers(writer, stop_event) + client.loop.run_until_complete(_run_sender_loop(client, writer, stop_event)) + client.loop.run_until_complete(client._cleanup_workers()) + + assert stop_event.is_set() + assert seen == list(range(total_payloads)) + assert len(writer.messages) == total_payloads + assert len(client.timed_sender_state) == 1 + finally: + client.loop.close() + + + +def test_timed_sender_recovers_after_worker_exception_on_previous_tick(): + client = SummonerClient("timed-worker-recovery") + writer = DummyWriter() + stop_event = asyncio.Event() + attempts: list[str] = [] + + try: + @client.hook(Direction.SEND, priority=(1,)) + async def stop_after_success(payload: dict) -> dict: + await client.quit() + return payload + + @client.send(route="tick", every=0.01) + async def unstable_sender() -> dict: + attempts.append("tick") + if len(attempts) == 1: + raise RuntimeError("boom") + return {"ok": True, "attempt": len(attempts)} + + client.loop.run_until_complete(client._wait_for_registration()) + _configure_runtime(client, workers=1, queue_size=8) + + client._start_send_workers(writer, stop_event) + client.loop.run_until_complete(_run_sender_loop(client, writer, stop_event)) + client.loop.run_until_complete(client._cleanup_workers()) + + assert stop_event.is_set() + assert len(attempts) >= 2 + assert len(writer.messages) == 1 + finally: + client.loop.close() + + + +def test_slow_async_run_while_does_not_block_reactive_timed_arming(): + client = SummonerClient("timed-run-while-lock-scope") + client.flow().activate() + Trigger = load_triggers(json_dict={"go": None}) + writer = DummyWriter() + stop_event = asyncio.Event() + gate_started = asyncio.Event() + timed_call_times: list[float] = [] + armed_at: Optional[float] = None + orchestrator = None + + try: + async def slow_gate() -> bool: + gate_started.set() + await asyncio.sleep(0.05) + return False + + @client.hook(Direction.SEND, priority=(1,)) + async def stop_after_first(payload: dict) -> dict: + if payload["kind"] == "timed": + await client.quit() + return payload + + @client.send(route="slowtick", every=0.01, run_while=slow_gate) + async def slow_sender() -> None: + return None + + @client.send( + route="request", + on_actions={Action.STAY}, + on_triggers={Trigger.go}, + every=0.01, + run_while=True, + ) + async def timed_sender() -> dict: + timed_call_times.append(client.loop.time()) + return {"kind": "timed"} + + async def orchestrate() -> None: + nonlocal armed_at + await gate_started.wait() + armed_at = client.loop.time() + await client._enqueue_sender_event(_event_for(client, "request", Trigger.go)) + + client.loop.run_until_complete(client._wait_for_registration()) + _configure_runtime(client, workers=1, queue_size=8) + + client._start_send_workers(writer, stop_event) + orchestrator = client.loop.create_task(orchestrate()) + client.loop.run_until_complete(_run_sender_loop(client, writer, stop_event, timeout=1.5)) + client.loop.run_until_complete(orchestrator) + client.loop.run_until_complete(client._cleanup_workers()) + + assert stop_event.is_set() + assert armed_at is not None + assert len(timed_call_times) == 1 + assert timed_call_times[0] - armed_at < 0.03 + finally: + if orchestrator is not None and not orchestrator.done(): + orchestrator.cancel() + client.loop.run_until_complete(asyncio.gather(orchestrator, return_exceptions=True)) + client.loop.close() + + + +def test_untimed_sender_batches_still_wait_for_slow_peer_sender(): + client = SummonerClient("timed-batch-gating") + writer = DummyWriter() + stop_event = asyncio.Event() + fast_call_times: list[float] = [] + + try: + @client.send(route="slow") + async def slow_sender() -> dict: + await asyncio.sleep(0.05) + return {"kind": "slow"} + + @client.send(route="fast") + async def fast_sender() -> dict: + fast_call_times.append(client.loop.time()) + if len(fast_call_times) == 2: + await client.quit() + return {"kind": "fast", "count": len(fast_call_times)} + + client.loop.run_until_complete(client._wait_for_registration()) + _configure_runtime(client, workers=2, queue_size=8) + + client._start_send_workers(writer, stop_event) + client.loop.run_until_complete(_run_sender_loop(client, writer, stop_event, timeout=1.5)) + client.loop.run_until_complete(client._cleanup_workers()) + + assert stop_event.is_set() + assert len(fast_call_times) == 2 + assert fast_call_times[1] - fast_call_times[0] >= 0.04 + finally: + client.loop.close() + + + +def test_timed_scheduler_is_decoupled_from_slow_untimed_sender_batch(): + client = SummonerClient("timed-batch-decoupled") + writer = DummyWriter() + stop_event = asyncio.Event() + timed_call_times: list[float] = [] + + try: + @client.send(route="slow") + async def slow_sender() -> dict: + await asyncio.sleep(0.05) + return {"kind": "slow"} + + @client.send(route="tick", every=0.01) + async def timed_sender() -> dict: + timed_call_times.append(client.loop.time()) + if len(timed_call_times) == 2: + await client.quit() + return {"kind": "timed", "count": len(timed_call_times)} + + client.loop.run_until_complete(client._wait_for_registration()) + _configure_runtime(client, workers=2, queue_size=8) + + client._start_send_workers(writer, stop_event) + client.loop.run_until_complete(_run_sender_loop(client, writer, stop_event, timeout=1.5)) + client.loop.run_until_complete(client._cleanup_workers()) + + assert stop_event.is_set() + assert len(timed_call_times) == 2 + assert timed_call_times[1] - timed_call_times[0] < 0.035 + finally: + client.loop.close() + + + +def test_client_merger_preserves_lambda_run_while_for_imported_client_sources(): + source = SummonerClient("source-lambda-run-while") + merged = None + flag = True + + try: + @source.send(route="tick", every=0.1, run_while=lambda: flag) + async def timed_sender() -> None: + return None + + source.loop.run_until_complete(source._wait_for_registration()) + + merged = ClientMerger([source], name="merged", close_subclients=False) + merged.initiate_senders() + merged.loop.run_until_complete(merged._wait_for_registration()) + + sender = merged.sender_index["tick"][0] + assert callable(sender.run_while) + assert sender.run_while() is True + finally: + source.loop.close() + if merged is not None: + merged.loop.close() + + + +def test_run_while_json_dna_portability_matrix(tmp_path: Path): + helper_module = tmp_path / "helper_gate.py" + export_script = tmp_path / "export_timed.py" + translate_script = tmp_path / "translate_timed.py" + dna_path = tmp_path / "timed_dna.json" + + helper_module.write_text( + textwrap.dedent( + """ + def importable_gate() -> bool: + return True + """ + ) + ) + + export_script.write_text( + textwrap.dedent( + f""" + from helper_gate import importable_gate + from summoner.client.client import SummonerClient + + def main_gate() -> bool: + return True + + client = SummonerClient('portable') + + @client.send(route='importable', every=0.1, run_while=importable_gate) + async def send_importable() -> None: + return None + + @client.send(route='main', every=0.1, run_while=main_gate) + async def send_main() -> None: + return None + + @client.send(route='lambda', every=0.1, run_while=lambda: True) + async def send_lambda() -> None: + return None + + client.loop.run_until_complete(client._wait_for_registration()) + with open({str(dna_path)!r}, 'w') as f: + f.write(client.dna()) + client.loop.close() + """ + ) + ) + + translate_script.write_text( + textwrap.dedent( + f""" + import json + from summoner.client import ClientTranslation + + with open({str(dna_path)!r}) as f: + dna = json.load(f) + + results = [] + for entry in dna: + route = entry['route'] + client = ClientTranslation([entry], name=f'translate-{{route}}') + try: + client.initiate_senders() + client.loop.run_until_complete(client._wait_for_registration()) + results.append((route, 'ok')) + except Exception as e: + results.append((route, type(e).__name__)) + finally: + client.loop.close() + + print(json.dumps(results)) + """ + ) + ) + + env = _pythonpath_env(tmp_path) + export_result = _run_subprocess(export_script, env=env) + assert export_result.returncode == 0, export_result.stderr + + translate_result = _run_subprocess(translate_script, env=env) + assert translate_result.returncode == 0, translate_result.stderr + + results = dict(json.loads(translate_result.stdout.strip())) + assert results == { + "importable": "ok", + "main": "ok", + "lambda": "ValueError", + } diff --git a/tests/test_client_timed_senders.py b/tests/test_client_timed_senders.py new file mode 100644 index 0000000..d329785 --- /dev/null +++ b/tests/test_client_timed_senders.py @@ -0,0 +1,370 @@ +import asyncio +import json +from typing import Any + +import pytest + +from summoner.client import ClientTranslation +from summoner.client.client import SummonerClient +from summoner.protocol import Action +from summoner.protocol.process import Direction +from summoner.protocol.triggers import load_triggers +from tests.helpers import DummyWriter + + +def top_level_gate() -> bool: + return True + + +RUN_WHILE_FLAG = True + + +def gate_using_module_flag() -> bool: + return RUN_WHILE_FLAG + + +async def async_gate_true() -> bool: + await asyncio.sleep(0) + return True + + +def test_send_timed_validation_requires_every_for_run_while(): + client = SummonerClient("timed-validation") + + try: + with pytest.raises(ValueError): + @client.send(route="tick", run_while=True) + async def bad_sender() -> None: + return None + finally: + client.loop.close() + + +def test_send_timed_validation_requires_use_data_for_data_mode(): + client = SummonerClient("timed-validation") + + try: + with pytest.raises(ValueError): + @client.send(route="tick", data_mode="snapshot") + async def bad_sender() -> None: + return None + finally: + client.loop.close() + + +def test_send_timed_reactive_sender_requires_flow_activation(): + client = SummonerClient("timed-validation") + Trigger = load_triggers(json_dict={"go": None}) + + try: + with pytest.raises(RuntimeError): + @client.send(route="request", on_actions={Action.STAY}, on_triggers={Trigger.go}, every=0.1) + async def bad_sender() -> None: + return None + finally: + client.loop.close() + + +def test_non_reactive_timed_sender_fires_immediately_and_then_repeats(): + client = SummonerClient("timed-non-reactive") + writer = DummyWriter() + stop_event = asyncio.Event() + call_times: list[float] = [] + + try: + @client.hook(Direction.SEND, priority=(1,)) + async def stop_after_second(payload: dict) -> dict: + if payload["count"] >= 2: + await client.quit() + return payload + + @client.send(route="tick", every=0.01) + async def tick_sender() -> dict: + call_times.append(client.loop.time()) + return {"count": len(call_times)} + + client.loop.run_until_complete(client._wait_for_registration()) + + client.max_concurrent_workers = 1 + client.send_queue_maxsize = 8 + client.event_bridge_maxsize = 8 + client.max_consecutive_worker_errors = 3 + client.batch_drain = True + client.send_queue = asyncio.Queue(maxsize=client.send_queue_maxsize) + client.event_bridge = asyncio.Queue(maxsize=client.event_bridge_maxsize) + + start = client.loop.time() + client._start_send_workers(writer, stop_event) + client.loop.run_until_complete(client.message_sender_loop(writer, stop_event)) + client.loop.run_until_complete(client._cleanup_workers()) + + assert stop_event.is_set() + assert len(call_times) == 2 + assert call_times[0] - start < 0.05 + assert call_times[1] >= call_times[0] + assert len(writer.messages) == 2 + finally: + client.loop.close() + + +def test_non_reactive_timed_sender_accepts_async_run_while(): + client = SummonerClient("timed-async-run-while") + writer = DummyWriter() + stop_event = asyncio.Event() + seen: list[str] = [] + + try: + @client.hook(Direction.SEND, priority=(1,)) + async def stop_after_first(payload: dict) -> dict: + await client.quit() + return payload + + @client.send(route="tick", every=0.01, run_while=async_gate_true) + async def tick_sender() -> dict: + seen.append("tick") + return {"ok": True} + + client.loop.run_until_complete(client._wait_for_registration()) + + client.max_concurrent_workers = 1 + client.send_queue_maxsize = 8 + client.event_bridge_maxsize = 8 + client.max_consecutive_worker_errors = 3 + client.batch_drain = True + client.send_queue = asyncio.Queue(maxsize=client.send_queue_maxsize) + client.event_bridge = asyncio.Queue(maxsize=client.event_bridge_maxsize) + + client._start_send_workers(writer, stop_event) + client.loop.run_until_complete(client.message_sender_loop(writer, stop_event)) + client.loop.run_until_complete(client._cleanup_workers()) + + assert stop_event.is_set() + assert seen == ["tick"] + assert len(writer.messages) == 1 + finally: + client.loop.close() + + +def test_non_reactive_timed_multi_sender_emits_all_payloads_per_tick(): + client = SummonerClient("timed-non-reactive-multi") + writer = DummyWriter() + stop_event = asyncio.Event() + call_counts: list[int] = [] + + try: + @client.hook(Direction.SEND, priority=(1,)) + async def stop_after_second_tick(payload: dict) -> dict: + if payload["tick"] == 2 and payload["part"] == 2: + await client.quit() + return payload + + @client.send(route="tick", every=0.01, multi=True) + async def tick_sender() -> list[dict]: + tick = len(call_counts) + 1 + call_counts.append(tick) + return [ + {"tick": tick, "part": 1}, + {"tick": tick, "part": 2}, + ] + + client.loop.run_until_complete(client._wait_for_registration()) + + client.max_concurrent_workers = 1 + client.send_queue_maxsize = 8 + client.event_bridge_maxsize = 8 + client.max_consecutive_worker_errors = 3 + client.batch_drain = True + client.send_queue = asyncio.Queue(maxsize=client.send_queue_maxsize) + client.event_bridge = asyncio.Queue(maxsize=client.event_bridge_maxsize) + + client._start_send_workers(writer, stop_event) + client.loop.run_until_complete(client.message_sender_loop(writer, stop_event)) + client.loop.run_until_complete(client._cleanup_workers()) + + assert stop_event.is_set() + assert call_counts == [1, 2] + assert len(writer.messages) == 4 + finally: + client.loop.close() + + +def test_reactive_timed_use_data_buffers_multiple_payloads_for_same_tick(): + client = SummonerClient("timed-reactive-data") + client.flow().activate() + Trigger = load_triggers(json_dict={"go": None}) + writer = DummyWriter() + stop_event = asyncio.Event() + seen: list[dict[str, Any]] = [] + + try: + @client.hook(Direction.SEND, priority=(1,)) + async def stop_after_second(payload: dict) -> dict: + if payload["id"] == 2: + await client.quit() + return payload + + @client.send( + route="request", + on_actions={Action.STAY}, + on_triggers={Trigger.go}, + use_data=True, + data_mode="snapshot", + every=0.01, + run_while=True, + ) + async def timed_sender(data: dict) -> dict: + seen.append({"id": data["id"]}) + return {"id": data["id"]} + + client.loop.run_until_complete(client._wait_for_registration()) + + client.max_concurrent_workers = 1 + client.send_queue_maxsize = 8 + client.event_bridge_maxsize = 8 + client.max_consecutive_worker_errors = 3 + client.batch_drain = True + client.send_queue = asyncio.Queue(maxsize=client.send_queue_maxsize) + client.event_bridge = asyncio.Queue(maxsize=client.event_bridge_maxsize) + + parsed_route = client.flow().parse_route("request") + client.loop.run_until_complete( + client._enqueue_sender_event( + ((1,), "tape:peer-a", parsed_route, Action.STAY(Trigger.go, data={"id": 1})) + ) + ) + client.loop.run_until_complete( + client._enqueue_sender_event( + ((1,), "tape:peer-a", parsed_route, Action.STAY(Trigger.go, data={"id": 2})) + ) + ) + + client._start_send_workers(writer, stop_event) + client.loop.run_until_complete(client.message_sender_loop(writer, stop_event)) + client.loop.run_until_complete(client._cleanup_workers()) + + assert stop_event.is_set() + assert seen == [{"id": 1}, {"id": 2}] + assert len(writer.messages) == 2 + assert len(client.timed_sender_state) == 1 + finally: + client.loop.close() + + +def test_reactive_timed_multi_use_data_on_triggers_only_emits_all_payloads(): + client = SummonerClient("timed-reactive-multi-data") + client.flow().activate() + Trigger = load_triggers(json_dict={"go": None}) + writer = DummyWriter() + stop_event = asyncio.Event() + seen: list[int] = [] + + try: + @client.hook(Direction.SEND, priority=(1,)) + async def stop_after_last(payload: dict) -> dict: + if payload["turn"] == 2 and payload["part"] == 2: + await client.quit() + return payload + + @client.send( + route="request", + on_triggers={Trigger.go}, + use_data=True, + data_mode="snapshot", + every=0.01, + run_while=True, + multi=True, + ) + async def timed_sender(data: dict) -> list[dict]: + seen.append(data["turn"]) + return [ + {"turn": data["turn"], "part": 1}, + {"turn": data["turn"], "part": 2}, + ] + + client.loop.run_until_complete(client._wait_for_registration()) + + client.max_concurrent_workers = 1 + client.send_queue_maxsize = 8 + client.event_bridge_maxsize = 8 + client.max_consecutive_worker_errors = 3 + client.batch_drain = True + client.send_queue = asyncio.Queue(maxsize=client.send_queue_maxsize) + client.event_bridge = asyncio.Queue(maxsize=client.event_bridge_maxsize) + + parsed_route = client.flow().parse_route("request") + client.loop.run_until_complete( + client._enqueue_sender_event( + ((1,), "tape:peer-a", parsed_route, Action.TEST(Trigger.go, data={"turn": 1})) + ) + ) + client.loop.run_until_complete( + client._enqueue_sender_event( + ((1,), "tape:peer-a", parsed_route, Action.TEST(Trigger.go, data={"turn": 2})) + ) + ) + + client._start_send_workers(writer, stop_event) + client.loop.run_until_complete(client.message_sender_loop(writer, stop_event)) + client.loop.run_until_complete(client._cleanup_workers()) + + assert stop_event.is_set() + assert seen == [1, 2] + assert len(writer.messages) == 4 + finally: + client.loop.close() + + +def test_client_translation_replays_timed_sender_fields(): + source = SummonerClient("timed-source") + translated = None + + try: + @source.send(route="tick", every=0.5, run_while=top_level_gate) + async def timed_sender() -> None: + return None + + source.loop.run_until_complete(source._wait_for_registration()) + + dna_entries = json.loads(source.dna()) + assert dna_entries[0]["every"] == 0.5 + assert dna_entries[0]["run_while_kind"] == "callable" + assert dna_entries[0]["run_while_name"] + assert "def top_level_gate" in dna_entries[0]["run_while_source"] + + translated = ClientTranslation(dna_entries, name="timed-translated") + translated.initiate_senders() + translated.loop.run_until_complete(translated._wait_for_registration()) + + sender = translated.sender_index["tick"][0] + assert sender.every == 0.5 + assert callable(sender.run_while) + assert sender.run_while() is True + finally: + source.loop.close() + if translated is not None: + translated.loop.close() + + +def test_client_translation_replays_run_while_source_with_context_globals(): + source = SummonerClient("timed-source-context") + translated = None + + try: + @source.send(route="tick", every=0.5, run_while=gate_using_module_flag) + async def timed_sender() -> None: + return None + + source.loop.run_until_complete(source._wait_for_registration()) + + dna_entries = json.loads(source.dna(include_context=True)) + translated = ClientTranslation(dna_entries, name="timed-translated-context") + translated.initiate_senders() + translated.loop.run_until_complete(translated._wait_for_registration()) + + sender = translated.sender_index["tick"][0] + assert callable(sender.run_while) + assert sender.run_while() is True + finally: + source.loop.close() + if translated is not None: + translated.loop.close() diff --git a/tests/test_process.py b/tests/test_process.py index e85c5b9..7976e68 100644 --- a/tests/test_process.py +++ b/tests/test_process.py @@ -169,3 +169,7 @@ def test_sender_responds_to_filters(): def test_sender_use_data_defaults_to_false(): sender = Sender(fn=lambda: None, multi=False, actions=None, triggers=None) assert sender.use_data is False + assert sender.data_mode is None + assert sender.every is None + assert sender.run_while is None + assert sender.registration_id is None From 22b3495a6f0c250b7c5d8fc8d3e5a5e10632b88a Mon Sep 17 00:00:00 2001 From: Remy Tuyeras Date: Sat, 4 Apr 2026 10:20:31 -0400 Subject: [PATCH 3/4] cleanup and polish structure --- summoner/client/client.py | 54 +++++++++++++++++++-------------------- 1 file changed, 27 insertions(+), 27 deletions(-) diff --git a/summoner/client/client.py b/summoner/client/client.py index cda443a..31b1212 100644 --- a/summoner/client/client.py +++ b/summoner/client/client.py @@ -108,57 +108,57 @@ class SummonerClient: def __init__(self, name: Optional[str] = None): - # Give a name to the server + # Give the client a name self.name = name if isinstance(name, str) else "" - # Create a bare logger (no handlers yet) + # Create a bare logger with no handlers yet self.logger: Logger = get_logger(self.name) # Create a new event loop self.loop = asyncio.new_event_loop() - # Set the new loop as the current thread + # Set the current event loop for this thread asyncio.set_event_loop(self.loop) - # Protect concurrent access to the set of active tasks + # Protect access to active tasks self.active_tasks: set[asyncio.Task] = set() self.tasks_lock = asyncio.Lock() - # Protect route registration and access for receive/send functions + # Protect route registration and lookup for receivers and senders self.receiver_index: dict[str, Receiver] = {} - self.sender_index: dict[str, list[Sender]] = {} # do not use defaultdict(list) because we use .get + self.sender_index: dict[str, list[Sender]] = {} # Do not use defaultdict(list) because we rely on .get(). self.routes_lock = asyncio.Lock() - # Dynamic routing configuration (can be changed at runtime) + # Store routing information that can change at runtime self.host: Optional[str] = None self.port: Optional[int] = None - self._travel = False # Flag to signal intent to travel - self._quit = False # Flag to signal intent to shutdown the client + self._travel = False # Flag the intent to travel. + self._quit = False # Flag the intent to shut down the client. self.connection_lock = asyncio.Lock() - # Safe registration of decorators (hooks, receivers, senders) + # Track decorator registration tasks self._registration_tasks: list[asyncio.Task] = [] - # One-time indexing of parsed routes + # Cache parsed routes self.receiver_parsed_routes: dict[str, ParsedRoute] = {} self.sender_parsed_routes: dict[str, ParsedRoute] = {} - # Flow representing the underlying finite state machine + # Store the client's flow self._flow = Flow() - # Functions to read and write the flow's active states in memory + # Store callbacks that read and write active states self._upload_states: Optional[Callable[[Any], Awaitable]] = None self._download_states: Optional[Callable[[Any], Awaitable]] = None self.event_bridge_maxsize = None - self.max_concurrent_workers = None # Limit the sending rate (will use 50 if None is given) + self.max_concurrent_workers = None # Limit sender concurrency. Uses the configured default when None. self.send_queue_maxsize = None self.max_bytes_per_line = None - self.read_timeout_seconds = None # None is prefered + self.read_timeout_seconds = None # Wait indefinitely when None. self.retry_delay_seconds = None self.batch_drain = None - # Pass Event information from the receiving end to the sending end + # Pass events from receivers to senders self.event_bridge: Optional[asyncio.Queue[tuple[tuple[int, ...], Optional[str], ParsedRoute, Event]]] = None # Sender-side orchestration runtime. These structures belong to the @@ -179,8 +179,8 @@ def __init__(self, name: Optional[str] = None): self.receiving_hooks: dict[tuple[int,...], Callable[[Union[str, dict]], Union[str, dict]]] = {} self.hooks_lock = asyncio.Lock() - # ─── DNA capture for merging ───────────────────────────────────────── - # lists of dicts, each entry records one decorated handler + # Store DNA entries for cloning and merging. + # Each list records decorated handlers of one kind. self._dna_receivers: list[dict] = [] self._dna_senders: list[dict] = [] self._dna_hooks: list[dict] = [] @@ -188,7 +188,7 @@ def __init__(self, name: Optional[str] = None): self._dna_upload_states: Optional[dict] = None self._dna_download_states: Optional[dict] = None - # ==== VERSION SPECIFIC ==== + # ==== CLIENT SETUP ==== def _apply_config(self, config: dict[str,Union[str,dict[str,Union[str,dict]]]]): @@ -236,6 +236,8 @@ def initialize(self): def flow(self) -> Flow: return self._flow + # ==== CLIENT CONTROL ==== + async def travel_to(self, host, port): async with self.connection_lock: self.host = host @@ -252,6 +254,8 @@ async def _reset_client_intent(self): self._quit = False self._travel = False + # ==== STATE HOOKS ==== + def upload_states(self): """ Decorator to supply a function that returns the current state snapshot. @@ -713,7 +717,7 @@ async def register(): return decorator - # ==== DNA PROCESSING ==== + # ==== DNA EXPORT ==== def _iter_registered_handler_functions(self): """ @@ -1928,7 +1932,7 @@ async def message_sender_loop( # swallow during shutdown break - # ==== HANDLE BOTH SENDING AND RECEIVING ENDS ==== + # ==== SESSION HANDLING ==== async def handle_session(self, host: str = '127.0.0.1', port: int = 8888): """ @@ -2036,17 +2040,13 @@ def shutdown(self): task.cancel() def set_termination_signals(self): - """ - Install SIGINT/SIGTERM handlers onto the loop: - - SIGINT: interupt signal for Ctrl+C | value = 2 - - SIGTERM: system/process-based termination | value = 15 - """ + """Install SIGINT and SIGTERM handlers on the event loop.""" if platform.system() != "Windows": for sig in (signal.SIGINT, signal.SIGTERM): self.loop.add_signal_handler(sig, self.shutdown) else: def _handler(sig, frame): - # thread-safe: schedule shutdown on the event loop + # Schedule shutdown on the event loop in a thread-safe way. try: self.loop.call_soon_threadsafe(self.shutdown) except RuntimeError: From 82a846c6f8723a175b9543a24553d5e7168a7dfe Mon Sep 17 00:00:00 2001 From: Remy Tuyeras Date: Sat, 4 Apr 2026 13:21:47 -0400 Subject: [PATCH 4/4] fix: import in some tests are not compatible with sdk installation --- tests/__init__.py | 1 + tests/test_client_runtime.py | 2 +- tests/test_client_send_data.py | 2 +- tests/test_client_send_stress.py | 2 +- tests/test_client_timed_senders.py | 2 +- 5 files changed, 5 insertions(+), 4 deletions(-) create mode 100644 tests/__init__.py diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/tests/__init__.py @@ -0,0 +1 @@ + diff --git a/tests/test_client_runtime.py b/tests/test_client_runtime.py index d133004..6c76c78 100644 --- a/tests/test_client_runtime.py +++ b/tests/test_client_runtime.py @@ -7,7 +7,7 @@ from summoner.protocol.payload import wrap_with_types from summoner.protocol.process import Direction, Node from summoner.protocol.triggers import load_triggers -from tests.helpers import DummyWriter +from .helpers import DummyWriter def test_message_receiver_loop_applies_receive_hooks_updates_state_and_bridges_events(): diff --git a/tests/test_client_send_data.py b/tests/test_client_send_data.py index 05cc8de..a09258e 100644 --- a/tests/test_client_send_data.py +++ b/tests/test_client_send_data.py @@ -8,7 +8,7 @@ from summoner.protocol import Action from summoner.protocol.process import Direction from summoner.protocol.triggers import load_triggers -from tests.helpers import DummyWriter +from .helpers import DummyWriter def test_send_use_data_requires_one_argument(): diff --git a/tests/test_client_send_stress.py b/tests/test_client_send_stress.py index c386f8f..af1e572 100644 --- a/tests/test_client_send_stress.py +++ b/tests/test_client_send_stress.py @@ -12,7 +12,7 @@ from summoner.protocol import Action from summoner.protocol.process import Direction from summoner.protocol.triggers import load_triggers -from tests.helpers import DummyWriter +from .helpers import DummyWriter REPO_ROOT = Path(__file__).resolve().parents[1] diff --git a/tests/test_client_timed_senders.py b/tests/test_client_timed_senders.py index d329785..3fc8d7b 100644 --- a/tests/test_client_timed_senders.py +++ b/tests/test_client_timed_senders.py @@ -9,7 +9,7 @@ from summoner.protocol import Action from summoner.protocol.process import Direction from summoner.protocol.triggers import load_triggers -from tests.helpers import DummyWriter +from .helpers import DummyWriter def top_level_gate() -> bool: