diff --git a/src/agentex/lib/core/temporal/services/temporal_task_service.py b/src/agentex/lib/core/temporal/services/temporal_task_service.py index 774de4d9e..f13bc3eca 100644 --- a/src/agentex/lib/core/temporal/services/temporal_task_service.py +++ b/src/agentex/lib/core/temporal/services/temporal_task_service.py @@ -10,6 +10,7 @@ from agentex.types.agent import Agent from agentex.types.event import Event from agentex.protocol.acp import SendEventParams, CreateTaskParams, InterruptTaskParams +from agentex.lib.core.tracing import sgp_evals from agentex.lib.environment_variables import EnvironmentVariables from agentex.lib.core.clients.temporal.types import WorkflowState, ConflictWorkflowPolicy from agentex.lib.core.temporal.types.workflow import SignalName @@ -91,6 +92,9 @@ async def submit_task(self, agent: Agent, task: Task, params: dict[str, Any] | N execution_timeout = timedelta(seconds=timeout_seconds) if timeout_seconds and timeout_seconds > 0 else None # USE_EXISTING makes task/create idempotent # If same task ID is already running Temporal returns a handle to the existing run instead of raising WorkflowAlreadyStarted + # Eval tasks carry their span attrs in the memo so the worker can stamp them. + eval_attrs = sgp_evals.span_attrs_from_task_metadata(task.task_metadata) + memo_kwargs: dict[str, Any] = {"memo": {sgp_evals.MEMO_KEY: eval_attrs}} if eval_attrs else {} with _acp_dispatch_span("acp.task_create", task_id=task.id): return await self._temporal_client.start_workflow( workflow=self._env_vars.WORKFLOW_NAME, @@ -103,6 +107,7 @@ async def submit_task(self, agent: Agent, task: Task, params: dict[str, Any] | N task_queue=self._env_vars.WORKFLOW_TASK_QUEUE, execution_timeout=execution_timeout, conflict_policy=ConflictWorkflowPolicy.USE_EXISTING, + **memo_kwargs, ) async def get_state(self, task_id: str) -> WorkflowState: diff --git a/src/agentex/lib/core/temporal/workers/worker.py b/src/agentex/lib/core/temporal/workers/worker.py index ba8f87de5..6c2a9afbe 100644 --- a/src/agentex/lib/core/temporal/workers/worker.py +++ b/src/agentex/lib/core/temporal/workers/worker.py @@ -34,6 +34,7 @@ from agentex.lib.core.tracing.span_queue import shutdown_default_span_queue from agentex.lib.core.compat.version_guard import assert_backend_compatible from agentex.lib.core.observability.sgp_obs_setup import init_sgp_obs, shutdown_sgp_obs +from agentex.lib.core.tracing.sgp_evals_interceptor import SGPEvalsInterceptor from agentex.lib.core.tracing.tracing_processor_manager import shutdown_sync_tracing_processors logger = make_logger(__name__) @@ -274,7 +275,7 @@ async def run( build_id=str(uuid.uuid4()), debug_mode=debug_enabled, # Disable deadlock detection in debug mode # Temporal inherits client tracing before these business interceptors. - interceptors=self.interceptors, + interceptors=[SGPEvalsInterceptor(), *self.interceptors], ) logger.info(f"Starting workers for task queue: {self.task_queue}") diff --git a/src/agentex/lib/core/tracing/processors/sgp_tracing_processor.py b/src/agentex/lib/core/tracing/processors/sgp_tracing_processor.py index 9ee269231..2ec8a625c 100644 --- a/src/agentex/lib/core/tracing/processors/sgp_tracing_processor.py +++ b/src/agentex/lib/core/tracing/processors/sgp_tracing_processor.py @@ -11,7 +11,7 @@ from scale_gp_beta.lib.tracing.span import Span as SGPSpan from agentex.types.span import Span -from agentex.lib.core.tracing import code_revision +from agentex.lib.core.tracing import sgp_evals, code_revision from agentex.lib.types.tracing import SGPTracingProcessorConfig from agentex.lib.utils.logging import make_logger from agentex.lib.core.observability import tracing_metrics_recording as _metrics @@ -71,7 +71,7 @@ def _add_source_to_span(span: Span, env_vars: EnvironmentVariables) -> None: def _sgp_metadata(span: Span) -> Any: - """Metadata for the SGP write: ``span.data`` plus the opt-in commit SHA. + """Metadata for the SGP write: ``span.data`` plus the opt-in commit SHA and evals run/row ids. Returns a COPY rather than mutating ``span``. ``trace.py`` hands the same Span instance to every registered processor, so anything written onto @@ -83,13 +83,19 @@ def _sgp_metadata(span: Span) -> Any: leak like that today. Left as-is: changing five long-shipped fields is not this change's business.) """ + eval_attrs = sgp_evals.attrs_for_span(span) + extra: dict[str, Any] = dict(eval_attrs) commit_sha = code_revision.commit_sha() - if commit_sha is None: + if commit_sha is not None: + extra[code_revision.COMMIT_SHA_KEY] = commit_sha + if not extra: return span.data if isinstance(span.data, dict): - return {**span.data, code_revision.COMMIT_SHA_KEY: commit_sha} - # List-shaped data is an accepted `data` shape and has nowhere to put a - # metadata key; leave it untouched rather than dropping the caller's data. + return {**span.data, **extra} + # List-shaped data has nowhere to put a key. Eval spans must stay searchable by run and row, + # so their list moves under "data". Otherwise it is left untouched rather than reshaped. + if eval_attrs and isinstance(span.data, list): + return {**extra, "data": span.data} return span.data @@ -148,6 +154,7 @@ def on_span_start(self, span: Span) -> None: def on_span_end(self, span: Span) -> None: sgp_span = _build_sgp_span(span, self.env_vars) sgp_span.end_time = span.end_time.isoformat() # type: ignore[union-attr] + sgp_evals.release_span(span.id) sgp_span.flush(blocking=False) @override @@ -251,6 +258,9 @@ async def on_spans_end(self, spans: list[Span]) -> None: sgp_span.end_time = span.end_time.isoformat() # type: ignore[union-attr] sgp_spans.append(sgp_span) await client.spans.upsert_batch(items=[s.to_request_params() for s in sgp_spans]) + # Released only after the upload so a queue retry rebuilds the span with the same ids. + for span in spans: + sgp_evals.release_span(span.id) _metrics.record_export_success( event_type="end", span_count=len(spans), processor="sgp" ) diff --git a/src/agentex/lib/core/tracing/sgp_evals.py b/src/agentex/lib/core/tracing/sgp_evals.py new file mode 100644 index 000000000..e624edcd1 --- /dev/null +++ b/src/agentex/lib/core/tracing/sgp_evals.py @@ -0,0 +1,119 @@ +"""Run/row attribution for spans of SGP evals generation-unit tasks. + +The evals service tags each unit's task with ``task_metadata`` carrying +``sgp_evals`` plus the run, row and attempt it belongs to. The ids are copied +onto the SGP copy of every span as flat ``sgp_evals_*`` keys so a spans search +on ``extra_metadata`` finds a unit's spans. Flat because sgp-traces treats +dotted keys differently on ClickHouse and Postgres. +""" + +from __future__ import annotations + +import threading +from typing import Any +from collections import OrderedDict + +from agentex.types.span import Span + +__all__ = ( + "MEMO_KEY", + "register_task", + "register_task_metadata", + "span_attrs_from_task_metadata", + "attrs_for_span", + "capture_for_span", + "release_span", +) + +TASK_METADATA_MARKER = "sgp_evals" +SPAN_KEY_PREFIX = "sgp_evals_" +_TASK_METADATA_KEYS = ("generation_run_id", "row_id", "attempt_idx") +# Temporal workflow memo key the ACP server sets so the worker process can stamp the same attrs. +MEMO_KEY = "sgp_evals_span_attrs" + +# Only eval tasks are stored, so this stays tiny. The LRU bound caps a long-lived agent process. +_MAX_TASKS = 10_000 +_attrs_by_task: OrderedDict[str, dict[str, Any]] = OrderedDict() +# Attrs captured when a span starts, so a later registry change cannot alter a span still queued for export. +_attrs_by_span: OrderedDict[str, dict[str, Any]] = OrderedDict() +# Spans pinned with no attrs live apart so a flood of plain spans cannot evict an eval span's pinned ids. +_plain_spans: OrderedDict[str, None] = OrderedDict() +_lock = threading.Lock() + + +def span_attrs_from_task_metadata(task_metadata: Any) -> dict[str, Any] | None: + """The flat span attrs for an evals generation-unit task, or None for any other task.""" + if not isinstance(task_metadata, dict) or task_metadata.get(TASK_METADATA_MARKER) is None: + return None + attrs = { + f"{SPAN_KEY_PREFIX}{key}": task_metadata[key] for key in _TASK_METADATA_KEYS if task_metadata.get(key) is not None + } + return attrs or None + + +def register_task(task_id: str, attrs: dict[str, Any]) -> None: + with _lock: + _attrs_by_task[task_id] = attrs + _attrs_by_task.move_to_end(task_id) + while len(_attrs_by_task) > _MAX_TASKS: + _attrs_by_task.popitem(last=False) + + +def register_task_metadata(task_id: str, task_metadata: Any) -> dict[str, Any] | None: + """Remember an eval task's span attrs. No-op (and no lookup cost) for every other task.""" + attrs = span_attrs_from_task_metadata(task_metadata) + if attrs is not None: + register_task(task_id, attrs) + else: + unregister_task(task_id) + return attrs + + +def unregister_task(task_id: str) -> None: + if _attrs_by_task: + with _lock: + _attrs_by_task.pop(task_id, None) + + +def _lookup(span: Span) -> dict[str, Any]: + for key in (span.task_id, span.trace_id): + if key and key in _attrs_by_task: + return dict(_attrs_by_task[key]) + return {} + + +def capture_for_span(span: Span) -> None: + """Pin the span's attrs at start, empty included, so an eval task registered later cannot claim it.""" + with _lock: + attrs = _lookup(span) + store: OrderedDict[str, Any] = _attrs_by_span if attrs else _plain_spans + store[span.id] = attrs or None + while len(store) > _MAX_TASKS: + store.popitem(last=False) + + +def release_span(span_id: str) -> None: + if _attrs_by_span or _plain_spans: + with _lock: + _attrs_by_span.pop(span_id, None) + _plain_spans.pop(span_id, None) + + +def attrs_for_span(span: Span) -> dict[str, Any]: + """Attrs captured at span start, else those of the task found by ``span.task_id`` then ``span.trace_id``.""" + if not _attrs_by_task and not _attrs_by_span and not _plain_spans: + return {} + with _lock: + if span.id in _attrs_by_span: + return dict(_attrs_by_span[span.id]) + if span.id in _plain_spans: + return {} + return _lookup(span) + + +def clear() -> None: + """Reset the registry (test isolation).""" + with _lock: + _attrs_by_task.clear() + _attrs_by_span.clear() + _plain_spans.clear() diff --git a/src/agentex/lib/core/tracing/sgp_evals_interceptor.py b/src/agentex/lib/core/tracing/sgp_evals_interceptor.py new file mode 100644 index 000000000..5463f4328 --- /dev/null +++ b/src/agentex/lib/core/tracing/sgp_evals_interceptor.py @@ -0,0 +1,83 @@ +"""Carry an eval task's span attrs from its Temporal workflow memo into the worker's activities. + +Spans are emitted by activities in the worker process, which never sees the ACP +server's task. The ACP server puts the attrs in the workflow memo, the outbound +interceptor copies them onto each activity's headers, and the activity +interceptor registers them for the task so the SGP processor can stamp spans. +Non-eval workflows have no memo entry and get no header. +""" + +from __future__ import annotations + +from typing import Any, override + +from temporalio import activity, workflow +from temporalio.worker import ( + Interceptor, + StartActivityInput, + ExecuteActivityInput, + StartLocalActivityInput, + ActivityInboundInterceptor, + WorkflowInboundInterceptor, + WorkflowOutboundInterceptor, +) +from temporalio.converter import default + +from agentex.lib.utils.logging import make_logger +from agentex.lib.core.tracing.sgp_evals import MEMO_KEY, register_task, unregister_task + +logger = make_logger(__name__) + +ATTRS_HEADER = "sgp-evals-span-attrs" +_converter = default().payload_converter + + +class SGPEvalsInterceptor(Interceptor): + @override + def intercept_activity(self, next: ActivityInboundInterceptor) -> ActivityInboundInterceptor: + return _ActivityInbound(next) + + @override + def workflow_interceptor_class(self, input: Any) -> type[WorkflowInboundInterceptor] | None: + return _WorkflowInbound + + +class _WorkflowInbound(WorkflowInboundInterceptor): + @override + def init(self, outbound: WorkflowOutboundInterceptor) -> None: + super().init(_WorkflowOutbound(outbound)) + + +def _add_attrs_header(input: StartActivityInput | StartLocalActivityInput) -> None: + attrs = workflow.memo_value(MEMO_KEY, default=None) + if isinstance(attrs, dict) and attrs: + input.headers = {**input.headers, ATTRS_HEADER: _converter.to_payload(attrs)} + + +class _WorkflowOutbound(WorkflowOutboundInterceptor): + @override + def start_activity(self, input: StartActivityInput) -> workflow.ActivityHandle[Any]: + _add_attrs_header(input) + return super().start_activity(input) + + @override + def start_local_activity(self, input: StartLocalActivityInput) -> workflow.ActivityHandle[Any]: + _add_attrs_header(input) + return super().start_local_activity(input) + + +class _ActivityInbound(ActivityInboundInterceptor): + @override + async def execute_activity(self, input: ExecuteActivityInput) -> Any: + payload = input.headers.get(ATTRS_HEADER) + # The workflow id is the task id (see TemporalTaskService.submit_task). + task_id = activity.info().workflow_id + if task_id and payload is None: + # A reused workflow id must not inherit a previous eval run's ids. + unregister_task(task_id) + elif task_id and payload is not None: + try: + register_task(task_id, _converter.from_payload(payload, dict)) + except Exception: + logger.warning("failed to read sgp evals span attrs from activity headers", exc_info=True) + return await super().execute_activity(input) diff --git a/src/agentex/lib/core/tracing/trace.py b/src/agentex/lib/core/tracing/trace.py index 8f5260913..0871905e6 100644 --- a/src/agentex/lib/core/tracing/trace.py +++ b/src/agentex/lib/core/tracing/trace.py @@ -10,6 +10,7 @@ from agentex import Agentex, AsyncAgentex from agentex.types.span import Span +from agentex.lib.core.tracing import sgp_evals from agentex.lib.utils.logging import make_logger from agentex.lib.utils.model_utils import recursive_model_dump from agentex.lib.core.tracing.obs_ids import obs_correlation, warn_on_backend_drift @@ -293,6 +294,7 @@ def start_span( if obs_handle is not None: _register_obs_handle(span.id, obs_handle) + sgp_evals.capture_for_span(span) for processor in self.processors: _run_on_span_start(processor, span) @@ -457,6 +459,8 @@ async def start_span( if obs_handle is not None: _register_obs_handle(span.id, obs_handle) + sgp_evals.capture_for_span(span) + # Enqueueing the START event must not crash the app path either (same # principle as _run_on_span_start): swallow so start_span still returns # and end_span cleans up the handle. The processors' on_span_start runs diff --git a/src/agentex/lib/sdk/fastacp/base/base_acp_server.py b/src/agentex/lib/sdk/fastacp/base/base_acp_server.py index 50c304c92..bec8ca787 100644 --- a/src/agentex/lib/sdk/fastacp/base/base_acp_server.py +++ b/src/agentex/lib/sdk/fastacp/base/base_acp_server.py @@ -24,6 +24,7 @@ SendMessageParams, InterruptTaskParams, ) +from agentex.lib.core.tracing import sgp_evals from agentex.lib.utils.logging import make_logger, ctx_var_request_id from agentex.protocol.json_rpc import JSONRPCError, JSONRPCRequest, JSONRPCResponse from agentex.lib.utils.model_utils import BaseModel @@ -351,6 +352,10 @@ async def _handle_jsonrpc(self, request: Request): params_data["request"] = {"headers": custom_headers} params = params_model.model_validate(params_data) + task = getattr(params, "task", None) + if task is not None: + sgp_evals.register_task_metadata(task.id, task.task_metadata) + if method in RPC_SYNC_METHODS: handler = self._handlers[method] result = await handler(params) diff --git a/tests/lib/core/tracing/test_sgp_evals.py b/tests/lib/core/tracing/test_sgp_evals.py new file mode 100644 index 000000000..85337b754 --- /dev/null +++ b/tests/lib/core/tracing/test_sgp_evals.py @@ -0,0 +1,397 @@ +"""Spans of an SGP evals generation-unit task carry flat ``sgp_evals_*`` keys on the SGP copy. + +Expectations come from the contract with sgp-evaluations: the task is created with +``task_metadata = {"sgp_evals": "generation-unit", "generation_run_id", "row_id", "attempt_idx"}`` +and its spans must be findable by ``sgp_evals_generation_run_id``, ``sgp_evals_row_id`` and +``sgp_evals_attempt_idx``. +""" + +from __future__ import annotations + +import uuid +from typing import Any +from datetime import UTC, datetime +from contextlib import ExitStack +from unittest.mock import Mock, AsyncMock, patch + +import pytest +from temporalio.worker import StartActivityInput, StartLocalActivityInput + +from agentex.types.span import Span +from agentex.types.task import Task +from agentex.types.agent import Agent +from agentex.protocol.acp import RPCMethod, CreateTaskParams, SendMessageParams +from agentex.lib.core.tracing import sgp_evals, sgp_evals_interceptor as interceptor +from agentex.lib.types.tracing import SGPTracingProcessorConfig +from agentex.lib.core.tracing.trace import Trace, AsyncTrace +from agentex.types.task_message_content import TextContent +from agentex.lib.sdk.fastacp.impl.sync_acp import SyncACP +from agentex.lib.core.clients.temporal.types import ConflictWorkflowPolicy +from agentex.lib.core.temporal.workers.worker import AgentexWorker +from agentex.lib.core.temporal.services.temporal_task_service import TemporalTaskService +from agentex.lib.core.tracing.processors.sgp_tracing_processor import ( + SGPSyncTracingProcessor, + SGPAsyncTracingProcessor, + _sgp_metadata, +) + +EVAL_METADATA = { + "sgp_evals": "generation-unit", + "generation_run_id": "run-42", + "row_id": "row-7", + "attempt_idx": 2, +} +EXPECTED_ATTRS = { + "sgp_evals_generation_run_id": "run-42", + "sgp_evals_row_id": "row-7", + "sgp_evals_attempt_idx": 2, +} + + +@pytest.fixture(autouse=True) +def _clean_registry(): + sgp_evals.clear() + yield + sgp_evals.clear() + + +def _span(trace_id: str = "task-1", task_id: str | None = None, data: Any = None) -> Span: + return Span( + id=str(uuid.uuid4()), + name="s", + start_time=datetime.now(UTC), + trace_id=trace_id, + task_id=task_id, + data=data, + ) + + +def _agent() -> Agent: + return Agent( + id="a1", + name="a", + description="a", + acp_type="async", + created_at="2023-01-01T00:00:00Z", + updated_at="2023-01-01T00:00:00Z", + ) + + +class TestAttrsFromTaskMetadata: + def test_eval_task_yields_flat_keys(self) -> None: + assert sgp_evals.span_attrs_from_task_metadata(EVAL_METADATA) == EXPECTED_ATTRS + + @pytest.mark.parametrize("metadata", [None, {}, {"generation_run_id": "run-42"}, {"other": 1}, "x"]) + def test_non_eval_task_yields_nothing(self, metadata: Any) -> None: + assert sgp_evals.span_attrs_from_task_metadata(metadata) is None + + +class TestSGPMetadata: + def test_registered_task_stamps_sgp_copy_not_span_data(self) -> None: + sgp_evals.register_task_metadata("task-1", EVAL_METADATA) + span = _span(data={"caller": "kept"}) + + metadata = _sgp_metadata(span) + + assert metadata == {"caller": "kept", **EXPECTED_ATTRS} + assert span.data == {"caller": "kept"} + + def test_matches_on_span_task_id_when_trace_id_differs(self) -> None: + sgp_evals.register_task_metadata("task-1", EVAL_METADATA) + assert _sgp_metadata(_span(trace_id="other", task_id="task-1", data={}))["sgp_evals_row_id"] == "row-7" + + def test_other_tasks_and_non_eval_tasks_are_untouched(self) -> None: + sgp_evals.register_task_metadata("task-1", EVAL_METADATA) + sgp_evals.register_task_metadata("task-2", {"unrelated": True}) + assert _sgp_metadata(_span(trace_id="task-2", data={"k": 1})) == {"k": 1} + assert _sgp_metadata(_span(trace_id="task-3", data={"k": 1})) == {"k": 1} + + def test_eval_span_with_list_data_keeps_list_and_gains_the_ids(self) -> None: + sgp_evals.register_task_metadata("task-1", EVAL_METADATA) + assert _sgp_metadata(_span(data=[{"a": 1}])) == {**EXPECTED_ATTRS, "data": [{"a": 1}]} + + def test_task_that_stops_being_an_eval_task_stops_being_stamped(self) -> None: + sgp_evals.register_task_metadata("task-1", EVAL_METADATA) + sgp_evals.register_task_metadata("task-1", {"team": "x"}) + assert _sgp_metadata(_span(data={"k": 1})) == {"k": 1} + + def test_registry_is_bounded(self) -> None: + with patch.object(sgp_evals, "_MAX_TASKS", 2): + for i in range(3): + sgp_evals.register_task_metadata(f"t{i}", EVAL_METADATA) + assert sgp_evals.attrs_for_span(_span(trace_id="t0")) == {} + assert sgp_evals.attrs_for_span(_span(trace_id="t2")) == EXPECTED_ATTRS + + +class TestQueuedSpans: + """A span keeps the ids its task had when it started, even if the task id is reused before export.""" + + def test_captured_ids_survive_the_task_being_unregistered(self) -> None: + sgp_evals.register_task_metadata("task-1", EVAL_METADATA) + span = _span(data={}) + sgp_evals.capture_for_span(span) + + sgp_evals.unregister_task("task-1") + + assert _sgp_metadata(span) == EXPECTED_ATTRS + + def test_captured_ids_survive_a_new_run_taking_the_task_id(self) -> None: + sgp_evals.register_task_metadata("task-1", EVAL_METADATA) + span = _span(data={}) + sgp_evals.capture_for_span(span) + + sgp_evals.register_task("task-1", {"sgp_evals_row_id": "row-99"}) + + assert _sgp_metadata(span) == EXPECTED_ATTRS + + def test_plain_span_is_not_stamped_by_a_run_registered_after_it_started(self) -> None: + sgp_evals.register_task_metadata("other", EVAL_METADATA) + span = _span(trace_id="task-1", data={"k": 1}) + sgp_evals.capture_for_span(span) + + sgp_evals.register_task_metadata("task-1", EVAL_METADATA) + + assert _sgp_metadata(span) == {"k": 1} + + def test_plain_span_started_with_an_empty_registry_is_not_stamped_later(self) -> None: + span = _span(trace_id="task-1", data={"k": 1}) + sgp_evals.capture_for_span(span) + + sgp_evals.register_task_metadata("task-1", EVAL_METADATA) + + assert _sgp_metadata(span) == {"k": 1} + + async def test_async_processor_keeps_the_capture_when_the_upload_fails(self) -> None: + sgp_evals.register_task_metadata("task-1", EVAL_METADATA) + span = _span(data={}) + sgp_evals.capture_for_span(span) + span.end_time = datetime.now(UTC) + module = "agentex.lib.core.tracing.processors.sgp_tracing_processor" + with patch(f"{module}.tracing.init"), patch(f"{module}.EnvironmentVariables"): + processor = SGPAsyncTracingProcessor(SGPTracingProcessorConfig(sgp_api_key="", sgp_account_id="")) + client = Mock() + client.spans.upsert_batch = AsyncMock(side_effect=[RuntimeError("boom"), None]) + with patch.object(processor, "_get_client", return_value=client): + with pytest.raises(RuntimeError): + await processor.on_spans_end([span]) + sgp_evals.unregister_task("task-1") + assert sgp_evals.attrs_for_span(span) == EXPECTED_ATTRS + + await processor.on_spans_end([span]) + + assert sgp_evals.attrs_for_span(span) == {} + + def test_release_drops_the_capture(self) -> None: + sgp_evals.register_task_metadata("task-1", EVAL_METADATA) + span = _span(data={}) + sgp_evals.capture_for_span(span) + sgp_evals.release_span(span.id) + sgp_evals.unregister_task("task-1") + + assert sgp_evals.attrs_for_span(span) == {} + + def test_captures_are_bounded(self) -> None: + sgp_evals.register_task_metadata("task-1", EVAL_METADATA) + spans = [_span(data={}) for _ in range(3)] + with patch.object(sgp_evals, "_MAX_TASKS", 2): + for span in spans: + sgp_evals.capture_for_span(span) + sgp_evals.unregister_task("task-1") + + assert sgp_evals.attrs_for_span(spans[0]) == {} + assert sgp_evals.attrs_for_span(spans[2]) == EXPECTED_ATTRS + + def test_plain_spans_cannot_evict_an_eval_span_capture(self) -> None: + sgp_evals.register_task_metadata("task-1", EVAL_METADATA) + eval_span = _span(data={}) + sgp_evals.capture_for_span(eval_span) + sgp_evals.unregister_task("task-1") + with patch.object(sgp_evals, "_MAX_TASKS", 2): + for _ in range(5): + sgp_evals.capture_for_span(_span(data={})) + + assert sgp_evals.attrs_for_span(eval_span) == EXPECTED_ATTRS + + def test_sync_trace_start_span_captures_ids(self) -> None: + sgp_evals.register_task_metadata("task-1", EVAL_METADATA) + span = Trace(processors=[], client=Mock(), trace_id="task-1").start_span(name="s") + + sgp_evals.unregister_task("task-1") + span.data = {} + + assert _sgp_metadata(span) == EXPECTED_ATTRS + + async def test_async_trace_start_span_captures_ids(self) -> None: + sgp_evals.register_task_metadata("task-1", EVAL_METADATA) + span = await AsyncTrace(processors=[], client=Mock(), trace_id="task-1", span_queue=Mock()).start_span(name="s") + + sgp_evals.unregister_task("task-1") + span.data = {} + + assert _sgp_metadata(span) == EXPECTED_ATTRS + + def test_sgp_processor_releases_the_capture_when_the_span_ends(self) -> None: + sgp_evals.register_task_metadata("task-1", EVAL_METADATA) + span = _span(data={}) + sgp_evals.capture_for_span(span) + span.end_time = datetime.now(UTC) + module = "agentex.lib.core.tracing.processors.sgp_tracing_processor" + with patch(f"{module}.tracing.init"), patch(f"{module}.EnvironmentVariables"): + processor = SGPSyncTracingProcessor(SGPTracingProcessorConfig(sgp_api_key="", sgp_account_id="")) + with patch(f"{module}._build_sgp_span", return_value=Mock()): + processor.on_span_end(span) + sgp_evals.unregister_task("task-1") + + assert sgp_evals.attrs_for_span(span) == {} + + +class TestAcpServer: + """Sync and async (base) agents: the ACP server sees the task and spans are emitted in-process.""" + + async def _dispatch(self, acp: SyncACP, method: RPCMethod, params: Any) -> Any: + class _Req: + headers: dict[str, str] = {} + + async def json(self) -> dict[str, Any]: + return {"jsonrpc": "2.0", "method": method.value, "params": params.model_dump(mode="json"), "id": "1"} + + return await acp._handle_jsonrpc(_Req()) # pyright: ignore[reportArgumentType] + + async def test_task_create_registers_eval_task(self) -> None: + acp = SyncACP() + seen: list[Any] = [] + + @acp.on_task_create + async def handler(params: CreateTaskParams) -> None: + seen.append(_sgp_metadata(_span(trace_id=params.task.id, data={}))) + + task = Task(id="task-1", status="RUNNING", task_metadata=EVAL_METADATA) + await self._dispatch(acp, RPCMethod.TASK_CREATE, CreateTaskParams(agent=_agent(), task=task, params=None)) + + assert seen == [EXPECTED_ATTRS] + + async def test_message_send_registers_eval_task(self) -> None: + acp = SyncACP() + seen: list[Any] = [] + + @acp.on_message_send + async def handler(params: SendMessageParams) -> Any: + seen.append(_sgp_metadata(_span(trace_id=params.task.id, data={}))) + return None + + task = Task(id="task-1", status="RUNNING", task_metadata=EVAL_METADATA) + params = SendMessageParams(agent=_agent(), task=task, content=TextContent(author="user", content="hi"), stream=False) + await self._dispatch(acp, RPCMethod.MESSAGE_SEND, params) + + assert seen == [EXPECTED_ATTRS] + + async def test_non_eval_task_is_not_registered(self) -> None: + acp = SyncACP() + + @acp.on_task_create + async def handler(params: CreateTaskParams) -> None: + return None + + task = Task(id="task-1", status="RUNNING", task_metadata={"team": "x"}) + await self._dispatch(acp, RPCMethod.TASK_CREATE, CreateTaskParams(agent=_agent(), task=task, params=None)) + + assert sgp_evals.attrs_for_span(_span(trace_id="task-1")) == {} + + +def _env_vars() -> Mock: + env_vars = Mock() + env_vars.WORKFLOW_NAME = "wf" + env_vars.WORKFLOW_TASK_QUEUE = "q" + env_vars.WORKFLOW_EXECUTION_TIMEOUT_SECONDS = 0 + return env_vars + + +class TestTemporal: + """Temporal agents: attrs ride the workflow memo, then activity headers, into the worker process.""" + + async def _submit(self, task_metadata: dict[str, Any] | None) -> dict[str, Any]: + client = Mock() + client.start_workflow = AsyncMock(return_value="task-1") + service = TemporalTaskService(temporal_client=client, env_vars=_env_vars()) + await service.submit_task( + agent=_agent(), task=Task(id="task-1", task_metadata=task_metadata), params=None + ) + return client.start_workflow.await_args.kwargs + + async def test_eval_task_workflow_starts_with_attrs_in_memo(self) -> None: + kwargs = await self._submit(EVAL_METADATA) + assert kwargs["memo"] == {"sgp_evals_span_attrs": EXPECTED_ATTRS} + assert kwargs["conflict_policy"] == ConflictWorkflowPolicy.USE_EXISTING + + async def test_non_eval_task_workflow_has_no_memo(self) -> None: + assert "memo" not in await self._submit({"team": "x"}) + assert "memo" not in await self._submit(None) + + async def test_memo_flows_through_activity_headers_to_worker_registry(self) -> None: + sent: dict[str, Any] = {} + next_outbound = Mock() + next_outbound.start_activity = lambda input: sent.update(headers=input.headers) + outbound = interceptor._WorkflowOutbound(next_outbound) + start_input = Mock(spec=StartActivityInput, headers={}) + + with patch.object(interceptor.workflow, "memo_value", return_value=EXPECTED_ATTRS): + outbound.start_activity(start_input) + + inbound_next = Mock() + inbound_next.execute_activity = AsyncMock(return_value="ok") + inbound = interceptor._ActivityInbound(inbound_next) + info = Mock(workflow_id="task-1") + with patch.object(interceptor.activity, "info", return_value=info): + await inbound.execute_activity(Mock(headers=sent["headers"])) + + assert sgp_evals.attrs_for_span(_span(trace_id="task-1")) == EXPECTED_ATTRS + + async def test_plain_activity_does_not_inherit_a_reused_workflow_ids_eval_attrs(self) -> None: + sgp_evals.register_task("task-1", EXPECTED_ATTRS) + inbound_next = Mock() + inbound_next.execute_activity = AsyncMock(return_value="ok") + with patch.object(interceptor.activity, "info", return_value=Mock(workflow_id="task-1")): + await interceptor._ActivityInbound(inbound_next).execute_activity(Mock(headers={})) + + assert sgp_evals.attrs_for_span(_span(trace_id="task-1")) == {} + + def test_workflow_without_memo_adds_no_header(self) -> None: + sent: dict[str, Any] = {} + next_outbound = Mock() + next_outbound.start_activity = lambda input: sent.update(headers=input.headers) + start_input = Mock(spec=StartActivityInput, headers={}) + + with patch.object(interceptor.workflow, "memo_value", return_value=None): + interceptor._WorkflowOutbound(next_outbound).start_activity(start_input) + + assert sent["headers"] == {} + + def test_local_activities_get_the_header_too(self) -> None: + sent: dict[str, Any] = {} + next_outbound = Mock() + next_outbound.start_local_activity = lambda input: sent.update(headers=input.headers) + start_input = Mock(spec=StartLocalActivityInput, headers={}) + + with patch.object(interceptor.workflow, "memo_value", return_value=EXPECTED_ATTRS): + interceptor._WorkflowOutbound(next_outbound).start_local_activity(start_input) + + assert interceptor.ATTRS_HEADER in sent["headers"] + + async def test_worker_runs_with_the_interceptor_ahead_of_agent_interceptors(self) -> None: + agent_interceptor = interceptor.SGPEvalsInterceptor() + module = "agentex.lib.core.temporal.workers.worker" + with ExitStack() as stack: + stack.enter_context(patch(f"{module}.EnvironmentVariables")) + stack.enter_context(patch(f"{module}.init_sgp_obs")) + for name in ("shutdown_sgp_obs", "shutdown_default_span_queue", "shutdown_sync_tracing_processors", "get_temporal_client"): + stack.enter_context(patch(f"{module}.{name}", new=AsyncMock())) + worker_cls = stack.enter_context(patch(f"{module}.Worker")) + worker_cls.return_value.run = AsyncMock() + worker = AgentexWorker(task_queue="q", interceptors=[agent_interceptor]) + stack.enter_context(patch.object(worker, "start_health_check_server", new=AsyncMock())) + stack.enter_context(patch.object(worker, "_register_agent", new=AsyncMock())) + await worker.run(activities=[], workflow=object) + + interceptors = worker_cls.call_args.kwargs["interceptors"] + assert isinstance(interceptors[0], interceptor.SGPEvalsInterceptor) + assert interceptors[1:] == [agent_interceptor]