diff --git a/src/agentex/lib/core/harness/auto_send.py b/src/agentex/lib/core/harness/auto_send.py index b645a4aae..22ab73d1b 100644 --- a/src/agentex/lib/core/harness/auto_send.py +++ b/src/agentex/lib/core/harness/auto_send.py @@ -6,6 +6,7 @@ from datetime import datetime from agentex.types.text_delta import TextDelta +from agentex.lib.utils.temporal import heartbeat_if_in_activity from agentex.types.text_content import TextContent from agentex.lib.core.harness.types import TurnUsage, TurnResult, StreamTaskMessage from agentex.lib.core.harness.tracer import SpanTracer @@ -86,6 +87,7 @@ async def _close_all() -> None: try: async for event in events: + heartbeat_if_in_activity("auto send") if deriver is not None and tracer is not None: for signal in deriver.observe(event): await tracer.handle(signal) diff --git a/src/agentex/lib/utils/temporal.py b/src/agentex/lib/utils/temporal.py index 03f22649b..ceec356b6 100644 --- a/src/agentex/lib/utils/temporal.py +++ b/src/agentex/lib/utils/temporal.py @@ -12,11 +12,26 @@ def in_temporal_workflow(): return False -def heartbeat_if_in_workflow(heartbeat_name: str): - if in_temporal_workflow(): +def heartbeat_if_in_activity(heartbeat_name: str) -> None: + """Records a Temporal activity heartbeat when called from inside an activity. + + ``activity.heartbeat`` is only legal inside an activity, so this stays a + silent no-op everywhere else: workflow code and sync agents reach the same + shared services and must not raise there. + """ + if activity.in_activity(): activity.heartbeat(heartbeat_name) +def heartbeat_if_in_workflow(heartbeat_name: str) -> None: + """Deprecated alias for :func:`heartbeat_if_in_activity`. + + Kept so the existing call sites keep working; prefer the activity-named + helper in new code. + """ + heartbeat_if_in_activity(heartbeat_name) + + def workflow_now_if_in_workflow() -> datetime | None: # Returns Temporal's deterministic workflow clock when called from inside a # workflow, otherwise None. Used to stamp messages with a monotonic diff --git a/tests/lib/core/harness/test_auto_send.py b/tests/lib/core/harness/test_auto_send.py index 8133a488c..b57298ac0 100644 --- a/tests/lib/core/harness/test_auto_send.py +++ b/tests/lib/core/harness/test_auto_send.py @@ -478,3 +478,35 @@ async def test_auto_send_created_at_forwarded(): await auto_send(_gen(events), task_id="task1", tracer=None, streaming=streaming, created_at=dt) assert all(ts == dt for ts in streaming.recorded_created_at) + + +@pytest.mark.asyncio +async def test_auto_send_heartbeats_per_event_inside_an_activity(): + """A long streamed turn must keep heartbeating, or a short heartbeat_timeout + times out an activity that is still delivering messages.""" + from temporalio import activity + from temporalio.testing import ActivityEnvironment + + streaming = _FakeStreaming() + events = [ + StreamTaskMessageStart(type="start", index=0, content=TextContent(type="text", author="agent", content="")), + StreamTaskMessageDelta(type="delta", index=0, delta=TextDelta(type="text", text_delta="a")), + StreamTaskMessageDelta(type="delta", index=0, delta=TextDelta(type="text", text_delta="b")), + StreamTaskMessageDone(type="done", index=0), + ] + beats: list[tuple[object, ...]] = [] + + @activity.defn(name="deliver_turn") + async def deliver_turn() -> None: + async def _source(): + for e in events: + yield e + + await auto_send(_source(), task_id="task1", tracer=None, streaming=streaming) + + env = ActivityEnvironment() + env.on_heartbeat = lambda *details: beats.append(details) + await env.run(deliver_turn) + + assert len(beats) == len(events) + diff --git a/tests/lib/test_temporal_utils.py b/tests/lib/test_temporal_utils.py index b06bdfb3c..a2d96fb44 100644 --- a/tests/lib/test_temporal_utils.py +++ b/tests/lib/test_temporal_utils.py @@ -2,12 +2,18 @@ from __future__ import annotations +from typing import Any from datetime import datetime from unittest.mock import patch +from temporalio import activity +from temporalio.testing import ActivityEnvironment + from agentex.lib.utils import temporal as _temporal_mod from agentex.lib.utils.temporal import ( in_temporal_workflow, + heartbeat_if_in_activity, + heartbeat_if_in_workflow, workflow_now_if_in_workflow, ) @@ -18,6 +24,45 @@ def test_in_temporal_workflow_returns_false_outside_workflow() -> None: assert in_temporal_workflow() is False +async def test_heartbeat_if_in_activity_heartbeats_inside_activity() -> None: + """The helper must reach Temporal's heartbeat channel from inside an activity.""" + recorded: list[tuple[Any, ...]] = [] + + @activity.defn(name="heartbeating_activity") + async def heartbeating_activity() -> None: + heartbeat_if_in_activity("doing slow work") + + env = ActivityEnvironment() + env.on_heartbeat = lambda *details: recorded.append(details) + + await env.run(heartbeating_activity) + + assert recorded == [("doing slow work",)] + + +async def test_heartbeat_if_in_workflow_alias_heartbeats_inside_activity() -> None: + """The legacy name stays wired to the same behaviour for existing call sites.""" + recorded: list[tuple[Any, ...]] = [] + + @activity.defn(name="legacy_heartbeating_activity") + async def legacy_heartbeating_activity() -> None: + heartbeat_if_in_workflow("doing slow work") + + env = ActivityEnvironment() + env.on_heartbeat = lambda *details: recorded.append(details) + + await env.run(legacy_heartbeating_activity) + + assert recorded == [("doing slow work",)] + + +def test_heartbeat_if_in_activity_is_noop_outside_activity() -> None: + """Sync agents and workflow code call the same helpers and must not raise there.""" + assert activity.in_activity() is False + heartbeat_if_in_activity("doing slow work") + heartbeat_if_in_workflow("doing slow work") + + def test_workflow_now_if_in_workflow_returns_none_outside_workflow() -> None: assert workflow_now_if_in_workflow() is None