Skip to content

Commit a348297

Browse files
fix(tracing): stamp sgp evals ids on local activity spans too
Local activities go through start_local_activity, which the outbound interceptor did not override, so their spans carried no run or row id. Also pin the worker's interceptor wiring. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
1 parent 43211ba commit a348297

2 files changed

Lines changed: 46 additions & 4 deletions

File tree

‎src/agentex/lib/core/tracing/sgp_evals_interceptor.py‎

Lines changed: 13 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616
Interceptor,
1717
StartActivityInput,
1818
ExecuteActivityInput,
19+
StartLocalActivityInput,
1920
ActivityInboundInterceptor,
2021
WorkflowInboundInterceptor,
2122
WorkflowOutboundInterceptor,
@@ -47,14 +48,23 @@ def init(self, outbound: WorkflowOutboundInterceptor) -> None:
4748
super().init(_WorkflowOutbound(outbound))
4849

4950

51+
def _add_attrs_header(input: StartActivityInput | StartLocalActivityInput) -> None:
52+
attrs = workflow.memo_value(MEMO_KEY, default=None)
53+
if isinstance(attrs, dict) and attrs:
54+
input.headers = {**input.headers, ATTRS_HEADER: _converter.to_payload(attrs)}
55+
56+
5057
class _WorkflowOutbound(WorkflowOutboundInterceptor):
5158
@override
5259
def start_activity(self, input: StartActivityInput) -> workflow.ActivityHandle[Any]:
53-
attrs = workflow.memo_value(MEMO_KEY, default=None)
54-
if isinstance(attrs, dict) and attrs:
55-
input.headers = {**input.headers, ATTRS_HEADER: _converter.to_payload(attrs)}
60+
_add_attrs_header(input)
5661
return super().start_activity(input)
5762

63+
@override
64+
def start_local_activity(self, input: StartLocalActivityInput) -> workflow.ActivityHandle[Any]:
65+
_add_attrs_header(input)
66+
return super().start_local_activity(input)
67+
5868

5969
class _ActivityInbound(ActivityInboundInterceptor):
6070
@override

‎tests/lib/core/tracing/test_sgp_evals.py‎

Lines changed: 33 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,10 +11,11 @@
1111
import uuid
1212
from typing import Any
1313
from datetime import UTC, datetime
14+
from contextlib import ExitStack
1415
from unittest.mock import Mock, AsyncMock, patch
1516

1617
import pytest
17-
from temporalio.worker import StartActivityInput
18+
from temporalio.worker import StartActivityInput, StartLocalActivityInput
1819

1920
from agentex.types.span import Span
2021
from agentex.types.task import Task
@@ -24,6 +25,7 @@
2425
from agentex.types.task_message_content import TextContent
2526
from agentex.lib.sdk.fastacp.impl.sync_acp import SyncACP
2627
from agentex.lib.core.clients.temporal.types import ConflictWorkflowPolicy
28+
from agentex.lib.core.temporal.workers.worker import AgentexWorker
2729
from agentex.lib.core.temporal.services.temporal_task_service import TemporalTaskService
2830
from agentex.lib.core.tracing.processors.sgp_tracing_processor import _sgp_metadata
2931

@@ -226,3 +228,33 @@ def test_workflow_without_memo_adds_no_header(self) -> None:
226228
interceptor._WorkflowOutbound(next_outbound).start_activity(start_input)
227229

228230
assert sent["headers"] == {}
231+
232+
def test_local_activities_get_the_header_too(self) -> None:
233+
sent: dict[str, Any] = {}
234+
next_outbound = Mock()
235+
next_outbound.start_local_activity = lambda input: sent.update(headers=input.headers)
236+
start_input = Mock(spec=StartLocalActivityInput, headers={})
237+
238+
with patch.object(interceptor.workflow, "memo_value", return_value=EXPECTED_ATTRS):
239+
interceptor._WorkflowOutbound(next_outbound).start_local_activity(start_input)
240+
241+
assert interceptor.ATTRS_HEADER in sent["headers"]
242+
243+
async def test_worker_runs_with_the_interceptor_ahead_of_agent_interceptors(self) -> None:
244+
agent_interceptor = interceptor.SGPEvalsInterceptor()
245+
module = "agentex.lib.core.temporal.workers.worker"
246+
with ExitStack() as stack:
247+
stack.enter_context(patch(f"{module}.EnvironmentVariables"))
248+
stack.enter_context(patch(f"{module}.init_sgp_obs"))
249+
for name in ("shutdown_sgp_obs", "shutdown_default_span_queue", "shutdown_sync_tracing_processors", "get_temporal_client"):
250+
stack.enter_context(patch(f"{module}.{name}", new=AsyncMock()))
251+
worker_cls = stack.enter_context(patch(f"{module}.Worker"))
252+
worker_cls.return_value.run = AsyncMock()
253+
worker = AgentexWorker(task_queue="q", interceptors=[agent_interceptor])
254+
stack.enter_context(patch.object(worker, "start_health_check_server", new=AsyncMock()))
255+
stack.enter_context(patch.object(worker, "_register_agent", new=AsyncMock()))
256+
await worker.run(activities=[], workflow=object)
257+
258+
interceptors = worker_cls.call_args.kwargs["interceptors"]
259+
assert isinstance(interceptors[0], interceptor.SGPEvalsInterceptor)
260+
assert interceptors[1:] == [agent_interceptor]

0 commit comments

Comments
 (0)