Skip to main content

serve/
eval.rs

1//! `#[eval]`: conversations with assertions, run in-process or against a URL.
2//!
3//! Decision: the same eval code runs against the local binary (fast, offline
4//! with the simulator) and, with `--against <url>`, against a deployed build
5//! over the wire API. That makes an eval run a promotion gate for a preview
6//! deploy without a second test harness.
7
8use 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/// What to do when a tool asks for approval during an eval.
30#[derive(Clone, Debug, PartialEq, Eq)]
31pub enum OnApproval {
32    Approve,
33    Deny,
34}
35
36/// One finished turn, as the eval saw it.
37#[derive(Clone, Debug, Default)]
38pub struct TurnRecord {
39    pub response: String,
40    pub success: bool,
41    pub error: Option<String>,
42    /// Tools called, in order.
43    pub tools: Vec<String>,
44    /// Tools that asked for approval.
45    pub approvals: Vec<String>,
46    /// `ask_user` question sets answered (with their declared defaults).
47    pub questions: usize,
48    /// The turn's durable canonical events, as the wire API sends them.
49    pub events: Vec<Value>,
50}
51
52/// The eval's handle on one session.
53pub struct EvalCx {
54    target: Target,
55    session: Option<String>,
56    agent: Option<String>,
57    /// Last durable sequence seen; the next turn's events come after it.
58    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    /// Talk to this agent instead of the default one. Call before `send`.
89    pub fn agent(&mut self, name: impl Into<String>) -> &mut Self {
90        self.agent = Some(name.into());
91        self
92    }
93
94    /// How approval requests are answered (default: approve). Questions from
95    /// `ask_user` are answered with their declared defaults.
96    pub fn on_approval(&mut self, policy: OnApproval) -> &mut Self {
97        self.on_approval = policy;
98        self
99    }
100
101    /// Send a message and wait for the turn to finish.
102    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    /// The last turn, asserting it completed successfully.
129    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    /// The last turn, whatever its outcome.
141    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    /// In-process: send, answer approvals and questions as the host parks
175    /// them, and take the outcome from the turn itself.
176    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                        // The turn may have moved on (cancel) in between.
197                        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    /// Over the wire: send, then poll the session, answering its pending
236    /// approvals and questions, until it is idle and the turn's terminal
237    /// event is in the log.
238    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
327/// Tool names from canonical `tool.started` events, in order.
328fn 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
337/// The text of the turn's last completed output message that has any.
338fn 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
357/// Assertions on a completed turn. Each returns `Result<Self>` so they chain
358/// with `?`.
359pub 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    /// Case-insensitive substring match on the reply.
394    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/// Outcome of an eval run.
416#[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
435/// Run the app's evals, optionally only those whose name contains `filter`.
436pub(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}