1use std::sync::Arc;
9use std::time::{Duration, Instant};
10
11use anyhow::{anyhow, bail};
12use everruns::ask_user::{AskUser, DefaultsResponder, Question, Status};
13use serde_json::{Value, json};
14
15use crate::app::{App, Mode};
16use crate::host::{Host, NewSession, Notice, wire_json};
17
18const TURN_TIMEOUT: Duration = Duration::from_secs(180);
19const POLL: Duration = Duration::from_millis(100);
20
21enum Target {
22 Local(Arc<Host>),
23 Remote {
24 base: String,
25 client: reqwest::Client,
26 },
27}
28
29#[derive(Clone, Debug, PartialEq, Eq)]
31pub enum OnApproval {
32 Approve,
33 Deny,
34}
35
36#[derive(Clone, Debug, Default)]
38pub struct TurnRecord {
39 pub response: String,
40 pub success: bool,
41 pub error: Option<String>,
42 pub tools: Vec<String>,
44 pub approvals: Vec<String>,
46 pub questions: usize,
48 pub events: Vec<Value>,
50}
51
52pub struct EvalCx {
54 target: Target,
55 session: Option<String>,
56 agent: Option<String>,
57 cursor: i32,
59 on_approval: OnApproval,
60 turns: Vec<TurnRecord>,
61}
62
63impl EvalCx {
64 fn new(target: Target) -> Self {
65 Self {
66 target,
67 session: None,
68 agent: None,
69 cursor: 0,
70 on_approval: OnApproval::Approve,
71 turns: Vec::new(),
72 }
73 }
74
75 #[cfg(test)]
76 pub(crate) fn local_for_test(host: Arc<Host>) -> Self {
77 Self::new(Target::Local(host))
78 }
79
80 #[cfg(test)]
81 pub(crate) fn remote_for_test(base: String) -> Self {
82 Self::new(Target::Remote {
83 base,
84 client: reqwest::Client::new(),
85 })
86 }
87
88 pub fn agent(&mut self, name: impl Into<String>) -> &mut Self {
90 self.agent = Some(name.into());
91 self
92 }
93
94 pub fn on_approval(&mut self, policy: OnApproval) -> &mut Self {
97 self.on_approval = policy;
98 self
99 }
100
101 pub async fn send(&mut self, text: impl Into<String>) -> crate::Result {
103 let text = text.into();
104 let session = match &self.session {
105 Some(session) => session.clone(),
106 None => {
107 let session = self.create().await?;
108 self.session = Some(session.clone());
109 session
110 }
111 };
112 let turn = match &self.target {
113 Target::Local(host) => self.local_turn(host.clone(), &session, text).await?,
114 Target::Remote { .. } => self.remote_turn(&session, text).await?,
115 };
116 if let Some(last) = turn
117 .events
118 .iter()
119 .filter_map(|e| e["sequence"].as_i64())
120 .max()
121 {
122 self.cursor = i32::try_from(last).unwrap_or(self.cursor);
123 }
124 self.turns.push(turn);
125 Ok(())
126 }
127
128 pub fn completed(&self) -> crate::Result<TurnCheck<'_>> {
130 let turn = self.last()?;
131 if !turn.success {
132 bail!(
133 "turn did not complete: {}",
134 turn.error.clone().unwrap_or_default()
135 );
136 }
137 Ok(TurnCheck { turn })
138 }
139
140 pub fn last(&self) -> crate::Result<&TurnRecord> {
142 self.turns
143 .last()
144 .ok_or_else(|| anyhow!("no turn yet; call send() first"))
145 }
146
147 async fn create(&self) -> crate::Result<String> {
148 match &self.target {
149 Target::Local(host) => {
150 host.create_session(NewSession {
151 agent: self.agent.clone(),
152 metadata: Some(json!({ "eval": true })),
153 ..NewSession::default()
154 })
155 .await
156 }
157 Target::Remote { base, client } => {
158 let body: Value = client
159 .post(format!("{base}/v1/sessions"))
160 .json(&json!({ "agent_name": self.agent, "metadata": { "eval": true } }))
161 .send()
162 .await?
163 .error_for_status()?
164 .json()
165 .await?;
166 body.get("id")
167 .and_then(Value::as_str)
168 .map(str::to_string)
169 .ok_or_else(|| anyhow!("create session returned no id: {body}"))
170 }
171 }
172 }
173
174 async fn local_turn(
177 &self,
178 host: Arc<Host>,
179 session: &str,
180 text: String,
181 ) -> crate::Result<TurnRecord> {
182 let mut notices = host.notices.subscribe();
183 let pending = host.send(session, text).await?;
184 let mut record = TurnRecord::default();
185 let wait = pending.wait();
186 tokio::pin!(wait);
187 let deadline = tokio::time::sleep(TURN_TIMEOUT);
188 tokio::pin!(deadline);
189 let outcome = loop {
190 tokio::select! {
191 outcome = &mut wait => break outcome?,
192 () = &mut deadline => bail!("turn did not finish within {TURN_TIMEOUT:?}"),
193 notice = notices.recv() => match notice {
194 Ok(Notice::ApprovalRequested(view)) if view.session_id == session => {
195 record.approvals.push(view.tool_name.clone());
196 let _ = host.resolve_approval(
198 session,
199 &view.tool_call_id,
200 self.on_approval == OnApproval::Approve,
201 );
202 }
203 Ok(Notice::QuestionAsked { session_id, tool_call_id, questions })
204 if session_id == session =>
205 {
206 record.questions += 1;
207 let outcome = DefaultsResponder.ask(&questions).await;
208 let _ = host.answer_questions(
209 session,
210 Some(&tool_call_id),
211 Status::Answered,
212 outcome.answers,
213 );
214 }
215 Err(tokio::sync::broadcast::error::RecvError::Closed) => {
216 bail!("host shut down mid-turn")
217 }
218 _ => {}
219 },
220 }
221 };
222 record.response = outcome.response;
223 record.success = outcome.success;
224 record.error = outcome.error;
225 record.events = host
226 .events_after(session, self.cursor)
227 .await?
228 .iter()
229 .map(wire_json)
230 .collect();
231 record.tools = tools_called(&record.events);
232 Ok(record)
233 }
234
235 async fn remote_turn(&self, session: &str, text: String) -> crate::Result<TurnRecord> {
239 let Target::Remote { base, client } = &self.target else {
240 bail!("not a remote eval");
241 };
242 client
243 .post(format!("{base}/v1/sessions/{session}/messages"))
244 .json(&json!({ "message": { "role": "user", "content": [{ "type": "text", "text": text }] } }))
245 .send()
246 .await?
247 .error_for_status()?;
248 let mut record = TurnRecord::default();
249 let deadline = Instant::now() + TURN_TIMEOUT;
250 loop {
251 if Instant::now() > deadline {
252 bail!("turn did not finish within {TURN_TIMEOUT:?}");
253 }
254 let state: Value = client
255 .get(format!("{base}/v1/sessions/{session}"))
256 .send()
257 .await?
258 .error_for_status()?
259 .json()
260 .await?;
261 for pending in state["pending_approvals"].as_array().into_iter().flatten() {
262 record.approvals.push(
263 pending["tool_name"]
264 .as_str()
265 .unwrap_or_default()
266 .to_string(),
267 );
268 let decision = match self.on_approval {
269 OnApproval::Approve => "approve",
270 OnApproval::Deny => "deny",
271 };
272 let call = pending["tool_call_id"].as_str().unwrap_or_default();
273 client
274 .post(format!("{base}/v1/sessions/{session}/approvals/{call}"))
275 .json(&json!({ "decision": decision, "note": "answered by eval" }))
276 .send()
277 .await?;
278 }
279 for pending in state["pending_questions"].as_array().into_iter().flatten() {
280 record.questions += 1;
281 let questions: Vec<Question> =
282 serde_json::from_value(pending["questions"].clone())?;
283 let outcome = DefaultsResponder.ask(&questions).await;
284 client
285 .post(format!("{base}/v1/sessions/{session}/question-answers"))
286 .json(&json!({
287 "tool_call_id": pending["tool_call_id"],
288 "status": "answered",
289 "answers": outcome.answers,
290 }))
291 .send()
292 .await?;
293 }
294 if state["status"] == "idle" {
295 let events: Value = client
296 .get(format!(
297 "{base}/v1/sessions/{session}/events?after_sequence={}",
298 self.cursor
299 ))
300 .send()
301 .await?
302 .error_for_status()?
303 .json()
304 .await?;
305 let events = events["data"].as_array().cloned().unwrap_or_default();
306 if let Some(terminal) = events.iter().rev().find(|event| is_terminal(event)) {
307 record.success = terminal["type"] == "turn.completed";
308 record.error = terminal["data"]["error"].as_str().map(str::to_string);
309 record.response = final_response(&events);
310 record.tools = tools_called(&events);
311 record.events = events;
312 return Ok(record);
313 }
314 }
315 tokio::time::sleep(POLL).await;
316 }
317 }
318}
319
320fn is_terminal(event: &Value) -> bool {
321 matches!(
322 event["type"].as_str(),
323 Some("turn.completed" | "turn.failed" | "turn.cancelled")
324 )
325}
326
327fn tools_called(events: &[Value]) -> Vec<String> {
329 events
330 .iter()
331 .filter(|event| event["type"] == "tool.started")
332 .filter_map(|event| event["data"]["tool_call"]["name"].as_str())
333 .map(str::to_string)
334 .collect()
335}
336
337fn final_response(events: &[Value]) -> String {
339 events
340 .iter()
341 .rev()
342 .filter(|event| event["type"] == "output.message.completed")
343 .map(|event| {
344 event["data"]["message"]["content"]
345 .as_array()
346 .into_iter()
347 .flatten()
348 .filter(|part| part["type"] == "text")
349 .filter_map(|part| part["text"].as_str())
350 .collect::<Vec<_>>()
351 .join("")
352 })
353 .find(|text| !text.is_empty())
354 .unwrap_or_default()
355}
356
357pub struct TurnCheck<'a> {
360 turn: &'a TurnRecord,
361}
362
363impl<'a> TurnCheck<'a> {
364 pub fn called_tool(self, name: &str) -> crate::Result<Self> {
365 if self.turn.tools.iter().any(|tool| tool == name) {
366 Ok(self)
367 } else {
368 bail!(
369 "expected a `{name}` call; tools called: {:?}",
370 self.turn.tools
371 )
372 }
373 }
374
375 pub fn did_not_call(self, name: &str) -> crate::Result<Self> {
376 if self.turn.tools.iter().any(|tool| tool == name) {
377 bail!("`{name}` was called but should not have been")
378 }
379 Ok(self)
380 }
381
382 pub fn asked_approval(self, tool: &str) -> crate::Result<Self> {
383 if self.turn.approvals.iter().any(|name| name == tool) {
384 Ok(self)
385 } else {
386 bail!(
387 "expected `{tool}` to ask for approval; approvals: {:?}",
388 self.turn.approvals
389 )
390 }
391 }
392
393 pub fn reply_contains(self, needle: &str) -> crate::Result<Self> {
395 if self
396 .turn
397 .response
398 .to_lowercase()
399 .contains(&needle.to_lowercase())
400 {
401 Ok(self)
402 } else {
403 bail!(
404 "reply does not contain {needle:?}: {:?}",
405 self.turn.response
406 )
407 }
408 }
409
410 pub fn reply(&self) -> &'a str {
411 &self.turn.response
412 }
413}
414
415#[derive(Clone, Debug, Default)]
417pub struct EvalReport {
418 pub results: Vec<EvalResult>,
419}
420
421#[derive(Clone, Debug)]
422pub struct EvalResult {
423 pub name: String,
424 pub passed: bool,
425 pub error: Option<String>,
426 pub duration: Duration,
427}
428
429impl EvalReport {
430 pub fn passed(&self) -> bool {
431 self.results.iter().all(|result| result.passed)
432 }
433}
434
435pub(crate) async fn run(
437 app: &App,
438 against: Option<&str>,
439 filter: Option<&str>,
440) -> crate::Result<EvalReport> {
441 let local = match against {
442 Some(_) => None,
443 None => Some(Host::new(app.clone(), Mode::Eval, None)?),
444 };
445 let mut report = EvalReport::default();
446 for eval in &app.inner.evals {
447 if filter.is_some_and(|filter| !eval.name.contains(filter)) {
448 continue;
449 }
450 let target = match (&local, against) {
451 (Some(host), _) => Target::Local(host.clone()),
452 (None, Some(base)) => Target::Remote {
453 base: base.trim_end_matches('/').to_string(),
454 client: reqwest::Client::new(),
455 },
456 (None, None) => bail!("no eval target"),
457 };
458 let mut cx = EvalCx::new(target);
459 let started = Instant::now();
460 let outcome = (eval.run)(&mut cx).await;
461 let result = EvalResult {
462 name: eval.name.to_string(),
463 passed: outcome.is_ok(),
464 error: outcome.err().map(|err| format!("{err:#}")),
465 duration: started.elapsed(),
466 };
467 match &result.error {
468 None => println!(" ✓ {} ({:.1?})", result.name, result.duration),
469 Some(err) => println!(" ✗ {} ({:.1?})\n {err}", result.name, result.duration),
470 }
471 report.results.push(result);
472 }
473 Ok(report)
474}
475
476#[cfg(test)]
477mod tests {
478 use super::*;
479
480 #[test]
481 fn remote_turns_read_tools_and_the_reply_from_canonical_events() {
482 let events = vec![
483 json!({ "type": "tool.started", "data": { "tool_call": { "name": "run_sql" } } }),
484 json!({ "type": "output.message.completed", "data": { "message": { "content": [{ "type": "text", "text": "Net of refunds." }] } } }),
485 json!({ "type": "output.message.completed", "data": { "message": { "content": [] } } }),
486 json!({ "type": "turn.completed", "data": {} }),
487 ];
488 assert_eq!(tools_called(&events), vec!["run_sql"]);
489 assert_eq!(final_response(&events), "Net of refunds.");
490 assert!(is_terminal(&events[3]));
491 }
492
493 #[test]
494 fn checks_explain_failures() {
495 let record = TurnRecord {
496 response: "Revenue was $10, net of refunds.".into(),
497 success: true,
498 tools: vec!["run_sql".into()],
499 ..TurnRecord::default()
500 };
501 let check = TurnCheck { turn: &record };
502 let check = check
503 .called_tool("run_sql")
504 .unwrap()
505 .reply_contains("NET OF REFUNDS")
506 .unwrap();
507 let err = check
508 .called_tool("delete_everything")
509 .err()
510 .unwrap()
511 .to_string();
512 assert!(err.contains("run_sql"), "{err}");
513 }
514}