use std::collections::HashMap;
use std::sync::Arc;
use tokio::io::{AsyncBufReadExt, AsyncRead, AsyncWrite, AsyncWriteExt, BufReader};
use tokio::sync::{Mutex, oneshot};
use tokio::task::JoinSet;
use crate::eval::Eval;
use crate::protocol::{
AxisInfo, CancelParams, CancelResult, EvalInfo, EventParams, ExecuteResult, InitializeResult,
ListResult, ListSamplesParams, ListSamplesResult, Notification, PROTOCOL_VERSION, Request,
Response, RpcError, RunParams, RunResult, SampleInfo, ScoreParams, TargetInfo,
TranscriptSummary, capabilities, codes, event,
};
use crate::registry::registered_evals;
use crate::runner::{aggregate_value, execute_case, run_case, score_transcript, verdict};
type SharedWriter = Arc<Mutex<Box<dyn AsyncWrite + Send + Unpin>>>;
type Inflight = Arc<std::sync::Mutex<HashMap<u64, oneshot::Sender<()>>>>;
pub const DEFAULT_PAGE_SIZE: usize = 500;
pub struct Study {
name: String,
evals: Vec<Eval>,
page_size: Option<usize>,
}
impl Default for Study {
fn default() -> Self {
Self::new()
}
}
impl Study {
pub fn new() -> Self {
Self {
name: env!("CARGO_PKG_NAME").into(),
evals: Vec::new(),
page_size: Some(DEFAULT_PAGE_SIZE),
}
}
pub fn registered() -> Self {
Self::new().evals(registered_evals())
}
pub fn eval(mut self, eval: Eval) -> Self {
self.evals.push(eval);
self
}
pub fn evals(mut self, evals: impl IntoIterator<Item = Eval>) -> Self {
self.evals.extend(evals);
self
}
pub fn named(mut self, name: impl Into<String>) -> Self {
self.name = name.into();
self
}
pub fn page_size(mut self, size: usize) -> Self {
self.page_size = (size > 0).then_some(size);
self
}
pub async fn serve(self) -> std::io::Result<()> {
self.serve_io(tokio::io::stdin(), tokio::io::stdout()).await
}
pub async fn serve_io<R, W>(self, reader: R, writer: W) -> std::io::Result<()>
where
R: AsyncRead + Unpin,
W: AsyncWrite + Send + Unpin + 'static,
{
let mut lines = BufReader::new(reader).lines();
let out: SharedWriter = Arc::new(Mutex::new(Box::new(writer)));
let me = Arc::new(self);
let mut tasks: JoinSet<()> = JoinSet::new();
let inflight: Inflight = Default::default();
while let Some(line) = lines.next_line().await? {
if line.trim().is_empty() {
continue;
}
let request: Request = match serde_json::from_str(&line) {
Ok(req) => req,
Err(e) => {
write_line(&out, &Notification::log(format!("bad request: {e}"), 0)).await?;
continue;
}
};
if request.method == "cancel" {
let response = cancel(&request, &inflight);
write_line(&out, &response).await?;
continue;
}
let (cancel_tx, cancel_rx) = oneshot::channel::<()>();
inflight
.lock()
.expect("inflight mutex poisoned")
.insert(request.id, cancel_tx);
let me = me.clone();
let out = out.clone();
let inflight = inflight.clone();
tasks.spawn(async move {
let id = request.id;
let response = tokio::select! {
resp = me.dispatch(&request, &out) => resp,
_ = cancel_rx => Response::err(id, "cancelled"),
};
inflight
.lock()
.expect("inflight mutex poisoned")
.remove(&id);
let _ = write_line(&out, &response).await;
});
}
while tasks.join_next().await.is_some() {}
Ok(())
}
async fn dispatch(&self, request: &Request, stdout: &SharedWriter) -> Response {
match request.method.as_str() {
"initialize" => Response::ok(
request.id,
json(&InitializeResult {
protocol_version: PROTOCOL_VERSION.into(),
study: self.name.clone(),
evals: self.evals.len(),
study_version: Some(env!("CARGO_PKG_VERSION").into()),
capabilities: vec![
capabilities::AXES.into(),
capabilities::EVENTS.into(),
capabilities::USAGE.into(),
capabilities::EXECUTE.into(),
capabilities::SCORE.into(),
capabilities::TRIALS.into(),
capabilities::CANCEL.into(),
capabilities::PAGINATE.into(),
capabilities::TRAJECTORY.into(),
],
capability_params: advertised_capability_params(),
}),
),
"list" => Response::ok(request.id, json(&self.list())),
"list_samples" => {
let params: ListSamplesParams = match serde_json::from_value(request.params.clone())
{
Ok(p) => p,
Err(e) => {
return Response::err(request.id, format!("bad list_samples params: {e}"));
}
};
match self.list_samples(¶ms) {
Ok(result) => Response::ok(request.id, json(&result)),
Err(e) => Response::err(request.id, e),
}
}
"run" => {
let params: RunParams = match serde_json::from_value(request.params.clone()) {
Ok(p) => p,
Err(e) => {
return Response::err_with(
request.id,
RpcError::new(format!("bad run params: {e}"))
.with_code(codes::INVALID_PARAMS),
);
}
};
let _ = write_line(stdout, &case_event(request.id, ¶ms, event::STARTED)).await;
let result = self.run(¶ms).await;
let _ = write_line(stdout, &case_event(request.id, ¶ms, event::FINISHED)).await;
match result {
Ok(result) => Response::ok(request.id, json(&result)),
Err(e) => Response::err(request.id, e),
}
}
"execute" => {
let params: RunParams = match serde_json::from_value(request.params.clone()) {
Ok(p) => p,
Err(e) => {
return Response::err_with(
request.id,
RpcError::new(format!("bad execute params: {e}"))
.with_code(codes::INVALID_PARAMS),
);
}
};
let _ = write_line(stdout, &case_event(request.id, ¶ms, event::STARTED)).await;
let result = self.execute(¶ms).await;
let _ = write_line(stdout, &case_event(request.id, ¶ms, event::FINISHED)).await;
match result {
Ok(result) => Response::ok(request.id, json(&result)),
Err(e) => Response::err(request.id, e),
}
}
"score" => {
let params: ScoreParams = match serde_json::from_value(request.params.clone()) {
Ok(p) => p,
Err(e) => {
return Response::err_with(
request.id,
RpcError::new(format!("bad score params: {e}"))
.with_code(codes::INVALID_PARAMS),
);
}
};
match self.score(¶ms).await {
Ok(result) => Response::ok(request.id, json(&result)),
Err(e) => Response::err(request.id, e),
}
}
other => Response::err_with(
request.id,
RpcError::new(format!("unknown method: {other}"))
.with_code(codes::METHOD_NOT_FOUND),
),
}
}
pub fn list(&self) -> ListResult {
let evals = self
.evals
.iter()
.map(|eval| {
let (samples, next_cursor) = self.sample_page(eval, 0);
EvalInfo {
name: eval.name.clone(),
description: eval.description.clone(),
samples,
next_cursor,
scorers: eval.scorers.iter().map(|s| s.name()).collect(),
targets: eval
.targets
.iter()
.map(|m| TargetInfo {
label: m.label.clone(),
provider: m.provider.clone(),
available: m.available,
metadata: m.metadata.clone(),
})
.collect(),
axes: eval
.axes
.iter()
.map(|a| AxisInfo {
name: a.name.clone(),
values: a.values.clone(),
})
.collect(),
max_turns: eval.max_turns,
trials: eval.trials,
seed: eval.seed,
metadata: eval.metadata.clone(),
}
})
.collect();
ListResult { evals }
}
pub fn list_samples(&self, params: &ListSamplesParams) -> Result<ListSamplesResult, String> {
let eval = self
.evals
.iter()
.find(|e| e.name == params.eval)
.ok_or_else(|| format!("no such eval: {}", params.eval))?;
let offset: usize = params
.cursor
.parse()
.map_err(|_| format!("bad cursor: {}", params.cursor))?;
let (samples, next_cursor) = self.sample_page(eval, offset);
Ok(ListSamplesResult {
samples,
next_cursor,
})
}
fn sample_page(&self, eval: &Eval, offset: usize) -> (Vec<SampleInfo>, Option<String>) {
let all = &eval.dataset.samples;
let start = offset.min(all.len());
let end = match self.page_size {
Some(p) => start.saturating_add(p).min(all.len()),
None => all.len(),
};
let page = all[start..end]
.iter()
.map(|s| SampleInfo {
id: s.id.clone(),
tags: s.tags.clone(),
metadata: s.metadata.clone(),
})
.collect();
let next = (end < all.len()).then(|| end.to_string());
(page, next)
}
async fn run(&self, params: &RunParams) -> Result<RunResult, String> {
let (eval, sample, target) = self.locate(¶ms.eval, ¶ms.sample, ¶ms.target)?;
if !target.available {
return Ok(skipped_result(params, sample));
}
let outcome = run_case(eval, sample, target, ¶ms.params, params.trial()).await;
Ok(RunResult {
eval: outcome.eval,
sample: outcome.sample_id,
target: outcome.target,
params: outcome.params,
trial: params.trial,
trials: params.trials,
seed: params.seed,
input: sample.input.clone(),
expected: sample.expected.clone(),
passed: outcome.passed,
aggregate: outcome.aggregate,
scores: outcome.scores,
transcript: TranscriptSummary::of(&outcome.transcript),
skipped: false,
})
}
async fn execute(&self, params: &RunParams) -> Result<ExecuteResult, String> {
let (eval, sample, target) = self.locate(¶ms.eval, ¶ms.sample, ¶ms.target)?;
if !target.available {
return Ok(ExecuteResult {
eval: params.eval.clone(),
sample: params.sample.clone(),
target: params.target.clone(),
params: params.params.clone(),
trial: params.trial,
trials: params.trials,
seed: params.seed,
transcript: Default::default(),
skipped: true,
});
}
let transcript = execute_case(eval, sample, target, ¶ms.params, params.trial()).await;
Ok(ExecuteResult {
eval: params.eval.clone(),
sample: params.sample.clone(),
target: params.target.clone(),
params: params.params.clone(),
trial: params.trial,
trials: params.trials,
seed: params.seed,
transcript,
skipped: false,
})
}
async fn score(&self, params: &ScoreParams) -> Result<RunResult, String> {
let eval = self
.evals
.iter()
.find(|e| e.name == params.eval)
.ok_or_else(|| format!("no such eval: {}", params.eval))?;
let sample = eval
.dataset
.samples
.iter()
.find(|s| s.id == params.sample)
.ok_or_else(|| format!("no such sample: {}/{}", params.eval, params.sample))?;
let mut transcript = params.transcript.clone();
transcript.project_trajectory();
let scores = score_transcript(eval, sample, &transcript).await;
Ok(RunResult {
eval: params.eval.clone(),
sample: params.sample.clone(),
target: params.target.clone(),
params: params.params.clone(),
trial: params.trial,
trials: params.trials,
seed: params.seed,
input: sample.input.clone(),
expected: sample.expected.clone(),
passed: verdict(&scores),
aggregate: aggregate_value(&scores),
scores,
transcript: TranscriptSummary::of(&transcript),
skipped: false,
})
}
fn locate(
&self,
eval: &str,
sample: &str,
target: &str,
) -> Result<(&Eval, &crate::Sample, &crate::Target), String> {
let e = self
.evals
.iter()
.find(|e| e.name == eval)
.ok_or_else(|| format!("no such eval: {eval}"))?;
let s = e
.dataset
.samples
.iter()
.find(|s| s.id == sample)
.ok_or_else(|| format!("no such sample: {eval}/{sample}"))?;
let m = e
.targets
.iter()
.find(|m| m.label == target)
.ok_or_else(|| format!("no such target: {eval}@{target}"))?;
Ok((e, s, m))
}
}
fn skipped_result(params: &RunParams, sample: &crate::Sample) -> RunResult {
RunResult {
eval: params.eval.clone(),
sample: params.sample.clone(),
target: params.target.clone(),
params: params.params.clone(),
trial: params.trial,
trials: params.trials,
seed: params.seed,
input: sample.input.clone(),
expected: sample.expected.clone(),
passed: false,
aggregate: 0.0,
scores: Vec::new(),
transcript: TranscriptSummary::default(),
skipped: true,
}
}
fn cancel(request: &Request, inflight: &Inflight) -> Response {
let params: CancelParams = match serde_json::from_value(request.params.clone()) {
Ok(p) => p,
Err(e) => {
return Response::err_with(
request.id,
RpcError::new(format!("bad cancel params: {e}")).with_code(codes::INVALID_PARAMS),
);
}
};
let cancelled = inflight
.lock()
.expect("inflight mutex poisoned")
.remove(¶ms.id)
.is_some_and(|tx| tx.send(()).is_ok());
Response::ok(request.id, json(&CancelResult { cancelled }))
}
fn json<T: serde::Serialize>(value: &T) -> serde_json::Value {
serde_json::to_value(value).unwrap_or(serde_json::Value::Null)
}
fn advertised_capability_params() -> crate::Metadata {
let modalities = serde_json::json!(["text", "image", "audio", "file", "json"]);
crate::Metadata::from([
(
capabilities::EVENTS.to_string(),
serde_json::json!({ "kinds": [event::STARTED, event::FINISHED] }),
),
(
"modalities".to_string(),
serde_json::json!({ "input": modalities, "output": modalities }),
),
(
capabilities::TRAJECTORY.to_string(),
serde_json::json!({
"format": crate::trajectory::ATIF_FORMAT,
"version": crate::trajectory::ATIF_VERSION.trim_start_matches("ATIF-v"),
}),
),
])
}
fn case_event(req_id: u64, p: &RunParams, kind: &str) -> Notification {
Notification::event(EventParams {
request_id: req_id,
eval: p.eval.clone(),
sample: p.sample.clone(),
target: p.target.clone(),
params: p.params.clone(),
kind: kind.into(),
..Default::default()
})
}
async fn write_line<T: serde::Serialize>(out: &SharedWriter, value: &T) -> std::io::Result<()> {
let mut buf = serde_json::to_vec(value).unwrap_or_default();
buf.push(b'\n');
let mut out = out.lock().await;
out.write_all(&buf).await?;
out.flush().await
}
#[cfg(test)]
mod tests {
use super::*;
use crate::scorer::contains;
use crate::subject::subject_fn;
use crate::{Eval, Sample, Target, Transcript};
use serde_json::json;
fn study() -> Study {
Study::new().eval(
Eval::new("greet")
.describe("greeting eval")
.meta("suite", "smoke")
.add_sample(
Sample::new("hi", "say hi")
.tag("smoke")
.meta("difficulty", "easy"),
)
.subject(subject_fn(|_, _| async {
Transcript::response("hi there")
}))
.scorer(contains("hi"))
.targets([Target::sim().meta("agent", "demo")])
.build(),
)
}
#[test]
fn initialize_advertises_capability_params() {
let init = advertised_capability_params();
let info = InitializeResult {
protocol_version: PROTOCOL_VERSION.into(),
study: "x".into(),
evals: 0,
study_version: None,
capabilities: vec![capabilities::EVENTS.into()],
capability_params: init,
};
let modalities = info.capability_param("modalities").unwrap();
assert!(
modalities["input"]
.as_array()
.unwrap()
.iter()
.any(|m| m == "image")
);
assert!(info.capability_param("events").unwrap()["kinds"][0] == "started");
assert!(info.capability_param("absent").is_none());
let back: InitializeResult =
serde_json::from_str(&serde_json::to_string(&info).unwrap()).unwrap();
assert_eq!(back.capability_params, info.capability_params);
}
#[test]
fn list_advertises_everything() {
let listing = study().list();
assert_eq!(listing.evals.len(), 1);
let e = &listing.evals[0];
assert_eq!(e.description, "greeting eval");
assert_eq!(e.metadata.get("suite").unwrap(), "smoke");
assert_eq!(e.samples[0].tags, vec!["smoke"]);
assert_eq!(e.samples[0].metadata.get("difficulty").unwrap(), "easy");
assert_eq!(e.targets[0].label, "sim");
assert!(e.targets[0].available);
assert_eq!(e.targets[0].metadata.get("agent").unwrap(), "demo");
}
fn big_study(samples: usize, page: usize) -> Study {
let mut eval = Eval::new("big")
.subject(subject_fn(|_, _| async { Transcript::response("ok") }))
.scorer(contains("ok"));
for i in 0..samples {
eval = eval.add_sample(Sample::new(format!("s{i}"), "go"));
}
Study::new().page_size(page).eval(eval.build())
}
#[test]
fn list_paginates_first_page_with_cursor() {
let s = big_study(250, 100);
let listing = s.list();
let e = &listing.evals[0];
assert_eq!(e.samples.len(), 100);
assert_eq!(e.samples[0].id, "s0");
assert_eq!(e.next_cursor.as_deref(), Some("100"));
}
#[test]
fn list_samples_walks_every_page_then_stops() {
let s = big_study(250, 100);
let mut ids: Vec<String> = s.list().evals[0]
.samples
.iter()
.map(|x| x.id.clone())
.collect();
let mut cursor = s.list().evals[0].next_cursor.clone();
while let Some(c) = cursor {
let page = s
.list_samples(&ListSamplesParams {
eval: "big".into(),
cursor: c,
})
.unwrap();
ids.extend(page.samples.into_iter().map(|x| x.id));
cursor = page.next_cursor;
}
assert_eq!(ids.len(), 250);
assert_eq!(ids[0], "s0");
assert_eq!(ids[249], "s249");
let last = s
.list_samples(&ListSamplesParams {
eval: "big".into(),
cursor: "200".into(),
})
.unwrap();
assert_eq!(last.samples.len(), 50);
assert!(last.next_cursor.is_none());
}
#[test]
fn page_size_zero_disables_pagination() {
let s = big_study(250, 0);
let e = &s.list().evals[0];
assert_eq!(e.samples.len(), 250);
assert!(e.next_cursor.is_none());
}
#[test]
fn list_samples_rejects_unknown_eval_and_bad_cursor() {
let s = big_study(10, 5);
assert!(
s.list_samples(&ListSamplesParams {
eval: "nope".into(),
cursor: "0".into(),
})
.is_err()
);
assert!(
s.list_samples(&ListSamplesParams {
eval: "big".into(),
cursor: "xyz".into(),
})
.is_err()
);
let past = s
.list_samples(&ListSamplesParams {
eval: "big".into(),
cursor: "999".into(),
})
.unwrap();
assert!(past.samples.is_empty());
assert!(past.next_cursor.is_none());
}
#[tokio::test]
async fn run_scores_a_case() {
let params = RunParams {
eval: "greet".into(),
sample: "hi".into(),
target: "sim".into(),
params: Default::default(),
trial: 0,
trials: 0,
seed: None,
};
let result = study().run(¶ms).await.unwrap();
assert!(result.passed);
assert_eq!(result.transcript.final_response, "hi there");
}
#[tokio::test]
async fn run_echoes_trial_identity_and_threads_seed() {
let s = Study::new().eval(
Eval::new("rng")
.sample("a", "x")
.trials(4)
.subject(subject_fn(|_, cx| async move {
Transcript::response(format!("seed={:?}", cx.seed()))
}))
.scorer(contains("seed="))
.build(),
);
let params = RunParams {
eval: "rng".into(),
sample: "a".into(),
target: "sim".into(),
params: Default::default(),
trial: 2,
trials: 4,
seed: Some(77),
};
let result = s.run(¶ms).await.unwrap();
assert_eq!(result.trial, 2);
assert_eq!(result.trials, 4);
assert_eq!(result.seed, Some(77));
assert_eq!(result.key(), "rng/a@sim#2");
assert!(result.transcript.final_response.contains("77"));
}
#[tokio::test]
async fn run_and_score_carry_sample_input_and_expected() {
let s = Study::new().eval(
Eval::new("qa")
.add_sample(Sample::new("a", "what is 6*7?").expected("42"))
.subject(subject_fn(|_, _| async { Transcript::response("42") }))
.scorer(contains("42"))
.targets([
Target::sim(),
Target::new("down", "anthropic", "claude").available(false),
])
.build(),
);
let case = |target: &str| RunParams {
eval: "qa".into(),
sample: "a".into(),
target: target.into(),
params: Default::default(),
trial: 0,
trials: 0,
seed: None,
};
let ran = s.run(&case("sim")).await.unwrap();
assert_eq!(ran.input, vec!["what is 6*7?".to_string()]);
assert_eq!(ran.expected, Some(json!("42")));
let skipped = s.run(&case("down")).await.unwrap();
assert!(skipped.skipped);
assert_eq!(skipped.input, vec!["what is 6*7?".to_string()]);
assert_eq!(skipped.expected, Some(json!("42")));
let scored = s
.score(&ScoreParams {
eval: "qa".into(),
sample: "a".into(),
target: "sim".into(),
params: Default::default(),
trial: 0,
trials: 0,
seed: None,
transcript: Transcript::response("42"),
})
.await
.unwrap();
assert_eq!(scored.input, vec!["what is 6*7?".to_string()]);
assert_eq!(scored.expected, Some(json!("42")));
}
#[tokio::test]
async fn run_rejects_unknown_eval() {
let params = RunParams {
eval: "nope".into(),
sample: "hi".into(),
target: "sim".into(),
params: Default::default(),
trial: 0,
trials: 0,
seed: None,
};
assert!(study().run(¶ms).await.is_err());
}
#[tokio::test]
async fn execute_returns_full_transcript_without_scoring() {
let params = RunParams {
eval: "greet".into(),
sample: "hi".into(),
target: "sim".into(),
params: Default::default(),
trial: 0,
trials: 0,
seed: None,
};
let captured = study().execute(¶ms).await.unwrap();
assert!(!captured.skipped);
assert_eq!(captured.transcript.final_response, "hi there");
}
#[tokio::test]
async fn execute_then_score_matches_run() {
let s = study();
let rp = RunParams {
eval: "greet".into(),
sample: "hi".into(),
target: "sim".into(),
params: Default::default(),
trial: 0,
trials: 0,
seed: None,
};
let fused = s.run(&rp).await.unwrap();
let captured = s.execute(&rp).await.unwrap();
let sp = ScoreParams {
eval: captured.eval.clone(),
sample: captured.sample.clone(),
target: captured.target.clone(),
params: captured.params.clone(),
trial: captured.trial,
trials: captured.trials,
seed: captured.seed,
transcript: captured.transcript.clone(),
};
let split = s.score(&sp).await.unwrap();
assert_eq!(split.passed, fused.passed);
assert_eq!(split.aggregate, fused.aggregate);
assert_eq!(split.scores, fused.scores);
assert_eq!(
split.transcript.final_response,
fused.transcript.final_response
);
}
#[tokio::test]
async fn score_is_repeatable_for_rescoring() {
let s = study();
let sp = ScoreParams {
eval: "greet".into(),
sample: "hi".into(),
target: "sim".into(),
params: Default::default(),
trial: 0,
trials: 0,
seed: None,
transcript: Transcript::response("hi there"),
};
let first = s.score(&sp).await.unwrap();
let second = s.score(&sp).await.unwrap();
assert_eq!(first.scores, second.scores);
assert!(first.passed && second.passed);
}
#[tokio::test]
async fn score_normalizes_a_trajectory_only_transcript() {
use crate::scorer::{tool_called, tool_calls_within};
use crate::trajectory::{Agent, Step, StepSource, ToolCall, Trajectory};
let s = Study::new().eval(
Eval::new("traj")
.sample("hi", "say hi")
.subject(subject_fn(|_, _| async { Transcript::response("unused") }))
.scorer(contains("hi there"))
.scorer(tool_called("search"))
.scorer(tool_calls_within(1))
.build(),
);
let mut trajectory = Trajectory::new(Agent::new("external-agent", "1.0"));
let mut step = Step::new(1, StepSource::Agent, "hi there");
step.tool_calls = vec![ToolCall::new(
"c1",
"search",
serde_json::json!({"q": "hi"}),
)];
trajectory.steps.push(step);
let wire = serde_json::json!({ "trajectory": trajectory });
let transcript: Transcript = serde_json::from_value(wire).unwrap();
let result = s
.score(&ScoreParams {
eval: "traj".into(),
sample: "hi".into(),
target: "sim".into(),
params: Default::default(),
trial: 0,
trials: 0,
seed: None,
transcript,
})
.await
.unwrap();
assert!(result.passed, "scores: {:?}", result.scores);
assert!(result.scores.iter().all(|sc| sc.pass));
assert_eq!(result.transcript.final_response, "hi there");
assert_eq!(result.transcript.tool_calls, vec!["search"]);
assert_eq!(result.transcript.tool_calls_count, 1);
}
#[tokio::test]
async fn score_rejects_unknown_eval() {
let sp = ScoreParams {
eval: "nope".into(),
sample: "hi".into(),
target: "sim".into(),
params: Default::default(),
trial: 0,
trials: 0,
seed: None,
transcript: Transcript::response("x"),
};
assert!(study().score(&sp).await.is_err());
}
fn slow_study() -> Study {
Study::new().eval(
Eval::new("slow")
.add_sample(Sample::new("s", "go"))
.subject(subject_fn(|_, _| async {
tokio::time::sleep(std::time::Duration::from_secs(30)).await;
Transcript::response("done")
}))
.scorer(contains("done"))
.build(),
)
}
#[tokio::test]
async fn cancel_aborts_inflight_run() {
use std::time::Duration;
use tokio::io::AsyncWriteExt;
let (mut host_w, study_r) = tokio::io::duplex(8192);
let (study_w, host_r) = tokio::io::duplex(8192);
let server = tokio::spawn(async move { slow_study().serve_io(study_r, study_w).await });
let mut reader = BufReader::new(host_r).lines();
host_w
.write_all(
b"{\"id\":1,\"method\":\"run\",\"params\":\
{\"eval\":\"slow\",\"sample\":\"s\",\"target\":\"sim\"}}\n",
)
.await
.unwrap();
host_w
.write_all(b"{\"id\":2,\"method\":\"cancel\",\"params\":{\"id\":1}}\n")
.await
.unwrap();
host_w.flush().await.unwrap();
let (mut run_resp, mut cancel_resp) = (None, None);
while run_resp.is_none() || cancel_resp.is_none() {
let line = tokio::time::timeout(Duration::from_secs(5), reader.next_line())
.await
.expect("response did not arrive — cancel did not abort the run")
.expect("read line")
.expect("study closed early");
let v: serde_json::Value = serde_json::from_str(&line).unwrap();
match v.get("id").and_then(|i| i.as_u64()) {
Some(1) => run_resp = Some(v),
Some(2) => cancel_resp = Some(v),
_ => {} }
}
assert_eq!(cancel_resp.unwrap()["result"]["cancelled"], json!(true));
let msg = run_resp.unwrap()["error"]["message"]
.as_str()
.unwrap()
.to_string();
assert!(msg.contains("cancelled"), "run error was {msg:?}");
drop(host_w);
let _ = server.await;
}
#[tokio::test]
async fn cancel_unknown_id_is_benign_miss() {
use tokio::io::AsyncWriteExt;
let (mut host_w, study_r) = tokio::io::duplex(8192);
let (study_w, host_r) = tokio::io::duplex(8192);
let server = tokio::spawn(async move { study().serve_io(study_r, study_w).await });
let mut reader = BufReader::new(host_r).lines();
host_w
.write_all(b"{\"id\":9,\"method\":\"cancel\",\"params\":{\"id\":123}}\n")
.await
.unwrap();
host_w.flush().await.unwrap();
let line = reader.next_line().await.unwrap().unwrap();
let v: serde_json::Value = serde_json::from_str(&line).unwrap();
assert_eq!(v["id"], json!(9));
assert_eq!(v["result"]["cancelled"], json!(false));
drop(host_w);
let _ = server.await;
}
}