From 13bab882d2493ce125fa29dbfe079d9dd23fde7b Mon Sep 17 00:00:00 2001 From: Remy Tuyeras Date: Thu, 23 Apr 2026 18:14:18 -0400 Subject: [PATCH 1/2] add: parameter for conditional send handlers; takes a function on the input data when using use_data --- benchmarks/benchmark_use_data_send.py | 563 +++++++++++++++++++++++++ benchmarks/benchmark_when_data_send.py | 420 ++++++++++++++++++ summoner/_version.py | 2 +- summoner/client/client.py | 225 +++++++++- summoner/client/merger.py | 74 +++- summoner/protocol/payload.py | 3 +- summoner/protocol/process.py | 4 + summoner/protocol/triggers.py | 4 +- tests/test_client_send_data.py | 295 +++++++++++++ tests/test_client_timed_senders.py | 62 +++ tests/test_process.py | 1 + 11 files changed, 1628 insertions(+), 25 deletions(-) create mode 100644 benchmarks/benchmark_use_data_send.py create mode 100644 benchmarks/benchmark_when_data_send.py diff --git a/benchmarks/benchmark_use_data_send.py b/benchmarks/benchmark_use_data_send.py new file mode 100644 index 0000000..e9425af --- /dev/null +++ b/benchmarks/benchmark_use_data_send.py @@ -0,0 +1,563 @@ +import argparse +import asyncio +import copy +import json +import os +import statistics +import sys +import time + +from dataclasses import dataclass +from typing import Any + + +target_path = os.path.abspath(os.path.join(os.path.dirname(os.path.abspath(__file__)), "..")) +if target_path not in sys.path: + sys.path.insert(0, target_path) + +from summoner.client.client import SummonerClient +from summoner.protocol import Action +from summoner.protocol.payload import wrap_with_types +from summoner.protocol.triggers import load_triggers + + +@dataclass(frozen=True) +class BenchmarkCase: + schedule: str + strategy: str + + +@dataclass +class RuntimeStats: + receiver_calls: int = 0 + sender_calls: int = 0 + accepted: int = 0 + user_queue_puts: int = 0 + user_queue_gets: int = 0 + user_copy_ops: int = 0 + + +class CountingWriter: + def __init__(self) -> None: + self.messages: list[bytes] = [] + self.total_bytes = 0 + self.drain_calls = 0 + + def write(self, data: bytes) -> None: + self.messages.append(data) + self.total_bytes += len(data) + + async def drain(self) -> None: + self.drain_calls += 1 + + +def _build_payloads(*, messages: int, payload_bytes: int) -> list[dict[str, Any]]: + blob = "x" * max(0, payload_bytes) + payloads: list[dict[str, Any]] = [] + + for index in range(messages): + payloads.append( + { + "id": index, + "blob": blob, + "items": [index, index + 1, index + 2], + "meta": {"group": index % 8}, + } + ) + + return payloads + + +def _build_raw_messages(client: SummonerClient, payloads: list[dict[str, Any]]) -> list[bytes]: + raw_messages: list[bytes] = [] + + for payload in payloads: + raw_messages.append( + ( + json.dumps( + { + "remote_addr": "peer-a", + "content": json.loads( + wrap_with_types(payload, version=client.core_version) + ), + } + ) + + "\n" + ).encode() + ) + + return raw_messages + + +def _configure_runtime( + client: SummonerClient, + *, + workers: int, + queue_size: int, + bridge_size: int, + ) -> None: + client.max_concurrent_workers = workers + client.send_queue_maxsize = max(1, queue_size) + client.event_bridge_maxsize = max(1, bridge_size) + client.max_consecutive_worker_errors = 3 + client.batch_drain = True + client.read_timeout_seconds = 0.05 + client.max_bytes_per_line = 4 * 1024 * 1024 + client.send_queue = asyncio.Queue(maxsize=client.send_queue_maxsize) + client.event_bridge = asyncio.Queue(maxsize=client.event_bridge_maxsize) + + +async def _wait_for(predicate, *, timeout: float, interval: float = 0.001) -> None: + deadline = asyncio.get_running_loop().time() + timeout + while True: + if predicate(): + return + if asyncio.get_running_loop().time() >= deadline: + raise TimeoutError("benchmark case did not finish before timeout") + await asyncio.sleep(interval) + + +def _sdk_expected_copy_ops(strategy: str, messages: int) -> int: + if strategy == "sdk_use_data_snapshot": + return 2 * messages + return 0 + + +def _sender_data_mode(strategy: str) -> str: + if strategy == "sdk_use_data_snapshot": + return "snapshot" + return "live" + + +def _is_queue_shim(strategy: str) -> bool: + return strategy in { + "queue_shim_live", + "queue_shim_copy_put", + "queue_shim_snapshot_strict", + } + + +def _is_sdk_direct(strategy: str) -> bool: + return strategy in {"sdk_use_data_live", "sdk_use_data_snapshot"} + + +def _run_case_once( + case: BenchmarkCase, + *, + payloads: list[dict[str, Any]], + workers: int, + queue_size: int, + bridge_size: int, + timed_every: float, + timeout: float, + ) -> dict[str, Any]: + client = SummonerClient(name=f"bench-{case.schedule}-{case.strategy}") + client.flow().activate() + Trigger = load_triggers(json_dict={"go": None}) + stats = RuntimeStats() + writer = CountingWriter() + stop_event = asyncio.Event() + route = "request" + user_queue: asyncio.Queue[Any] | None = None + + try: + if _is_queue_shim(case.strategy): + user_queue = asyncio.Queue(maxsize=max(1, len(payloads) * 2 + 8)) + + @client.upload_states() + async def upload_states(payload: dict) -> dict: + return {payload["remote_addr"]: route} + + @client.receive(route=route) + async def receive_request(payload: dict) -> Any: + stats.receiver_calls += 1 + content = payload["content"] + + if case.strategy == "sdk_use_data_live": + return Action.STAY(Trigger.go, data=content) + + if case.strategy == "sdk_use_data_snapshot": + return Action.STAY(Trigger.go, data=content) + + assert user_queue is not None + + if case.strategy == "queue_shim_live": + await user_queue.put(content) + stats.user_queue_puts += 1 + elif case.strategy == "queue_shim_copy_put": + await user_queue.put(copy.deepcopy(content)) + stats.user_queue_puts += 1 + stats.user_copy_ops += 1 + elif case.strategy == "queue_shim_snapshot_strict": + await user_queue.put(copy.deepcopy(content)) + stats.user_queue_puts += 1 + stats.user_copy_ops += 1 + else: + raise ValueError(f"Unknown strategy {case.strategy!r}") + + # Carry a lightweight token in the event so the sender still receives + # one activation per receive, matching how agents would preserve + # event-by-event semantics without direct payload handoff. + return Action.STAY(Trigger.go, data=content["id"]) + + send_kwargs: dict[str, Any] = { + "route": route, + "on_actions": {Action.STAY}, + "on_triggers": {Trigger.go}, + "use_data": True, + "data_mode": _sender_data_mode(case.strategy), + } + if case.schedule == "timed": + send_kwargs["every"] = timed_every + send_kwargs["run_while"] = True + elif case.schedule != "untimed": + raise ValueError(f"Unknown schedule {case.schedule!r}") + + @client.send(**send_kwargs) + async def sender(data: Any) -> dict: + stats.sender_calls += 1 + + if _is_sdk_direct(case.strategy): + payload = data + elif _is_queue_shim(case.strategy): + assert user_queue is not None + try: + payload = user_queue.get_nowait() + except asyncio.QueueEmpty as exc: + raise RuntimeError( + "queue shim sender was triggered without a queued payload" + ) from exc + stats.user_queue_gets += 1 + + if case.strategy == "queue_shim_snapshot_strict": + payload = copy.deepcopy(payload) + stats.user_copy_ops += 1 + + if payload["id"] != data: + raise RuntimeError( + f"queue shim token/payload mismatch: token={data!r}, " + f"payload_id={payload['id']!r}" + ) + else: + raise ValueError(f"Unknown strategy {case.strategy!r}") + + stats.accepted += 1 + return {"id": payload["id"], "blob": payload["blob"]} + + client.loop.run_until_complete(client._wait_for_registration()) + _configure_runtime( + client, + workers=workers, + queue_size=queue_size, + bridge_size=bridge_size, + ) + + raw_messages = _build_raw_messages(client, payloads) + expected_messages = len(payloads) + + async def drive_case() -> dict[str, Any]: + reader = asyncio.StreamReader() + for raw_message in raw_messages: + reader.feed_data(raw_message) + + receiver_task = asyncio.create_task(client.message_receiver_loop(reader, stop_event)) + sender_task = asyncio.create_task(client.message_sender_loop(writer, stop_event)) + client._start_send_workers(writer, stop_event) + + started_at = time.perf_counter() + try: + await _wait_for( + lambda: ( + stats.receiver_calls >= expected_messages + and stats.sender_calls >= expected_messages + and stats.accepted >= expected_messages + and len(writer.messages) >= expected_messages + and (user_queue is None or user_queue.qsize() == 0) + ), + timeout=timeout, + ) + elapsed = time.perf_counter() - started_at + finally: + stop_event.set() + receiver_task.cancel() + sender_task.cancel() + await asyncio.gather(receiver_task, sender_task, return_exceptions=True) + await client._cleanup_workers() + + return { + "elapsed_seconds": elapsed, + "receiver_calls": stats.receiver_calls, + "sender_calls": stats.sender_calls, + "accepted": stats.accepted, + "user_queue_puts": stats.user_queue_puts, + "user_queue_gets": stats.user_queue_gets, + "user_copy_ops": stats.user_copy_ops, + "sdk_expected_copy_ops": _sdk_expected_copy_ops(case.strategy, len(payloads)), + "written_messages": len(writer.messages), + "written_bytes": writer.total_bytes, + } + + return client.loop.run_until_complete(drive_case()) + finally: + client.loop.close() + + +def run_case( + case: BenchmarkCase, + *, + messages: int, + payload_bytes: int, + rounds: int, + warmup: int, + workers: int, + queue_size: int, + bridge_size: int, + timed_every: float, + timeout: float, + ) -> dict[str, Any]: + payloads = _build_payloads(messages=messages, payload_bytes=payload_bytes) + durations: list[float] = [] + last_result: dict[str, Any] | None = None + + for round_index in range(warmup + rounds): + result = _run_case_once( + case, + payloads=payloads, + workers=workers, + queue_size=queue_size, + bridge_size=bridge_size, + timed_every=timed_every, + timeout=timeout, + ) + last_result = result + if round_index >= warmup: + durations.append(result["elapsed_seconds"]) + + assert last_result is not None + + best = min(durations) + mean = statistics.mean(durations) + + return { + "schedule": case.schedule, + "strategy": case.strategy, + "messages": len(payloads), + "best_seconds": best, + "mean_seconds": mean, + "messages_per_second": len(payloads) / best, + **last_result, + } + + +def _format_rate(value: float) -> str: + return f"{value:,.0f}" + + +def _find_result( + results: list[dict[str, Any]], + *, + schedule: str, + strategy: str, + ) -> dict[str, Any]: + return next( + result for result in results + if result["schedule"] == schedule and result["strategy"] == strategy + ) + + +def main(argv: list[str]) -> int: + parser = argparse.ArgumentParser( + description=( + "Benchmark fair agent-style receive-to-send payload handoff: " + "direct SDK use_data versus a user-managed queue shim." + ) + ) + parser.add_argument("--messages", type=int, default=20000, help="Payloads to feed into each case.") + parser.add_argument( + "--payload-bytes", + type=int, + default=256, + help="Approximate payload blob size echoed by the sender.", + ) + parser.add_argument("--rounds", type=int, default=5, help="Measured rounds per case.") + parser.add_argument("--warmup", type=int, default=1, help="Warmup rounds per case.") + parser.add_argument("--workers", type=int, default=1, help="Send workers to use.") + parser.add_argument( + "--queue-size", + type=int, + default=0, + help="Optional send queue maxsize. Default scales with messages.", + ) + parser.add_argument( + "--bridge-size", + type=int, + default=0, + help="Optional event bridge maxsize. Default scales with messages.", + ) + parser.add_argument( + "--timeout", + type=float, + default=10.0, + help="Per-round timeout in seconds.", + ) + parser.add_argument( + "--timed-every", + type=float, + default=0.01, + help="Cadence for the timed benchmark rows.", + ) + schedule_group = parser.add_mutually_exclusive_group() + schedule_group.add_argument( + "--include-timed", + action="store_true", + help="Include timed reactive sender cases in addition to untimed cases.", + ) + schedule_group.add_argument( + "--timed-only", + action="store_true", + help="Run only the timed reactive sender cases.", + ) + schedule_group.add_argument( + "--untimed-only", + action="store_true", + help="Run only the untimed sender cases. This is also the default.", + ) + args = parser.parse_args(argv) + + queue_size = args.queue_size if args.queue_size > 0 else max(8, args.messages * 2 + 8) + bridge_size = args.bridge_size if args.bridge_size > 0 else max(8, args.messages * 2 + 8) + + if args.timed_only: + schedules = ["timed"] + elif args.include_timed: + schedules = ["untimed", "timed"] + else: + schedules = ["untimed"] + + strategies = [ + "sdk_use_data_live", + "sdk_use_data_snapshot", + "queue_shim_live", + "queue_shim_copy_put", + "queue_shim_snapshot_strict", + ] + cases = [ + BenchmarkCase(schedule=schedule, strategy=strategy) + for schedule in schedules + for strategy in strategies + ] + + results = [ + run_case( + case, + messages=max(1, args.messages), + payload_bytes=max(0, args.payload_bytes), + rounds=max(1, args.rounds), + warmup=max(0, args.warmup), + workers=max(1, args.workers), + queue_size=queue_size, + bridge_size=bridge_size, + timed_every=max(0.0001, args.timed_every), + timeout=max(0.1, args.timeout), + ) + for case in cases + ] + + print( + "schedule strategy messages recv_calls sender_calls user_q_put user_q_get user_copy sdk_copy written_kib best_s mean_s msg/s" + ) + print( + "-------- ---------------------- -------- ---------- ------------ ---------- ---------- --------- -------- ----------- ------- ------- -------" + ) + for result in results: + print( + f"{result['schedule']:<8} " + f"{result['strategy']:<22} " + f"{result['messages']:>8} " + f"{result['receiver_calls']:>10} " + f"{result['sender_calls']:>12} " + f"{result['user_queue_puts']:>10} " + f"{result['user_queue_gets']:>10} " + f"{result['user_copy_ops']:>9} " + f"{result['sdk_expected_copy_ops']:>8} " + f"{(result['written_bytes'] / 1024):>11.1f} " + f"{result['best_seconds']:>7.4f} " + f"{result['mean_seconds']:>7.4f} " + f"{_format_rate(result['messages_per_second']):>7}" + ) + + print() + print("delta vs queue shims") + print( + "schedule comparison user_q_ops_saved user_copy_ops_saved best_delta_s mean_delta_s best_ratio_vs_sdk mean_ratio_vs_sdk" + ) + print( + "-------- ----------------- ---------------- ------------------- ------------ ------------ ----------------- -----------------" + ) + + for schedule in sorted({result["schedule"] for result in results}): + live_sdk = _find_result(results, schedule=schedule, strategy="sdk_use_data_live") + live_shim = _find_result(results, schedule=schedule, strategy="queue_shim_live") + snapshot_sdk = _find_result(results, schedule=schedule, strategy="sdk_use_data_snapshot") + snapshot_shim = _find_result(results, schedule=schedule, strategy="queue_shim_copy_put") + snapshot_strict_shim = _find_result( + results, + schedule=schedule, + strategy="queue_shim_snapshot_strict", + ) + + comparisons = [ + ("live", live_sdk, live_shim), + ("snapshot", snapshot_sdk, snapshot_shim), + ("snapshot_strict", snapshot_sdk, snapshot_strict_shim), + ] + + for label, sdk_result, shim_result in comparisons: + sdk_queue_ops = sdk_result["user_queue_puts"] + sdk_result["user_queue_gets"] + shim_queue_ops = shim_result["user_queue_puts"] + shim_result["user_queue_gets"] + queue_ops_saved = shim_queue_ops - sdk_queue_ops + user_copy_ops_saved = shim_result["user_copy_ops"] - sdk_result["user_copy_ops"] + best_time_saved = shim_result["best_seconds"] - sdk_result["best_seconds"] + mean_time_saved = shim_result["mean_seconds"] - sdk_result["mean_seconds"] + best_time_ratio = ( + shim_result["best_seconds"] / sdk_result["best_seconds"] + if sdk_result["best_seconds"] > 0 + else float("inf") + ) + mean_time_ratio = ( + shim_result["mean_seconds"] / sdk_result["mean_seconds"] + if sdk_result["mean_seconds"] > 0 + else float("inf") + ) + print( + f"{schedule:<8} " + f"{label:<17} " + f"{queue_ops_saved:>16} " + f"{user_copy_ops_saved:>19} " + f"{best_time_saved:>12.4f} " + f"{mean_time_saved:>12.4f} " + f"{best_time_ratio:>17.3f} " + f"{mean_time_ratio:>17.3f}" + ) + + print() + print("Interpretation:") + print("- This benchmark exercises the local path agents actually use: wrapped inbound message, receive handler, emitted action, and sender invocation.") + print("- `queue_shim_*` is the fair userland comparison: receive stores the payload in an application queue and emits a lightweight token event so the sender still runs once per received message.") + print("- `queue_shim_copy_put` models the intuitive user-level snapshot approximation: copy once on queue put, then hand the queued object to the sender.") + print("- `queue_shim_snapshot_strict` is a stricter control row: it copies on queue put and again on queue get to approximate the SDK's stronger per-delivery snapshot isolation more closely.") + print("- Sender-call counts stay aligned across direct and queue-shim rows, so timing differences are no longer caused by bulk-drain batching shortcuts.") + print("- `best_delta_s` and `mean_delta_s` are `queue_shim - sdk`, so positive values mean the SDK row was faster and negative values mean the queue shim was faster.") + print("- When `best` and `mean` disagree on direction, treat the result as parity/noise rather than a reliable win for either side.") + print("- If sender output is the same, written wire bytes should also stay the same.") + print("- The main SDK gain is avoiding user-managed queue put/get operations and extra queue bookkeeping while keeping payload handoff inside the runtime.") + print("- `sdk_use_data_snapshot` still intentionally performs internal copies for stronger delivery semantics; compare it against both snapshot queue rows, not just the intuitive one.") + if any(result["schedule"] == "timed" for result in results): + print("- Timed rows feed a burst of inbound messages and measure buffered timed handoff behavior rather than one-interval-per-message latency.") + print("- Tiny runs are timing-noisy; use larger `--messages` values when comparing elapsed time.") + + return 0 + + +if __name__ == "__main__": + raise SystemExit(main(sys.argv[1:])) diff --git a/benchmarks/benchmark_when_data_send.py b/benchmarks/benchmark_when_data_send.py new file mode 100644 index 0000000..ead07fd --- /dev/null +++ b/benchmarks/benchmark_when_data_send.py @@ -0,0 +1,420 @@ +import argparse +import asyncio +import os +import statistics +import sys +import time + +from dataclasses import dataclass +from typing import Any + + +target_path = os.path.abspath(os.path.join(os.path.dirname(os.path.abspath(__file__)), "..")) +if target_path not in sys.path: + sys.path.insert(0, target_path) + +from summoner.client.client import SummonerClient +from summoner.protocol import Action +from summoner.protocol.triggers import load_triggers + + +@dataclass(frozen=True) +class BenchmarkCase: + schedule: str + strategy: str + + +@dataclass +class RuntimeStats: + filter_evals: int = 0 + sender_calls: int = 0 + accepted: int = 0 + rejected_in_handler: int = 0 + + +class CountingWriter: + def __init__(self) -> None: + self.messages: list[bytes] = [] + self.total_bytes = 0 + self.drain_calls = 0 + + def write(self, data: bytes) -> None: + self.messages.append(data) + self.total_bytes += len(data) + + async def drain(self) -> None: + self.drain_calls += 1 + + +def _build_payloads(*, messages: int, accept_every: int, payload_bytes: int) -> list[dict[str, Any]]: + blob = "x" * max(0, payload_bytes) + payloads: list[dict[str, Any]] = [] + stride = max(1, accept_every) + + for index in range(messages): + payloads.append( + { + "id": index, + "keep": (index % stride) == 0, + "blob": blob, + } + ) + + return payloads + + +def _expected_accepted(payloads: list[dict[str, Any]]) -> int: + return sum(1 for payload in payloads if payload["keep"]) + + +def _configure_runtime( + client: SummonerClient, + *, + workers: int, + queue_size: int, + bridge_size: int, + ) -> None: + client.max_concurrent_workers = workers + client.send_queue_maxsize = max(1, queue_size) + client.event_bridge_maxsize = max(1, bridge_size) + 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) + + +async def _wait_for(predicate, *, timeout: float, interval: float = 0.001) -> None: + deadline = asyncio.get_running_loop().time() + timeout + while True: + if predicate(): + return + if asyncio.get_running_loop().time() >= deadline: + raise TimeoutError("benchmark case did not finish before timeout") + await asyncio.sleep(interval) + + +def _run_case_once( + case: BenchmarkCase, + *, + payloads: list[dict[str, Any]], + workers: int, + queue_size: int, + bridge_size: int, + timed_every: float, + timeout: float, + ) -> dict[str, Any]: + client = SummonerClient(name=f"bench-{case.schedule}-{case.strategy}") + client.flow().activate() + Trigger = load_triggers(json_dict={"go": None}) + stats = RuntimeStats() + writer = CountingWriter() + stop_event = asyncio.Event() + route = "request" + + def payload_filter(data: dict[str, Any]) -> bool: + stats.filter_evals += 1 + return bool(data["keep"]) + + try: + send_kwargs: dict[str, Any] = { + "route": route, + "on_actions": {Action.STAY}, + "on_triggers": {Trigger.go}, + "use_data": True, + "data_mode": "snapshot", + } + if case.schedule == "timed": + send_kwargs["every"] = timed_every + send_kwargs["run_while"] = True + elif case.schedule != "untimed": + raise ValueError(f"Unknown schedule {case.schedule!r}") + + if case.strategy == "sdk_when_data": + send_kwargs["when_data"] = payload_filter + + @client.send(**send_kwargs) + async def sender(data: dict[str, Any]) -> dict: + stats.sender_calls += 1 + stats.accepted += 1 + return {"id": data["id"], "blob": data["blob"]} + + elif case.strategy == "handler_guard": + @client.send(**send_kwargs) + async def sender(data: dict[str, Any]) -> Any: + stats.sender_calls += 1 + stats.filter_evals += 1 + if not data["keep"]: + stats.rejected_in_handler += 1 + return None + stats.accepted += 1 + return {"id": data["id"], "blob": data["blob"]} + + else: + raise ValueError(f"Unknown strategy {case.strategy!r}") + + client.loop.run_until_complete(client._wait_for_registration()) + _configure_runtime( + client, + workers=workers, + queue_size=queue_size, + bridge_size=bridge_size, + ) + + expected_accepted = _expected_accepted(payloads) + expected_filter_evals = len(payloads) + expected_sender_calls = ( + expected_accepted if case.strategy == "sdk_when_data" else len(payloads) + ) + + async def drive_case() -> dict[str, Any]: + message_task = asyncio.create_task(client.message_sender_loop(writer, stop_event)) + client._start_send_workers(writer, stop_event) + + started_at = time.perf_counter() + try: + parsed_route = client.flow().parse_route(route) + for payload in payloads: + await client._enqueue_sender_event( + ((1,), "tape:peer-a", parsed_route, Action.STAY(Trigger.go, data=payload)) + ) + + await _wait_for( + lambda: ( + stats.filter_evals >= expected_filter_evals + and stats.sender_calls >= expected_sender_calls + and len(writer.messages) >= expected_accepted + ), + timeout=timeout, + ) + elapsed = time.perf_counter() - started_at + finally: + message_task.cancel() + await asyncio.gather(message_task, return_exceptions=True) + await client._cleanup_workers() + + return { + "elapsed_seconds": elapsed, + "filter_evals": stats.filter_evals, + "sender_calls": stats.sender_calls, + "accepted": stats.accepted, + "rejected_in_handler": stats.rejected_in_handler, + "written_messages": len(writer.messages), + "written_bytes": writer.total_bytes, + "drain_calls": writer.drain_calls, + } + + return client.loop.run_until_complete(drive_case()) + finally: + client.loop.close() + + +def run_case( + case: BenchmarkCase, + *, + messages: int, + accept_every: int, + payload_bytes: int, + rounds: int, + warmup: int, + workers: int, + queue_size: int, + bridge_size: int, + timed_every: float, + timeout: float, + ) -> dict[str, Any]: + payloads = _build_payloads( + messages=messages, + accept_every=accept_every, + payload_bytes=payload_bytes, + ) + durations: list[float] = [] + last_result: dict[str, Any] | None = None + + for round_index in range(warmup + rounds): + result = _run_case_once( + case, + payloads=payloads, + workers=workers, + queue_size=queue_size, + bridge_size=bridge_size, + timed_every=timed_every, + timeout=timeout, + ) + last_result = result + if round_index >= warmup: + durations.append(result["elapsed_seconds"]) + + assert last_result is not None + + best = min(durations) + mean = statistics.mean(durations) + accepted = _expected_accepted(payloads) + + return { + "schedule": case.schedule, + "strategy": case.strategy, + "messages": len(payloads), + "accepted_expected": accepted, + "best_seconds": best, + "mean_seconds": mean, + "messages_per_second": len(payloads) / best, + "accepted_per_second": accepted / best if accepted else 0.0, + **last_result, + } + + +def _format_rate(value: float) -> str: + return f"{value:,.0f}" + + +def main(argv: list[str]) -> int: + parser = argparse.ArgumentParser( + description=( + "Benchmark reactive sender admission with SDK-side when_data " + "versus handler-side early return." + ) + ) + parser.add_argument("--messages", type=int, default=20000, help="Payloads to feed into each case.") + parser.add_argument( + "--accept-every", + type=int, + default=4, + help="Keep one payload out of every N. Example: 4 -> keep 25%%.", + ) + parser.add_argument( + "--payload-bytes", + type=int, + default=256, + help="Approximate payload blob size echoed by accepted sends.", + ) + parser.add_argument("--rounds", type=int, default=5, help="Measured rounds per case.") + parser.add_argument("--warmup", type=int, default=1, help="Warmup rounds per case.") + parser.add_argument("--workers", type=int, default=1, help="Send workers to use.") + parser.add_argument( + "--queue-size", + type=int, + default=0, + help="Optional send queue maxsize. Default scales with messages.", + ) + parser.add_argument( + "--bridge-size", + type=int, + default=0, + help="Optional event bridge maxsize. Default scales with messages.", + ) + parser.add_argument( + "--timed-every", + type=float, + default=0.01, + help="Cadence for the timed benchmark rows.", + ) + parser.add_argument( + "--timeout", + type=float, + default=10.0, + help="Per-round timeout in seconds.", + ) + parser.add_argument( + "--include-timed", + action="store_true", + help="Include timed reactive sender cases in addition to untimed cases.", + ) + args = parser.parse_args(argv) + + queue_size = args.queue_size if args.queue_size > 0 else max(8, args.messages * 2 + 8) + bridge_size = args.bridge_size if args.bridge_size > 0 else max(8, args.messages * 2 + 8) + + cases = [ + BenchmarkCase(schedule="untimed", strategy="sdk_when_data"), + BenchmarkCase(schedule="untimed", strategy="handler_guard"), + ] + if args.include_timed: + cases.extend( + [ + BenchmarkCase(schedule="timed", strategy="sdk_when_data"), + BenchmarkCase(schedule="timed", strategy="handler_guard"), + ] + ) + + results = [ + run_case( + case, + messages=max(1, args.messages), + accept_every=max(1, args.accept_every), + payload_bytes=max(0, args.payload_bytes), + rounds=max(1, args.rounds), + warmup=max(0, args.warmup), + workers=max(1, args.workers), + queue_size=queue_size, + bridge_size=bridge_size, + timed_every=max(0.0001, args.timed_every), + timeout=max(0.1, args.timeout), + ) + for case in cases + ] + + print( + "schedule strategy messages accepted filter_evals sender_calls written_msgs written_kib best_s mean_s msg/s accepted/s" + ) + print( + "-------- -------------- -------- -------- ------------ ------------ ------------ ----------- ------- ------- ------- ----------" + ) + for result in results: + print( + f"{result['schedule']:<8} " + f"{result['strategy']:<14} " + f"{result['messages']:>8} " + f"{result['accepted']:>8} " + f"{result['filter_evals']:>12} " + f"{result['sender_calls']:>12} " + f"{result['written_messages']:>12} " + f"{(result['written_bytes'] / 1024):>11.1f} " + f"{result['best_seconds']:>7.4f} " + f"{result['mean_seconds']:>7.4f} " + f"{_format_rate(result['messages_per_second']):>7} " + f"{_format_rate(result['accepted_per_second']):>10}" + ) + + print() + print("delta vs handler_guard") + print("schedule sender_calls_saved wire_bytes_saved best_time_saved_s best_time_ratio") + print("-------- ------------------ ---------------- ----------------- ---------------") + + for schedule in sorted({result["schedule"] for result in results}): + sdk = next( + result for result in results + if result["schedule"] == schedule and result["strategy"] == "sdk_when_data" + ) + handler = next( + result for result in results + if result["schedule"] == schedule and result["strategy"] == "handler_guard" + ) + sender_calls_saved = handler["sender_calls"] - sdk["sender_calls"] + wire_bytes_saved = handler["written_bytes"] - sdk["written_bytes"] + best_time_saved = handler["best_seconds"] - sdk["best_seconds"] + best_time_ratio = ( + handler["best_seconds"] / sdk["best_seconds"] + if sdk["best_seconds"] > 0 + else float("inf") + ) + print( + f"{schedule:<8} " + f"{sender_calls_saved:>18} " + f"{wire_bytes_saved:>16} " + f"{best_time_saved:>17.4f} " + f"{best_time_ratio:>15.3f}" + ) + + print() + print("Interpretation:") + print("- `sender_calls_saved` approximates avoided SendInvocation queueing and worker execution.") + print("- `wire_bytes_saved` is usually 0 versus handler-side early return, because both approaches suppress rejected sends.") + print("- The main win from `when_data` is local orchestration efficiency: fewer sender coroutine calls, less worker occupancy, and lower queue pressure.") + print("- Timed rows still buffer reactive payloads before admission, so their gain begins at sender-task admission rather than event buffering.") + print("- Tiny runs are timing-noisy; use larger `--messages` values when comparing elapsed time.") + + return 0 + + +if __name__ == "__main__": + raise SystemExit(main(sys.argv[1:])) diff --git a/summoner/_version.py b/summoner/_version.py index c68196d..67bc602 100644 --- a/summoner/_version.py +++ b/summoner/_version.py @@ -1 +1 @@ -__version__ = "1.2.0" +__version__ = "1.3.0" diff --git a/summoner/client/client.py b/summoner/client/client.py index 31b1212..889aef0 100644 --- a/summoner/client/client.py +++ b/summoner/client/client.py @@ -142,6 +142,8 @@ def __init__(self, name: Optional[str] = None): # Cache parsed routes self.receiver_parsed_routes: dict[str, ParsedRoute] = {} self.sender_parsed_routes: dict[str, ParsedRoute] = {} + self.snapshot_capture_sender_index: dict[str, list[Sender]] = {} + self.snapshot_capture_sender_parsed_routes: dict[str, ParsedRoute] = {} # Store the client's flow self._flow = Flow() @@ -514,40 +516,110 @@ def _normalize_data_mode(self, use_data: bool, data_mode: Optional[str]) -> Opti raise ValueError( "Argument `data_mode` must be `None`, 'live', or 'snapshot'. " f"Provided: {data_mode!r}" - ) + ) return data_mode - def _serialize_run_while_spec( + def _serialize_callable_spec( self, - run_while: Any, + spec_name: str, + spec: Any, + *, + allow_bool: bool = False, ) -> tuple[str, Optional[bool], Optional[str], Optional[str]]: - if run_while is None: + if spec 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) + if allow_bool and isinstance(spec, bool): + return ("bool", spec, None, None) + if callable(spec): + module_name = getattr(spec, "__module__", None) + qualname = getattr(spec, "__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) + fallback_name = getattr(spec, "__name__", None) if isinstance(fallback_name, str) and fallback_name: serialized_name = fallback_name try: - source = inspect.getsource(run_while) + source = inspect.getsource(spec) except Exception: - source = getattr(run_while, "__dna_source__", None) + source = getattr(spec, "__dna_source__", None) if not (isinstance(source, str) and source.strip()): source = None return ("callable", None, serialized_name, source) + if allow_bool: + allowed = "`None`, a bool, or a callable" + else: + allowed = "`None` or a callable" raise TypeError( - "Argument `run_while` must be `None`, a bool, or a callable returning " - f"bool/awaitable bool. Provided: {run_while!r}" + f"Argument `{spec_name}` must be {allowed}. Provided: {spec!r}" + ) + + def _is_async_callable(self, fn: Any) -> bool: + if inspect.iscoroutinefunction(fn): + return True + call = getattr(fn, "__call__", None) + return inspect.iscoroutinefunction(call) + + def _validate_callable_accepts_positional_args( + self, + spec_name: str, + fn: Callable[..., Any], + arg_count: int, + ) -> None: + try: + signature = inspect.signature(fn) + except (TypeError, ValueError): + return + + probe_args = [object()] * arg_count + try: + signature.bind(*probe_args) + except TypeError as e: + raise TypeError( + f"Argument `{spec_name}` callable must accept {arg_count} positional " + f"argument(s). Provided: {fn!r}" + ) from e + + def _serialize_run_while_spec( + self, + run_while: Any, + ) -> tuple[str, Optional[bool], Optional[str], Optional[str]]: + return self._serialize_callable_spec( + "run_while", + run_while, + allow_bool=True, ) + def _serialize_when_data_spec( + self, + when_data: Any, + ) -> tuple[str, Optional[bool], Optional[str], Optional[str]]: + return self._serialize_callable_spec("when_data", when_data) + + def _validate_when_data_spec( + self, + when_data: Any, + *, + use_data: bool, + ) -> None: + if when_data is None: + return + if not use_data: + raise ValueError("Argument `when_data` requires `use_data=True`") + if not callable(when_data): + raise TypeError( + "Argument `when_data` must be `None` or a callable receiving the " + f"sender payload. Provided: {when_data!r}" + ) + if self._is_async_callable(when_data): + raise TypeError( + "Argument `when_data` must be a synchronous callable returning " + f"bool. Provided: {when_data!r}" + ) + self._validate_callable_accepts_positional_args("when_data", when_data, 1) + def send( self, route: str, @@ -558,6 +630,7 @@ def send( data_mode: Optional[str] = None, every: Optional[float] = None, run_while: Any = None, + when_data: Any = None, ): if not isinstance(route, str): raise TypeError(f"Argument `route` must be string. Provided: {route}") @@ -644,8 +717,10 @@ def decorator(fn: Callable[..., Awaitable]): if data_mode is not None and not use_data: raise ValueError("Argument `data_mode` requires `use_data=True`") + self._validate_when_data_spec(when_data, use_data=use_data) 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) + when_data_kind, when_data_value, when_data_name, when_data_source = self._serialize_when_data_spec(when_data) registration_id = self._allocate_sender_registration_id() # ----[ DNA capture ]---- @@ -663,6 +738,11 @@ def decorator(fn: Callable[..., Awaitable]): "run_while_value": run_while_value, "run_while_name": run_while_name, "run_while_source": run_while_source, + "when_data": when_data, + "when_data_kind": when_data_kind, + "when_data_value": when_data_value, + "when_data_name": when_data_name, + "when_data_source": when_data_source, "source": inspect.getsource(fn), }) @@ -676,6 +756,7 @@ async def register(): triggers=on_triggers, use_data=use_data, data_mode=normalized_data_mode, + when_data=when_data, every=every, run_while=run_while, registration_id=registration_id, @@ -705,6 +786,13 @@ async def register(): if parsed_route is not None and (actions_exist or triggers_exist): self.sender_parsed_routes.setdefault(normalized_route, parsed_route) + if use_data and normalized_data_mode == "snapshot": + self.snapshot_capture_sender_index.setdefault(normalized_route, []) + self.snapshot_capture_sender_index[normalized_route].append(sender) + self.snapshot_capture_sender_parsed_routes.setdefault( + normalized_route, + parsed_route, + ) else: self.sender_index.setdefault(route, []) self.sender_index[route].append(sender) @@ -751,6 +839,9 @@ def _iter_registered_handler_functions(self): run_while = d.get("run_while") if callable(run_while): yield run_while + when_data = d.get("when_data") + if callable(when_data): + yield when_data for d in self._dna_hooks: fn = d.get("fn") @@ -926,6 +1017,10 @@ def dna(self, include_context: bool = False) -> str: "run_while_value": dna["run_while_value"], "run_while_name": dna["run_while_name"], "run_while_source": dna.get("run_while_source", None), + "when_data_kind": dna.get("when_data_kind", "none"), + "when_data_value": dna.get("when_data_value", None), + "when_data_name": dna.get("when_data_name", None), + "when_data_source": dna.get("when_data_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 [])], @@ -1285,6 +1380,33 @@ def _materialize_send_data(self, value: Any, data_mode: Optional[str]) -> Any: return copy.deepcopy(value) raise TypeError(f"Unknown data_mode: {data_mode!r}") + def _passes_when_data(self, sender: Sender, route: str, payload: Any) -> bool: + spec = sender.when_data + if spec is None: + return True + if not callable(spec): + self.logger.warning( + f"Invalid when_data specification for '{sender.fn.__name__}' " + f"on route '{route}': {spec!r}" + ) + return False + + try: + value = spec(payload) + if inspect.isawaitable(value): + self.logger.warning( + f"when_data predicate for '{sender.fn.__name__}' on route '{route}' " + "returned an awaitable; async when_data is not supported" + ) + return False + return bool(value) + except Exception as e: + self.logger.warning( + f"when_data predicate failed for '{sender.fn.__name__}' on route '{route}': " + f"{type(e).__name__}: {e}" + ) + return False + # ==== TIMED SENDER GUARD HELPERS ==== async def _await_run_while_value(self, awaitable: Any) -> bool: @@ -1347,6 +1469,12 @@ def _has_reactive_filters(sender: Sender) -> bool: (sender.triggers and isinstance(sender.triggers, set)) ) + @staticmethod + def _event_snapshot_data(event: Event) -> tuple[bool, Any]: + if getattr(event, "_has_snapshot_data", False): + return True, getattr(event, "_snapshot_data", None) + return False, None + async def _snapshot_sender_registry(self) -> tuple[dict[str, list[Sender]], dict[str, ParsedRoute]]: async with self.routes_lock: sender_index = { @@ -1356,6 +1484,15 @@ async def _snapshot_sender_registry(self) -> tuple[dict[str, list[Sender]], dict sender_parsed_routes = self.sender_parsed_routes.copy() return sender_index, sender_parsed_routes + async def _snapshot_capture_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.snapshot_capture_sender_index.items() + } + sender_parsed_routes = self.snapshot_capture_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: @@ -1370,10 +1507,50 @@ def _drain_pending_events(self) -> list[tuple[tuple[int, ...], Optional[str], Pa pending.sort(key=lambda it: hook_priority_order(it[0])) return pending + async def _needs_snapshot_capture_for_event( + self, + parsed_route: ParsedRoute, + event: Event, + ) -> bool: + if not self._flow.in_use: + return False + + sender_index, sender_parsed_routes = await self._snapshot_capture_registry() + + for route, routed_senders in sender_index.items(): + sender_parsed_route = sender_parsed_routes.get(route) + if sender_parsed_route is None: + continue + if not self._route_accepts(sender_parsed_route, parsed_route): + continue + + for sender in routed_senders: + if sender.responds_to(event): + return True + + return False + async def _enqueue_sender_event( self, item: tuple[tuple[int, ...], Optional[str], ParsedRoute, Event], ) -> None: + priority, key, parsed_route, event = item + + if ( + self._flow.in_use + and isinstance(event, Event) + and (not getattr(event, "_has_snapshot_data", False)) + and await self._needs_snapshot_capture_for_event(parsed_route, event) + ): + try: + event._snapshot_data = copy.deepcopy(event.data) + event._has_snapshot_data = True + except Exception as e: + self.logger.warning( + f"Failed to capture sender snapshot data for event on route " + f"'{parsed_route}': {type(e).__name__}: {e}" + ) + await self.event_bridge.put(item) if self._flow.in_use: @@ -1446,8 +1623,11 @@ def _arm_timed_senders_from_pending_locked( if sender.use_data: try: + has_snapshot_data, snapshot_data = self._event_snapshot_data(event) runtime.pending_payloads.append( - self._capture_send_data(event.data, sender.data_mode) + snapshot_data + if sender.data_mode == "snapshot" and has_snapshot_data + else self._capture_send_data(event.data, sender.data_mode) ) except Exception as e: self.logger.warning( @@ -1654,12 +1834,20 @@ async def _message_sender_batch_loop( if sender.use_data: try: - captured_data = self._capture_send_data(event.data, sender.data_mode) + has_snapshot_data, snapshot_data = self._event_snapshot_data(event) + captured_data = ( + snapshot_data + if sender.data_mode == "snapshot" and has_snapshot_data + else self._capture_send_data(event.data, sender.data_mode) + ) + payload = self._materialize_send_data(captured_data, sender.data_mode) + if not self._passes_when_data(sender, route, payload): + continue invocations.append( self._make_send_invocation( route, sender, - data=self._materialize_send_data(captured_data, sender.data_mode), + data=payload, track_completion=True, ) ) @@ -1810,6 +1998,9 @@ async def _message_sender_timed_loop( ) continue + if not self._passes_when_data(sender, route, payload): + continue + timed_batch.append( self._make_send_invocation( route, diff --git a/summoner/client/merger.py b/summoner/client/merger.py index d46428e..0f43296 100644 --- a/summoner/client/merger.py +++ b/summoner/client/merger.py @@ -223,7 +223,7 @@ def _resolve_callable_reference_from_source( try: if "__builtins__" not in globals_dict: globals_dict["__builtins__"] = __builtins__ - exec(compile(textwrap.dedent(source), filename="", mode="exec"), globals_dict) + exec(compile(textwrap.dedent(source), filename="", mode="exec"), globals_dict) except Exception: return None @@ -233,16 +233,19 @@ def _resolve_callable_reference_from_source( return None -def _resolve_run_while_spec( +def _resolve_callable_spec( globals_dict: dict[str, Any], kind: str, value: Any, name: Optional[str], source: Optional[str] = None, + *, + spec_name: str, + allow_bool: bool = False, ) -> Any: if kind == "none": return None - if kind == "bool": + if allow_bool and kind == "bool": return bool(value) if kind == "callable": resolved = _resolve_callable_reference(globals_dict, name) @@ -251,10 +254,45 @@ def _resolve_run_while_spec( if callable(resolved): return resolved raise ValueError( - "Could not resolve serialized run_while callable " + f"Could not resolve serialized {spec_name} callable " f"{name!r} from available replay context" ) - raise ValueError(f"Unknown run_while kind {kind!r}") + raise ValueError(f"Unknown {spec_name} kind {kind!r}") + + +def _resolve_run_while_spec( + globals_dict: dict[str, Any], + kind: str, + value: Any, + name: Optional[str], + source: Optional[str] = None, + ) -> Any: + return _resolve_callable_spec( + globals_dict, + kind, + value, + name, + source, + spec_name="run_while", + allow_bool=True, + ) + + +def _resolve_when_data_spec( + globals_dict: dict[str, Any], + kind: str, + value: Any, + name: Optional[str], + source: Optional[str] = None, + ) -> Any: + return _resolve_callable_spec( + globals_dict, + kind, + value, + name, + source, + spec_name="when_data", + ) class ClientMerger(SummonerClient): @@ -930,6 +968,15 @@ def initiate_senders(self): dna.get("run_while_name", None), dna.get("run_while_source", None), ) + when_data = dna.get("when_data") + if when_data is None: + when_data = _resolve_when_data_spec( + fn_clone.__globals__, + dna.get("when_data_kind", "none"), + dna.get("when_data_value", None), + dna.get("when_data_name", None), + dna.get("when_data_source", None), + ) self.send( route, multi=dna.get("multi", False), @@ -939,6 +986,7 @@ def initiate_senders(self): data_mode=dna.get("data_mode", None), every=dna.get("every", None), run_while=run_while, + when_data=when_data, )(fn_clone) except Exception as e: self.logger.warning( @@ -971,6 +1019,13 @@ def initiate_senders(self): entry.get("run_while_name", None), entry.get("run_while_source", None), ) + when_data = _resolve_when_data_spec( + g, + entry.get("when_data_kind", "none"), + entry.get("when_data_value", None), + entry.get("when_data_name", None), + entry.get("when_data_source", None), + ) dec = self.send( route, multi=entry.get("multi", False), @@ -980,6 +1035,7 @@ def initiate_senders(self): data_mode=entry.get("data_mode", None), every=entry.get("every", None), run_while=run_while, + when_data=when_data, ) self._apply_with_source_patch(dec, fn, entry["source"]) @@ -1341,6 +1397,13 @@ def initiate_senders(self): entry.get("run_while_name", None), entry.get("run_while_source", None), ) + when_data = _resolve_when_data_spec( + g, + entry.get("when_data_kind", "none"), + entry.get("when_data_value", None), + entry.get("when_data_name", None), + entry.get("when_data_source", None), + ) dec = self.send( route, multi=entry.get("multi", False), @@ -1350,6 +1413,7 @@ def initiate_senders(self): data_mode=entry.get("data_mode", None), every=entry.get("every", None), run_while=run_while, + when_data=when_data, ) self._apply_with_source_patch(dec, fn, entry["source"]) diff --git a/summoner/protocol/payload.py b/summoner/protocol/payload.py index 258e1d2..4e6cfb4 100644 --- a/summoner/protocol/payload.py +++ b/summoner/protocol/payload.py @@ -159,7 +159,8 @@ def cast_v0_0_1(val: Any, expected: Any) -> Any: 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.2.0", 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("1.3.0", parse_v0_0_1, cast_v0_0_1) register_envelope_version(core_version, parse_v0_0_1, cast_v0_0_1) diff --git a/summoner/protocol/process.py b/summoner/protocol/process.py index cf1f86c..8775bc6 100644 --- a/summoner/protocol/process.py +++ b/summoner/protocol/process.py @@ -313,6 +313,7 @@ class Sender: 'triggers', 'use_data', 'data_mode', + 'when_data', 'every', 'run_while', 'registration_id', @@ -323,6 +324,7 @@ class Sender: triggers: Optional[set[Signal]] use_data: bool data_mode: Optional[str] + when_data: Any every: Optional[float] run_while: Any registration_id: Optional[str] @@ -335,6 +337,7 @@ def __init__( triggers: Optional[set[Signal]], use_data: bool = False, data_mode: Optional[str] = None, + when_data: Any = None, every: Optional[float] = None, run_while: Any = None, registration_id: Optional[str] = None, @@ -345,6 +348,7 @@ def __init__( object.__setattr__(self, "triggers", triggers) object.__setattr__(self, "use_data", use_data) object.__setattr__(self, "data_mode", data_mode) + object.__setattr__(self, "when_data", when_data) object.__setattr__(self, "every", every) object.__setattr__(self, "run_while", run_while) object.__setattr__(self, "registration_id", registration_id) diff --git a/summoner/protocol/triggers.py b/summoner/protocol/triggers.py index b31a726..0fbdbfc 100644 --- a/summoner/protocol/triggers.py +++ b/summoner/protocol/triggers.py @@ -222,10 +222,12 @@ def name_of(*args): class Event: - __slots__ = ("signal", "data") + __slots__ = ("signal", "data", "_snapshot_data", "_has_snapshot_data") def __init__(self, signal: Signal, data: Any = None) -> None: self.signal = signal self.data = data + self._snapshot_data = None + self._has_snapshot_data = False def __repr__(self) -> str: if self.data is None: return f"{type(self).__name__}({self.signal!r})" diff --git a/tests/test_client_send_data.py b/tests/test_client_send_data.py index a09258e..8fc0222 100644 --- a/tests/test_client_send_data.py +++ b/tests/test_client_send_data.py @@ -11,6 +11,13 @@ from .helpers import DummyWriter +WHEN_DATA_FLAG = True + + +def when_data_uses_module_flag(data: dict) -> bool: + return WHEN_DATA_FLAG and data.get("ready", False) + + def test_send_use_data_requires_one_argument(): client = SummonerClient("send-data") client.flow().activate() @@ -53,12 +60,16 @@ async def plain_sender() -> 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 + assert client._dna_senders[0]["when_data_kind"] == "none" + assert client._dna_senders[0]["when_data_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 + assert dna_entries[0]["when_data_kind"] == "none" + assert dna_entries[0]["when_data_source"] is None finally: client.loop.close() @@ -195,6 +206,39 @@ async def bad_sender(data: dict) -> None: client.loop.close() +def test_send_when_data_requires_use_data(): + client = SummonerClient("send-data") + + try: + with pytest.raises(ValueError): + @client.send(route="request", when_data=lambda data: True) + async def bad_sender() -> None: + return None + finally: + client.loop.close() + + +def test_send_when_data_rejects_async_predicates(): + client = SummonerClient("send-data") + client.flow().activate() + + try: + async def bad_when_data(data: dict) -> bool: + return True + + with pytest.raises(TypeError): + @client.send( + route="request", + on_actions={Action.STAY}, + use_data=True, + when_data=bad_when_data, + ) + 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") @@ -345,6 +389,223 @@ async def good_sender(data: dict) -> None: client.loop.close() +def test_send_use_data_snapshot_freezes_mutations_after_enqueue(): + 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}, + on_triggers={Trigger.go}, + use_data=True, + data_mode="snapshot", + ) + async def snapshot_sender(data: dict) -> None: + seen.append({"turn": data["turn"], "items": list(data["items"])}) + await client.quit() + 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 + client.event_bridge_maxsize = 8 + + parsed_route = client.flow().parse_route("request") + payload = {"turn": 1, "items": []} + + client.loop.run_until_complete( + client._enqueue_sender_event( + ((), "peer-a", parsed_route, Action.STAY(Trigger.go, data=payload)) + ) + ) + payload["items"].append("mutated") + + 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 seen == [{"turn": 1, "items": []}] + finally: + client.loop.close() + + +def test_send_use_data_live_observes_mutations_after_enqueue(): + 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}, + on_triggers={Trigger.go}, + use_data=True, + data_mode="live", + ) + async def live_sender(data: dict) -> None: + seen.append({"turn": data["turn"], "items": list(data["items"])}) + await client.quit() + 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 + client.event_bridge_maxsize = 8 + + parsed_route = client.flow().parse_route("request") + payload = {"turn": 1, "items": []} + + client.loop.run_until_complete( + client._enqueue_sender_event( + ((), "peer-a", parsed_route, Action.STAY(Trigger.go, data=payload)) + ) + ) + payload["items"].append("mutated") + + 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 seen == [{"turn": 1, "items": ["mutated"]}] + finally: + client.loop.close() + + +def test_send_use_data_snapshot_isolates_multiple_matching_senders(): + client = SummonerClient("send-data") + client.flow().activate() + + Trigger = load_triggers(json_dict={"go": None}) + first_seen: list[list[str]] = [] + second_seen: list[list[str]] = [] + + try: + @client.send( + route="request", + on_actions={Action.STAY}, + on_triggers={Trigger.go}, + use_data=True, + data_mode="snapshot", + ) + async def first_sender(data: dict) -> None: + data["items"].append("sender-one") + first_seen.append(list(data["items"])) + return None + + @client.send( + route="request", + on_actions={Action.STAY}, + on_triggers={Trigger.go}, + use_data=True, + data_mode="snapshot", + ) + async def second_sender(data: dict) -> None: + second_seen.append(list(data["items"])) + await client.quit() + 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 + client.event_bridge_maxsize = 8 + + parsed_route = client.flow().parse_route("request") + payload = {"items": []} + + client.loop.run_until_complete( + client._enqueue_sender_event( + ((), "peer-a", parsed_route, Action.STAY(Trigger.go, data=payload)) + ) + ) + + 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 first_seen == [["sender-one"]] + assert second_seen == [[]] + finally: + client.loop.close() + + +def test_send_when_data_filters_payloads_before_sender_runs(): + 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, + when_data=lambda data: data["turn"] % 2 == 0, + ) + async def good_sender(data: dict) -> None: + seen.append(data) + 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": "skip"} + second = {"turn": 2, "text": "keep"} + + 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 == [second] + finally: + client.loop.close() + + def test_send_multi_use_data_on_triggers_only_emits_all_payloads(): client = SummonerClient("send-data") client.flow().activate() @@ -396,6 +657,40 @@ async def good_sender(data: dict) -> list[dict]: client.loop.close() +def test_client_translation_replays_when_data_source_with_context_globals(): + source = SummonerClient("source-when-data-context") + translated = None + + try: + source.flow().activate() + + @source.send( + route="request", + on_actions={Action.STAY}, + use_data=True, + when_data=when_data_uses_module_flag, + ) + async def source_sender(data: dict) -> 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="translated-when-data") + translated.flow().activate() + translated.initiate_senders() + translated.loop.run_until_complete(translated._wait_for_registration()) + + sender = translated.sender_index["request"][0] + assert callable(sender.when_data) + assert sender.when_data({"ready": True}) is True + assert sender.when_data({"ready": False}) is False + finally: + source.loop.close() + if translated is not None: + translated.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_timed_senders.py b/tests/test_client_timed_senders.py index 3fc8d7b..0fe5518 100644 --- a/tests/test_client_timed_senders.py +++ b/tests/test_client_timed_senders.py @@ -250,6 +250,68 @@ async def timed_sender(data: dict) -> dict: client.loop.close() +def test_reactive_timed_when_data_filters_buffered_payloads(): + client = SummonerClient("timed-reactive-when-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_kept_payload(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", + when_data=lambda data: data["id"] % 2 == 0, + 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()) + + 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 == [2] + assert len(writer.messages) == 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() diff --git a/tests/test_process.py b/tests/test_process.py index 7976e68..702f181 100644 --- a/tests/test_process.py +++ b/tests/test_process.py @@ -170,6 +170,7 @@ 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.when_data is None assert sender.every is None assert sender.run_while is None assert sender.registration_id is None From ff47020e4bed558d57f3cf3898df4bce9d13965a Mon Sep 17 00:00:00 2001 From: Remy Tuyeras Date: Fri, 24 Apr 2026 16:47:30 -0400 Subject: [PATCH 2/2] changed python-dotenv to 1.2.0 --- requirements.txt | 2 +- setup.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/requirements.txt b/requirements.txt index 5c56ea8..e1a4053 100644 --- a/requirements.txt +++ b/requirements.txt @@ -2,5 +2,5 @@ pytest~=9.0.3 asyncio==3.4.3 aioconsole==0.8.1 maturin==1.8.3 -python-dotenv~=1.2.2 +python-dotenv~=1.2.1 typing_extensions==4.15.0 \ No newline at end of file diff --git a/setup.py b/setup.py index ac1f159..388b1fd 100644 --- a/setup.py +++ b/setup.py @@ -15,7 +15,7 @@ python_requires=">=3.9", install_requires=[ "aioconsole==0.8.1", - "python-dotenv~=1.2.2", + "python-dotenv~=1.2.1", "typing_extensions==4.15.0; python_version < '3.13'", ], extras_require={