Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 16 additions & 0 deletions src/agent/stirrup_agent/tests/test_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,9 +51,16 @@ class _TC:
tool_call_id: str


@dataclass
class _Reasoning:
content: str
signature: str | None = None


@dataclass
class _Assistant:
content: str
reasoning: _Reasoning | None = None
tool_calls: list = field(default_factory=list)
token_usage: _Usage = field(default_factory=_Usage)
request_start_time: float | None = None
Expand Down Expand Up @@ -192,6 +199,10 @@ def test_build_trajectory_maps_turns_calls_and_outputs():
[
_Assistant(
content="let me check work orders",
reasoning=_Reasoning(
content="I should query the work-order tool.",
signature="sig-1",
),
tool_calls=[_TC("wo__get_work_order", '{"asset": "CWC04013"}', "t1")],
token_usage=_Usage(input=20, answer=8, reasoning=2),
request_start_time=1.0,
Expand All @@ -217,6 +228,11 @@ def test_build_trajectory_maps_turns_calls_and_outputs():
assert len(traj.turns) == 2
assert traj.total_input_tokens == 25
assert traj.total_output_tokens == 16 # (8+2) + (6+0)
assert traj.turns[0].reasoning == {
"signature": "sig-1",
"content": "I should query the work-order tool.",
}
assert traj.turns[1].reasoning is None

call = traj.all_tool_calls[0]
assert call.name == "wo__get_work_order"
Expand Down
24 changes: 22 additions & 2 deletions src/agent/stirrup_agent/trajectory.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,8 @@

Mapping:
* each ``AssistantMessage`` -> one :class:`~agent.models.TurnRecord`
(its ``content`` text, ``tool_calls``, ``token_usage``, request timing);
(its ``content`` text, ``reasoning``, ``tool_calls``, ``token_usage``,
request timing);
* each ``ToolMessage`` -> the ``output`` of the matching :class:`ToolCall`,
joined by ``tool_call_id``.

Expand All @@ -20,6 +21,7 @@
from __future__ import annotations

import json
from dataclasses import dataclass
from typing import Any, Iterable

from ..models import ToolCall, Trajectory, TurnRecord
Expand All @@ -31,6 +33,13 @@
_WEB_TOOL_NAMES = {"web_search", "web_fetch"}


@dataclass
class StirrupTurnRecord(TurnRecord):
"""A shared turn record plus Stirrup's optional reasoning payload."""

reasoning: dict[str, Any] | None = None


def classify_tool(tool_name: str, domain_servers: set[str]) -> str:
"""Label a Stirrup tool call ``"domain"`` / ``"code"`` / ``"other"``.

Expand Down Expand Up @@ -79,6 +88,16 @@ def _parse_arguments(arguments: Any) -> dict:
return {}


def _reasoning_payload(reasoning: Any) -> dict[str, Any] | None:
"""Return Stirrup ``AssistantMessage.reasoning`` in JSON-friendly form."""
if reasoning is None:
return None
return {
"signature": getattr(reasoning, "signature", None),
"content": getattr(reasoning, "content", "") or "",
}


def _ms(start: float | None, end: float | None) -> float | None:
if start is None or end is None:
return None
Expand Down Expand Up @@ -123,7 +142,7 @@ def build_trajectory(history: Iterable[Any]) -> Trajectory:
out_tok = getattr(usage, "output", 0) if usage else 0

trajectory.turns.append(
TurnRecord(
StirrupTurnRecord(
index=turn_index,
text=_content_text(getattr(msg, "content", "")),
tool_calls=tool_calls,
Expand All @@ -133,6 +152,7 @@ def build_trajectory(history: Iterable[Any]) -> Trajectory:
getattr(msg, "request_start_time", None),
getattr(msg, "request_end_time", None),
),
reasoning=_reasoning_payload(getattr(msg, "reasoning", None)),
)
)
turn_index += 1
Expand Down