Skip to content
Open
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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:
Expand Down
3 changes: 2 additions & 1 deletion src/agentex/lib/core/temporal/workers/worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)
Expand Down Expand Up @@ -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}")
Expand Down
22 changes: 16 additions & 6 deletions src/agentex/lib/core/tracing/processors/sgp_tracing_processor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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


Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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"
)
Expand Down
119 changes: 119 additions & 0 deletions src/agentex/lib/core/tracing/sgp_evals.py
Original file line number Diff line number Diff line change
@@ -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()
83 changes: 83 additions & 0 deletions src/agentex/lib/core/tracing/sgp_evals_interceptor.py
Original file line number Diff line number Diff line change
@@ -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]:
Comment thread
mohammadatallah-scale marked this conversation as resolved.
_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)
Comment thread
mohammadatallah-scale marked this conversation as resolved.
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)
Comment thread
mohammadatallah-scale marked this conversation as resolved.
4 changes: 4 additions & 0 deletions src/agentex/lib/core/tracing/trace.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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
Expand Down
5 changes: 5 additions & 0 deletions src/agentex/lib/sdk/fastacp/base/base_acp_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
Loading
Loading