Skip to main content

aether_evals/agents/
transcript.rs

1use super::{AgentRunResult, RunError};
2use crate::EvalRunError;
3use aether_core::events::{AgentEvent, ToolEvent};
4use futures::{Stream, StreamExt};
5use llm::SessionUsageTotals;
6use std::fmt::Debug;
7use thiserror::Error;
8
9pub struct Transcript {
10    events: Vec<AgentEvent>,
11}
12
13pub struct ToolCall<'a> {
14    pub name: &'a str,
15    pub arguments: &'a str,
16}
17
18#[derive(Error)]
19#[error("{error}")]
20pub struct TranscriptError {
21    transcript: Transcript,
22    #[source]
23    error: EvalRunError,
24}
25
26impl Transcript {
27    pub fn new(events: Vec<AgentEvent>) -> Self {
28        Self { events }
29    }
30
31    pub async fn from_stream<T: Stream<Item = AgentRunResult>>(stream: T) -> Result<Self, TranscriptError> {
32        let mut transcript = Self::default();
33        futures::pin_mut!(stream);
34        while let Some(result) = stream.next().await {
35            match result {
36                Ok(event) => {
37                    transcript.add(event);
38                }
39                Err(error) => return Err(TranscriptError::new(transcript, error)),
40            }
41        }
42        Ok(transcript)
43    }
44
45    pub fn add(&mut self, event: AgentEvent) {
46        self.events.push(event);
47    }
48
49    pub fn events(&self) -> &[AgentEvent] {
50        &self.events
51    }
52
53    pub fn all_tool_calls(&self) -> impl Iterator<Item = ToolCall<'_>> + '_ {
54        self.events.iter().filter_map(|event| match event {
55            AgentEvent::Tool(ToolEvent::Result { result, .. }) => {
56                Some(ToolCall { name: &result.name, arguments: &result.arguments })
57            }
58            AgentEvent::Tool(ToolEvent::Error { error, .. }) => {
59                Some(ToolCall { name: &error.name, arguments: error.arguments.as_deref().unwrap_or("") })
60            }
61            _ => None,
62        })
63    }
64
65    pub fn tool_calls<'a>(&'a self, name: &'a str) -> impl Iterator<Item = ToolCall<'a>> + 'a {
66        self.all_tool_calls().filter(move |call| call.name == name)
67    }
68
69    pub fn tool_called(&self, name: &str) -> bool {
70        self.tool_calls(name).next().is_some()
71    }
72
73    pub fn tool_call_count(&self, name: &str) -> usize {
74        self.tool_calls(name).count()
75    }
76
77    /// Session-wide token totals and estimated cost from the last usage event,
78    /// or zeroed totals if no usage was recorded.
79    pub fn usage(&self) -> SessionUsageTotals {
80        self.events
81            .iter()
82            .rev()
83            .find_map(|event| match event {
84                AgentEvent::SessionUsage(usage) => Some(usage.totals.clone()),
85                _ => None,
86            })
87            .unwrap_or_default()
88    }
89}
90
91impl Default for Transcript {
92    fn default() -> Self {
93        Self::new(Vec::new())
94    }
95}
96
97impl From<Vec<AgentEvent>> for Transcript {
98    fn from(events: Vec<AgentEvent>) -> Self {
99        Self::new(events)
100    }
101}
102
103impl ToolCall<'_> {
104    pub fn arguments_json(&self) -> Result<serde_json::Value, serde_json::Error> {
105        serde_json::from_str(self.arguments)
106    }
107}
108
109impl TranscriptError {
110    fn new(transcript: Transcript, error: RunError) -> Self {
111        Self { transcript, error: EvalRunError::from(error) }
112    }
113
114    pub fn transcript(&self) -> &Transcript {
115        &self.transcript
116    }
117
118    pub fn error(&self) -> &EvalRunError {
119        &self.error
120    }
121
122    pub fn into_parts(self) -> (Transcript, EvalRunError) {
123        (self.transcript, self.error)
124    }
125}
126
127impl Debug for TranscriptError {
128    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
129        formatter.debug_struct("TranscriptError").field("error", &self.error).finish_non_exhaustive()
130    }
131}
132
133pub(crate) fn is_terminal(event: &AgentEvent) -> bool {
134    event.turn_outcome().is_some()
135}
136
137#[cfg(test)]
138mod tests {
139    use super::*;
140    use crate::{Agent, FakeAgent, Task};
141    use aether_core::events::TurnEvent;
142    use llm::testing::session_usage_event;
143    use llm::{TokenUsage, ToolCallRequest, ToolCallResult};
144
145    #[tokio::test]
146    async fn transcript_from_stream() {
147        let agent = FakeAgent::with_tool_call("bash", "success");
148        let stream = agent.run(Task::new("do the thing"));
149        let transcript = Transcript::from_stream(stream).await.unwrap();
150
151        assert!(transcript.tool_called("bash"));
152        assert!(matches!(transcript.events().last(), Some(AgentEvent::Turn(TurnEvent::Ended { .. }))));
153    }
154
155    #[test]
156    fn tool_call_count_counts_matching_tool_calls() {
157        let transcript = transcript_with_events(vec![tool_call("bash"), tool_call("read"), tool_result("bash")]);
158
159        assert!(transcript.tool_called("bash"));
160        assert!(!transcript.tool_called("read"));
161        assert!(!transcript.tool_called("write"));
162        assert_eq!(transcript.tool_call_count("bash"), 1);
163        assert_eq!(transcript.tool_call_count("read"), 0);
164    }
165
166    #[test]
167    fn tool_call_arguments_json_parses_arguments() {
168        let call = ToolCall { name: "bash", arguments: r#"{"command":"pwd"}"# };
169
170        assert_eq!(call.arguments_json().unwrap(), serde_json::json!({ "command": "pwd" }));
171    }
172
173    #[test]
174    fn tool_call_arguments_json_returns_error_for_invalid_json() {
175        let call = ToolCall { name: "bash", arguments: "not json" };
176
177        assert!(call.arguments_json().is_err());
178    }
179
180    #[test]
181    fn usage_returns_zeroed_totals_when_no_usage_was_recorded() {
182        let transcript = transcript_with_events(vec![tool_call("bash")]);
183        assert_eq!(transcript.usage(), SessionUsageTotals::default());
184    }
185
186    #[test]
187    fn usage_extracts_the_final_session_totals() {
188        let mut last = session_usage_event(2, TokenUsage::new(2000, 500));
189        last.totals.tokens = TokenUsage::new(3000, 600);
190        last.totals.unpriced_calls = 2;
191        let transcript = transcript_with_events(vec![
192            AgentEvent::SessionUsage(session_usage_event(1, TokenUsage::new(1000, 100))),
193            AgentEvent::SessionUsage(last),
194        ]);
195
196        let usage = transcript.usage();
197        assert_eq!(usage.tokens.input_tokens.get(), 3000);
198        assert_eq!(usage.tokens.output_tokens.get(), 600);
199        assert_eq!(usage.tokens.total_tokens().get(), 3600);
200        assert_eq!(usage.unpriced_calls, 2);
201        assert!(!usage.is_fully_priced());
202    }
203
204    fn transcript_with_events(events: Vec<AgentEvent>) -> Transcript {
205        Transcript::new(events)
206    }
207
208    fn tool_call(name: &str) -> AgentEvent {
209        AgentEvent::Tool(ToolEvent::Call {
210            request: ToolCallRequest { id: name.to_string(), name: name.to_string(), arguments: "{}".to_string() },
211        })
212    }
213
214    fn tool_result(name: &str) -> AgentEvent {
215        AgentEvent::Tool(ToolEvent::Result {
216            result: ToolCallResult {
217                id: name.to_string(),
218                name: name.to_string(),
219                arguments: "{}".to_string(),
220                result: "ok".to_string(),
221            },
222            result_meta: None,
223        })
224    }
225}