diff --git a/sdks/python/README.md b/sdks/python/README.md index 9522cf2e..94785665 100644 --- a/sdks/python/README.md +++ b/sdks/python/README.md @@ -84,6 +84,54 @@ async with agent_control.AgentControlClient() as client: The existing `evaluate_controls` helper remains available for callers that prefer its field-based convenience arguments. +## Building trace/session steps + +Trace- and session-level controls evaluate the aggregate of several +already-executed child steps, so `Step.children` has to be populated by the +caller. Building that tree by hand means hand-rolling a side-channel record +for every leaf call and converting it into `Step` objects afterward: + +```python +# Before: a hand-built dict record next to the real return value +def execute_policy_search(query): + docs = search_policy_documents(query) + record = {"type": "retriever", "query": query, "docs": docs} + return StepExecution(value=docs, record=record) + +spans = [] +result = execute_policy_search(query) +spans.append(result.record) +... +trace_step = build_agent_control_step({"type": "trace", "spans": spans, ...}) +``` + +`agent_control.record_step()` builds the same tree incrementally, in `Step` +vocabulary, with no intermediate dict and no converter: + +```python +with agent_control.record_step("trace", "banking_trace", input={"request": req}) as trace: + with trace.child("retriever", "policy_lookup", input=query) as span: + span.output = search_policy_documents(query) + + account = trace.call(lookup_account, account_id="acct-1001", step_type="tool") + plan = trace.call(run_banking_model, req, step_type="llm", tools=TOOL_DEFINITIONS) + trace.output = {"status": "planned", "message": plan["content"]} + +result = await trace.evaluate(stage="post") +``` + +`trace.call(...)`/`await trace.acall(...)` run the function and record it as +a child using the same capture logic as `@control()` (input from bound +arguments, output from the return value), then return the real result - +removing the need for a separate `StepExecution`-style wrapper. A failed +call is still recorded (with the error in `context`) before the exception is +re-raised. `trace.child(...)` nests another recorder the same way, so +sessions nest traces with `session.child("trace", ...)`. + +`trace.build()` produces the frozen `Step`; `trace.evaluate(...)` builds it +and evaluates it in one call via `evaluate_step()` - the same function +`evaluate_controls()` uses internally once its `Step` is built. + ## Sharing an OpenTelemetry provider with Google ADK When Google ADK and Agent Control should export through the same OpenTelemetry diff --git a/sdks/python/src/agent_control/__init__.py b/sdks/python/src/agent_control/__init__.py index 04d040f6..9cd90511 100644 --- a/sdks/python/src/agent_control/__init__.py +++ b/sdks/python/src/agent_control/__init__.py @@ -92,7 +92,7 @@ async def handle_input(user_message: str) -> str: ) from .client import AgentControlClient from .control_decorators import ControlSteerError, ControlViolationError, control -from .evaluation import check_evaluation_with_local, evaluate_controls +from .evaluation import check_evaluation_with_local, evaluate_controls, evaluate_step from .observability import ( LogConfig, add_event, @@ -116,6 +116,7 @@ async def handle_input(user_message: str) -> str: ) from .otel_sink import control_event_to_otel_span from .runtime_auth import validate_http_field_name +from .step_recorder import StepRecorder, record_step from .tracing import ( get_current_span_id, get_current_trace_id, @@ -1619,6 +1620,10 @@ async def main(): # Local evaluation "check_evaluation_with_local", "evaluate_controls", + "evaluate_step", + # Step recorder (incremental trace/session Step tree builder) + "record_step", + "StepRecorder", # Tracing "get_trace_and_span_ids", "get_current_trace_id", diff --git a/sdks/python/src/agent_control/evaluation.py b/sdks/python/src/agent_control/evaluation.py index 800b4bf4..749c4f00 100644 --- a/sdks/python/src/agent_control/evaluation.py +++ b/sdks/python/src/agent_control/evaluation.py @@ -1,6 +1,6 @@ """Evaluation check operations for Agent Control SDK.""" -from collections.abc import Awaitable, Callable +from collections.abc import Awaitable, Callable, Mapping, Sequence from dataclasses import dataclass from inspect import iscoroutinefunction from typing import Any, Literal, cast @@ -515,6 +515,58 @@ def _with_parse_errors(result: EvaluationResult) -> EvaluationResult: return _with_parse_errors(EvaluationResult(is_safe=True, confidence=1.0)) +async def evaluate_step( + step: Step, + *, + agent_name: str, + stage: Literal["pre", "post"] = "pre", + target_type: str | None = None, + target_id: str | None = None, + trace_id: str | None = None, + span_id: str | None = None, +) -> EvaluationResult: + """Evaluate controls for an already-built ``Step``. + + This is the shared tail of :func:`evaluate_controls`: resolve the + session target, open a client, and run local/server evaluation. Use + it directly when the ``Step`` - including any ``children`` - was + already assembled, e.g. via :class:`~agent_control.step_recorder.StepRecorder`. + + When ``target_type`` and ``target_id`` are both supplied, the request + is target-bearing: the server merges target bindings into the + effective control set. If they are omitted, the SDK falls back to the + target context fixed at ``init()`` time when present. A per-call + override that disagrees with the session target is rejected because + the cached controls were fetched for the session target and would + otherwise drive stale local-first evaluation. + """ + if state.server_url is None: + raise RuntimeError("Server URL not configured. Call agent_control.init() first.") + + target_type, target_id = _resolve_session_target(target_type, target_id) + resolved_controls = state.server_controls or [] + + async with AgentControlClient( + base_url=state.server_url, + api_key=state.api_key, + api_key_header=state.api_key_header, + runtime_token_header=state.runtime_token_header, + runtime_token_cache=state.runtime_token_cache, + ) as client: + return await check_evaluation_with_local( + client=client, + agent_name=agent_name, + step=step, + stage=stage, + controls=resolved_controls, + target_type=target_type, + target_id=target_id, + trace_id=trace_id, + span_id=span_id, + event_agent_name=agent_name, + ) + + async def evaluate_controls( step_name: str, *, @@ -523,7 +575,7 @@ async def evaluate_controls( context: dict[str, Any] | None = None, tools: list[dict[str, JSONValue]] | None = None, ground_truth: JSONValue | None = None, - children: list[Step] | None = None, + children: Sequence[Step | Mapping[str, Any]] | None = None, step_type: str = "llm", stage: Literal["pre", "post"] = "pre", agent_name: str, @@ -544,11 +596,6 @@ async def evaluate_controls( """ step_type = ensure_step_type(step_type) - if state.server_url is None: - raise RuntimeError("Server URL not configured. Call agent_control.init() first.") - - target_type, target_id = _resolve_session_target(target_type, target_id) - default_value = {} if step_type == "tool" else "" step_dict: dict[str, Any] = { "type": step_type, @@ -566,24 +613,13 @@ async def evaluate_controls( step_dict["children"] = children step_obj = Step(**step_dict) # type: ignore[arg-type] - resolved_controls = state.server_controls or [] - async with AgentControlClient( - base_url=state.server_url, - api_key=state.api_key, - api_key_header=state.api_key_header, - runtime_token_header=state.runtime_token_header, - runtime_token_cache=state.runtime_token_cache, - ) as client: - return await check_evaluation_with_local( - client=client, - agent_name=agent_name, - step=step_obj, - stage=stage, - controls=resolved_controls, - target_type=target_type, - target_id=target_id, - trace_id=trace_id, - span_id=span_id, - event_agent_name=agent_name, - ) + return await evaluate_step( + step_obj, + agent_name=agent_name, + stage=stage, + target_type=target_type, + target_id=target_id, + trace_id=trace_id, + span_id=span_id, + ) diff --git a/sdks/python/src/agent_control/step_recorder.py b/sdks/python/src/agent_control/step_recorder.py new file mode 100644 index 00000000..160d5ed0 --- /dev/null +++ b/sdks/python/src/agent_control/step_recorder.py @@ -0,0 +1,215 @@ +"""Incremental builder for ``Step`` trees (trace/session-level evaluation). + +Manually populating ``Step.children`` for a trace or session means building +the whole tree by hand before calling ``evaluate_controls``/``evaluate_step``. +``StepRecorder`` lets callers build that tree as they execute instead: each +nested call is appended to its parent as it happens, with no intermediate +dict representation and no separate converter. +""" + +from __future__ import annotations + +from collections.abc import Awaitable, Callable, Mapping +from typing import Any, Literal, TypeVar + +from agent_control_models import EvaluationResult, JSONObject, JSONValue, Step + +from ._state import state +from .control_decorators import ToolsConfiguration, _create_evaluation_payload +from .evaluation import evaluate_step +from .validation import ensure_step_type + +T = TypeVar("T") + +# Step types whose children are reported as a span/trace tree; an absent +# recording means "nothing happened", which the record factory treats +# differently from "not applicable" (None). See models/agent.py and +# records/factory.py for the trace/session child-shape rules this mirrors. +_EMPTY_CHILDREN_STEP_TYPES = frozenset({"trace", "session"}) + + +def record_step( + step_type: str, + name: str, + *, + input: Any | None = None, + context: dict[str, Any] | None = None, + tools: list[JSONObject] | None = None, + ground_truth: JSONValue | None = None, +) -> StepRecorder: + """Start building a ``Step`` tree. Use as a context manager.""" + return StepRecorder( + step_type, + name, + input=input, + context=context, + tools=tools, + ground_truth=ground_truth, + ) + + +class StepRecorder: + """Mutable builder for one ``Step`` node and its children. + + ``Step`` itself is frozen, so this accumulates state and only + constructs the real (frozen) ``Step`` in :meth:`build`. + """ + + def __init__( + self, + step_type: str, + name: str, + *, + input: Any | None = None, + context: dict[str, Any] | None = None, + tools: list[JSONObject] | None = None, + ground_truth: JSONValue | None = None, + _parent: StepRecorder | None = None, + ) -> None: + self.type = ensure_step_type(step_type) + self.name = name + self.input = input + self.output: Any | None = None + self.context: dict[str, Any] | None = dict(context) if context is not None else None + self.tools = tools + self.ground_truth = ground_truth + self._children: list[Step] = [] + self._parent = _parent + + def __enter__(self) -> StepRecorder: + return self + + def __exit__(self, *exc_info: object) -> None: + if self._parent is not None: + self._parent._children.append(self.build()) + + def child( + self, + step_type: str, + name: str, + *, + input: Any | None = None, + context: dict[str, Any] | None = None, + tools: list[JSONObject] | None = None, + ground_truth: JSONValue | None = None, + ) -> StepRecorder: + """Start a nested recorder; it attaches to this node on ``__exit__``.""" + return StepRecorder( + step_type, + name, + input=input, + context=context, + tools=tools, + ground_truth=ground_truth, + _parent=self, + ) + + def add(self, step: Step | Mapping[str, Any]) -> None: + """Attach an already-built child step.""" + self._children.append(step if isinstance(step, Step) else Step.model_validate(step)) + + def _record_call_result( + self, + func: Callable[..., Any], + args: tuple[Any, ...], + kwargs: dict[str, Any], + output: Any, + step_name: str | None, + step_type: str | None, + tools: ToolsConfiguration, + error: BaseException | None, + ) -> None: + payload = _create_evaluation_payload( + func, args, kwargs, output, step_name, step_type, tools + ) + if error is not None: + payload["context"] = {**(payload.get("context") or {}), "error": repr(error)} + self._children.append(Step.model_validate(payload)) + + def call( + self, + func: Callable[..., T], + *args: Any, + step_type: str | None = None, + step_name: str | None = None, + tools: ToolsConfiguration = None, + **kwargs: Any, + ) -> T: + """Run ``func``, record it via ``@control()``'s capture logic, return its output.""" + try: + output = func(*args, **kwargs) + except Exception as exc: + self._record_call_result(func, args, kwargs, None, step_name, step_type, tools, exc) + raise + self._record_call_result(func, args, kwargs, output, step_name, step_type, tools, None) + return output + + async def acall( + self, + func: Callable[..., Awaitable[T]], + *args: Any, + step_type: str | None = None, + step_name: str | None = None, + tools: ToolsConfiguration = None, + **kwargs: Any, + ) -> T: + """Async counterpart to :meth:`call`.""" + try: + output = await func(*args, **kwargs) + except Exception as exc: + self._record_call_result(func, args, kwargs, None, step_name, step_type, tools, exc) + raise + self._record_call_result(func, args, kwargs, output, step_name, step_type, tools, None) + return output + + def build(self) -> Step: + """Recursively construct the frozen ``Step`` for this node and its children.""" + step_dict: dict[str, Any] = { + "type": self.type, + "name": self.name, + "input": self.input, + "output": self.output, + } + if self.context is not None: + step_dict["context"] = self.context + if self.tools is not None: + step_dict["tools"] = self.tools + if self.ground_truth is not None: + step_dict["ground_truth"] = self.ground_truth + if self._children: + step_dict["children"] = list(self._children) + elif self.type in _EMPTY_CHILDREN_STEP_TYPES: + step_dict["children"] = [] + return Step(**step_dict) # type: ignore[arg-type] + + async def evaluate( + self, + *, + agent_name: str | None = None, + stage: Literal["pre", "post"] = "pre", + target_type: str | None = None, + target_id: str | None = None, + trace_id: str | None = None, + span_id: str | None = None, + ) -> EvaluationResult: + """Build this node into a ``Step`` and evaluate it via :func:`evaluate_step`. + + ``agent_name`` defaults to the agent registered via ``init()``. + """ + resolved_agent_name = agent_name or ( + state.current_agent.agent_name if state.current_agent is not None else None + ) + if resolved_agent_name is None: + raise RuntimeError( + "agent_name not supplied and no agent registered. " + "Call agent_control.init() first or pass agent_name explicitly." + ) + return await evaluate_step( + self.build(), + agent_name=resolved_agent_name, + stage=stage, + target_type=target_type, + target_id=target_id, + trace_id=trace_id, + span_id=span_id, + ) diff --git a/sdks/python/tests/test_evaluation.py b/sdks/python/tests/test_evaluation.py index 94f596fc..9657545a 100644 --- a/sdks/python/tests/test_evaluation.py +++ b/sdks/python/tests/test_evaluation.py @@ -4,11 +4,12 @@ from uuid import UUID import pytest -from agent_control import evaluation -from agent_control.evaluation import EvaluationResult from agent_control_models import Step from pydantic import ValidationError +from agent_control import evaluation +from agent_control.evaluation import EvaluationResult + @pytest.mark.asyncio async def test_check_evaluation_requires_step_name_before_server_call(): @@ -187,6 +188,55 @@ async def test_evaluate_controls_with_explicit_agent_name(monkeypatch): assert result.is_safe is True assert result.confidence == 1.0 + + +@pytest.mark.asyncio +async def test_evaluate_controls_delegates_to_evaluate_step(monkeypatch): + """evaluate_controls builds the Step then routes through evaluate_step.""" + mock_result = EvaluationResult(is_safe=True, confidence=1.0) + mock_evaluate_step = AsyncMock(return_value=mock_result) + monkeypatch.setattr(evaluation, "evaluate_step", mock_evaluate_step) + + result = await evaluation.evaluate_controls( + step_name="chat", + input="hello", + stage="pre", + agent_name="test-bot", + ) + + assert result is mock_result + mock_evaluate_step.assert_awaited_once() + args, kwargs = mock_evaluate_step.call_args + step_arg = args[0] + assert isinstance(step_arg, Step) + assert step_arg.name == "chat" + assert step_arg.input == "hello" + assert kwargs["agent_name"] == "test-bot" + assert kwargs["stage"] == "pre" + + +@pytest.mark.asyncio +async def test_evaluate_controls_coerces_dict_children_to_step(monkeypatch): + """children= accepts plain dicts; pydantic coerces them to Step instances.""" + mock_check = AsyncMock(return_value=EvaluationResult(is_safe=True, confidence=1.0)) + monkeypatch.setattr(evaluation, "check_evaluation_with_local", mock_check) + + with patch("agent_control.state.server_url", "http://localhost:8000"): + await evaluation.evaluate_controls( + step_name="trace1", + step_type="trace", + input={"request": "r1"}, + children=[{"type": "llm", "name": "x", "input": "hi", "output": "bye"}], + stage="post", + agent_name="test-bot", + ) + + step = mock_check.call_args.kwargs["step"] + assert step.children is not None + assert len(step.children) == 1 + assert isinstance(step.children[0], Step) + assert step.children[0].type == "llm" + assert step.children[0].name == "x" mock_check.assert_called_once() diff --git a/sdks/python/tests/test_step_recorder.py b/sdks/python/tests/test_step_recorder.py new file mode 100644 index 00000000..ca0e7e6e --- /dev/null +++ b/sdks/python/tests/test_step_recorder.py @@ -0,0 +1,220 @@ +"""Tests for StepRecorder, the incremental Step tree builder.""" + +from unittest.mock import AsyncMock, patch + +import pytest +from agent_control_models import EvaluationResult, Step + +from agent_control import record_step, step_recorder +from agent_control.control_decorators import _create_evaluation_payload + + +def test_nested_trace_builds_expected_step(): + """A trace with span children and a leaf retriever builds the matching tree.""" + with record_step("trace", "banking_trace", input={"request": "r1"}) as trace: + with trace.child("retriever", "policy_lookup", input="refund policy?") as retriever: + retriever.output = {"hits": 1} + trace.output = {"status": "planned"} + + step = trace.build() + + assert step.type == "trace" + assert step.name == "banking_trace" + assert step.input == {"request": "r1"} + assert step.output == {"status": "planned"} + assert step.children is not None + assert len(step.children) == 1 + child = step.children[0] + assert child.type == "retriever" + assert child.name == "policy_lookup" + assert child.input == "refund policy?" + assert child.output == {"hits": 1} + + +def test_session_nests_traces(): + """Sessions nest traces the same way traces nest spans.""" + with record_step("session", "banking_session") as session: + with session.child("trace", "turn_1") as trace: + with trace.child("llm", "respond") as llm: + llm.output = "hi" + + step = session.build() + + assert step.type == "session" + assert len(step.children) == 1 + assert step.children[0].type == "trace" + assert step.children[0].name == "turn_1" + assert len(step.children[0].children) == 1 + assert step.children[0].children[0].type == "llm" + + +def test_call_matches_decorator_payload_for_tool(): + """StepRecorder.call() must capture the same payload _create_evaluation_payload would.""" + + def lookup_account(account_id: str) -> dict: + return {"balance": 100} + + with record_step("trace", "t") as trace: + result = trace.call( + lookup_account, account_id="acct-1", step_type="tool", step_name="lookup_account" + ) + + assert result == {"balance": 100} + expected_payload = _create_evaluation_payload( + lookup_account, (), {"account_id": "acct-1"}, result, "lookup_account", "tool", None + ) + [child] = trace.build().children + assert child.model_dump(exclude_none=True) == { + k: v for k, v in expected_payload.items() if v is not None + } + + +@pytest.mark.asyncio +async def test_acall_runs_async_function_and_records_child(): + """StepRecorder.acall() awaits the function and records the result.""" + + async def run_banking_model(prompt: str) -> str: + return f"answer: {prompt}" + + with record_step("trace", "t") as trace: + result = await trace.acall( + run_banking_model, "hello", step_type="llm", step_name="respond" + ) + + assert result == "answer: hello" + [child] = trace.build().children + assert child.type == "llm" + assert child.name == "respond" + assert child.input == "hello" + assert child.output == "answer: hello" + + +def test_call_records_exception_and_reraises(): + """A failing call is still recorded as a child (with the error in context), then re-raised.""" + + def flaky(x: int) -> int: + raise ValueError("boom") + + with record_step("trace", "t") as trace: + with pytest.raises(ValueError, match="boom"): + trace.call(flaky, 1, step_type="tool", step_name="flaky") + + [child] = trace.build().children + assert child.name == "flaky" + assert child.output is None + assert child.context == {"error": "ValueError('boom')"} + + +@pytest.mark.asyncio +async def test_acall_records_exception_and_reraises(): + """acall()'s failure path mirrors call(): record (error in context), then re-raise.""" + + async def flaky(x: int) -> int: + raise ValueError("boom") + + with record_step("trace", "t") as trace: + with pytest.raises(ValueError, match="boom"): + await trace.acall(flaky, 1, step_type="tool", step_name="flaky") + + [child] = trace.build().children + assert child.name == "flaky" + assert child.output is None + assert child.context == {"error": "ValueError('boom')"} + + +def test_build_includes_context_tools_ground_truth(): + """build() carries context/tools/ground_truth through when supplied.""" + with record_step( + "llm", + "respond", + input="hi", + context={"locale": "en-US"}, + tools=[{"type": "function", "function": {"name": "search"}}], + ground_truth="hello", + ) as leaf: + leaf.output = "hello" + + step = leaf.build() + assert step.context == {"locale": "en-US"} + assert step.tools == [{"type": "function", "function": {"name": "search"}}] + assert step.ground_truth == "hello" + + +def test_empty_trace_and_session_build_empty_children_list(): + """A trace/session with no recorded children gets children=[], not None.""" + with record_step("trace", "empty_trace") as trace: + pass + with record_step("session", "empty_session") as session: + pass + + assert trace.build().children == [] + assert session.build().children == [] + + +def test_llm_step_without_children_stays_none(): + """Non trace/session step types are untouched when no children were recorded.""" + with record_step("llm", "respond", input="hi") as leaf: + leaf.output = "hello" + + assert leaf.build().children is None + + +def test_add_attaches_prebuilt_step_or_dict(): + """add() accepts both a real Step and a plain dict (coerced via Step.model_validate).""" + with record_step("trace", "t") as trace: + trace.add(Step(type="llm", name="a", input="x", output="y")) + trace.add({"type": "tool", "name": "b", "input": {}, "output": {}}) + + children = trace.build().children + assert [c.name for c in children] == ["a", "b"] + assert all(isinstance(c, Step) for c in children) + + +@pytest.mark.asyncio +async def test_evaluate_forwards_to_evaluate_step(): + """StepRecorder.evaluate() builds the Step and delegates to evaluate_step().""" + mock_result = EvaluationResult(is_safe=True, confidence=1.0) + mock_evaluate_step = AsyncMock(return_value=mock_result) + + with record_step("trace", "t", input="hi") as trace: + pass + + with patch.object(step_recorder, "evaluate_step", mock_evaluate_step): + result = await trace.evaluate(agent_name="test-bot", stage="post") + + assert result is mock_result + mock_evaluate_step.assert_awaited_once() + _, kwargs = mock_evaluate_step.call_args + assert kwargs["agent_name"] == "test-bot" + assert kwargs["stage"] == "post" + + +@pytest.mark.asyncio +async def test_evaluate_defaults_agent_name_from_current_agent(): + """When agent_name is omitted, evaluate() falls back to the agent registered via init().""" + mock_result = EvaluationResult(is_safe=True, confidence=1.0) + mock_evaluate_step = AsyncMock(return_value=mock_result) + + with record_step("trace", "t") as trace: + pass + + class _FakeAgent: + agent_name = "test-bot-0123456789" + + with patch.object(step_recorder, "evaluate_step", mock_evaluate_step): + with patch.object(step_recorder.state, "current_agent", _FakeAgent()): + await trace.evaluate(stage="post") + + _, kwargs = mock_evaluate_step.call_args + assert kwargs["agent_name"] == "test-bot-0123456789" + + +@pytest.mark.asyncio +async def test_evaluate_without_agent_name_or_init_raises(): + """Without an explicit agent_name or a prior init(), evaluate() fails loudly.""" + with record_step("trace", "t") as trace: + pass + + with patch.object(step_recorder.state, "current_agent", None): + with pytest.raises(RuntimeError, match="agent_name not supplied"): + await trace.evaluate(stage="post") diff --git a/sdks/python/tests/test_step_recorder_galileo_integration.py b/sdks/python/tests/test_step_recorder_galileo_integration.py new file mode 100644 index 00000000..6b6d3bc2 --- /dev/null +++ b/sdks/python/tests/test_step_recorder_galileo_integration.py @@ -0,0 +1,87 @@ +"""Cross-package check: StepRecorder trees feed the Galileo record factory. + +The galileo extras are normally installed in the dev environment (see +``test_evaluators_optional_imports.py``), so this skips cleanly when they +are not available instead of failing the suite. +""" + +from __future__ import annotations + +import importlib.util + +import pytest + +from agent_control import record_step + + +def _module_available(name: str) -> bool: + try: + return importlib.util.find_spec(name) is not None + except (ImportError, ValueError): + return False + + +_GALILEO_INSTALLED = _module_available("agent_control_evaluator_galileo.records") + +pytestmark = pytest.mark.skipif( + not _GALILEO_INSTALLED, + reason="agent-control-evaluator-galileo extras not installed in this environment", +) + + +def _record_from_step(step): + from agent_control_evaluator_galileo.records.factory import record_from_step + + return record_from_step(step) + + +def test_recorder_built_trace_matches_demo_shape(): + """A trace with one llm/tool/retriever child each builds a 3-span Trace.""" + + def policy_search(query: str) -> dict: + return {"docs": ["policy-1"]} + + def account_lookup(account_id: str) -> dict: + return {"balance": 500} + + def banking_llm(prompt: str) -> str: + return f"response: {prompt}" + + with record_step("trace", "banking_trace", input={"request": "r1"}) as trace: + trace.call( + policy_search, query="refund policy", step_type="retriever", step_name="policy_search" + ) + trace.call( + account_lookup, account_id="acct-1", step_type="tool", step_name="account_lookup" + ) + trace.call(banking_llm, "draft a reply", step_type="llm", step_name="banking_llm") + trace.output = {"status": "done"} + + record = _record_from_step(trace.build()) + + assert type(record).__name__ == "Trace" + assert len(record.spans) == 3 + assert {type(span).__name__ for span in record.spans} == { + "LlmSpan", + "ToolSpan", + "RetrieverSpan", + } + + +def test_recorder_built_session_matches_demo_shape(): + """A session with 2 traces, each with >=2 spans, builds the matching Session.""" + with record_step("session", "banking_session") as session: + for i in range(2): + with session.child("trace", f"turn_{i}", input={"turn": i}) as trace: + with trace.child("llm", "respond", input="hi") as llm: + llm.output = "hello" + with trace.child("tool", "lookup", input={}) as tool: + tool.output = {} + trace.output = {"turn": i, "status": "ok"} + + record = _record_from_step(session.build()) + + assert type(record).__name__ == "Session" + assert len(record.traces) == 2 + for sub_trace in record.traces: + assert len(sub_trace.spans) >= 2