|
11 | 11 | import uuid |
12 | 12 | from typing import Any |
13 | 13 | from datetime import UTC, datetime |
| 14 | +from contextlib import ExitStack |
14 | 15 | from unittest.mock import Mock, AsyncMock, patch |
15 | 16 |
|
16 | 17 | import pytest |
17 | | -from temporalio.worker import StartActivityInput |
| 18 | +from temporalio.worker import StartActivityInput, StartLocalActivityInput |
18 | 19 |
|
19 | 20 | from agentex.types.span import Span |
20 | 21 | from agentex.types.task import Task |
|
24 | 25 | from agentex.types.task_message_content import TextContent |
25 | 26 | from agentex.lib.sdk.fastacp.impl.sync_acp import SyncACP |
26 | 27 | from agentex.lib.core.clients.temporal.types import ConflictWorkflowPolicy |
| 28 | +from agentex.lib.core.temporal.workers.worker import AgentexWorker |
27 | 29 | from agentex.lib.core.temporal.services.temporal_task_service import TemporalTaskService |
28 | 30 | from agentex.lib.core.tracing.processors.sgp_tracing_processor import _sgp_metadata |
29 | 31 |
|
@@ -226,3 +228,33 @@ def test_workflow_without_memo_adds_no_header(self) -> None: |
226 | 228 | interceptor._WorkflowOutbound(next_outbound).start_activity(start_input) |
227 | 229 |
|
228 | 230 | 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