use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::Mutex;
use std::time::Instant;
use io_harness::provider::{CompletionRequest, CompletionResponse, Usage};
use io_harness::{
run_with_observed, ApproveAll, EventKind, Flow, Observer, Policy, Provider, RunEvent, Session,
Store,
};
const CHUNKS: [&str; 5] = ["the ", "retry poli", "cy retries a ", "503", ""];
fn joined() -> String {
CHUNKS.concat()
}
#[derive(Default)]
struct Streamer {
calls: AtomicUsize,
streamed: AtomicUsize,
finished_at: Mutex<Option<Instant>>,
}
impl Provider for Streamer {
async fn complete(&self, _req: CompletionRequest) -> io_harness::Result<CompletionResponse> {
self.calls.fetch_add(1, Ordering::SeqCst);
Ok(response())
}
async fn complete_streaming(
&self,
_req: CompletionRequest,
on_token: &(dyn Fn(&str) + Send + Sync),
) -> io_harness::Result<CompletionResponse> {
self.calls.fetch_add(1, Ordering::SeqCst);
self.streamed.fetch_add(1, Ordering::SeqCst);
for chunk in CHUNKS {
if chunk.is_empty() {
continue;
}
on_token(chunk);
tokio::task::yield_now().await;
}
*self.finished_at.lock().unwrap() = Some(Instant::now());
Ok(response())
}
fn name(&self) -> &str {
"streamer"
}
}
#[derive(Default)]
struct Silent {
calls: AtomicUsize,
}
impl Provider for Silent {
async fn complete(&self, _req: CompletionRequest) -> io_harness::Result<CompletionResponse> {
self.calls.fetch_add(1, Ordering::SeqCst);
Ok(response())
}
fn name(&self) -> &str {
"silent"
}
}
fn response() -> CompletionResponse {
CompletionResponse {
text: Some(joined()),
usage: Some(Usage {
total_tokens: 5,
..Default::default()
}),
..Default::default()
}
}
#[derive(Default)]
struct Listener {
tokens: Mutex<Vec<String>>,
kinds: Mutex<Vec<String>>,
first_token_at: Mutex<Option<Instant>>,
stepped: AtomicBool,
tokens_before_step: AtomicUsize,
}
impl Observer for Listener {
fn event(&self, event: &RunEvent) -> Flow {
match &event.kind {
EventKind::Token { text } => {
self.tokens.lock().unwrap().push(text.clone());
let mut first = self.first_token_at.lock().unwrap();
if first.is_none() {
*first = Some(Instant::now());
}
if !self.stepped.load(Ordering::SeqCst) {
self.tokens_before_step.fetch_add(1, Ordering::SeqCst);
}
self.kinds.lock().unwrap().push("token".into());
}
EventKind::Step { .. } => {
self.stepped.store(true, Ordering::SeqCst);
self.kinds.lock().unwrap().push("step".into());
}
other => self
.kinds
.lock()
.unwrap()
.push(kind_name(other).to_string()),
}
Flow::Continue
}
}
fn kind_name(kind: &EventKind) -> &'static str {
match kind {
EventKind::Started { .. } => "started",
EventKind::Finished { .. } => "finished",
_ => "other",
}
}
fn workspace() -> tempfile::TempDir {
let dir = tempfile::tempdir().unwrap();
std::fs::write(dir.path().join("notes.md"), "# notes\n").unwrap();
dir
}
fn policy() -> Policy {
Policy::default().layer("stream-test").allow_read("*")
}
#[tokio::test]
async fn deltas_concatenate_to_the_final_text_and_arrive_before_the_step() {
let ws = workspace();
let store = Store::memory().unwrap();
let provider = Streamer::default();
let listener = Listener::default();
let mut session = Session::open(&store, ws.path()).unwrap();
let turn = session
.turn_observed(
"what does the retry policy retry?",
&provider,
&store,
&policy(),
&ApproveAll,
&listener,
)
.await
.unwrap();
let seen = listener.tokens.lock().unwrap().concat();
assert_eq!(
seen,
joined(),
"the deltas do not reconstruct the answer the provider returned"
);
assert_eq!(turn.reply.as_deref(), Some(joined().as_str()));
assert_eq!(
listener.tokens_before_step.load(Ordering::SeqCst),
listener.tokens.lock().unwrap().len(),
"a delta arrived after the step had already been committed"
);
let first = listener.first_token_at.lock().unwrap().expect("a delta");
let finished = provider.finished_at.lock().unwrap().expect("streamed");
assert!(
first < finished,
"the first delta was observed after the stream had already ended"
);
assert_eq!(provider.streamed.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn a_one_shot_run_emits_no_deltas_and_a_session_turn_does() {
let ws = workspace();
let store = Store::memory().unwrap();
let provider = Streamer::default();
let quiet = Listener::default();
let contract = io_harness::TaskContract::workspace("read the notes", ws.path());
run_with_observed(&contract, &provider, &store, &policy(), &ApproveAll, &quiet)
.await
.unwrap();
assert!(
quiet.tokens.lock().unwrap().is_empty(),
"a one-shot run started emitting Token events"
);
assert_eq!(
provider.streamed.load(Ordering::SeqCst),
0,
"a one-shot run asked the provider to stream"
);
assert_eq!(
*quiet.kinds.lock().unwrap(),
vec!["started", "step", "finished"]
);
let loud = Listener::default();
let mut session = Session::open(&store, ws.path()).unwrap();
session
.turn_observed(
"read the notes",
&provider,
&store,
&policy(),
&ApproveAll,
&loud,
)
.await
.unwrap();
assert_eq!(loud.tokens.lock().unwrap().len(), CHUNKS.len() - 1);
assert_eq!(provider.streamed.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn a_provider_without_streaming_yields_one_delta_carrying_the_whole_text() {
let ws = workspace();
let store = Store::memory().unwrap();
let provider = Silent::default();
let listener = Listener::default();
let mut session = Session::open(&store, ws.path()).unwrap();
let turn = session
.turn_observed(
"read the notes",
&provider,
&store,
&policy(),
&ApproveAll,
&listener,
)
.await
.unwrap();
assert_eq!(*listener.tokens.lock().unwrap(), vec![joined()]);
assert_eq!(turn.reply.as_deref(), Some(joined().as_str()));
assert_eq!(provider.calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn a_quiet_turn_never_enters_the_streaming_path() {
let ws = workspace();
let store = Store::memory().unwrap();
let provider = Streamer::default();
let mut session = Session::open(&store, ws.path()).unwrap();
session
.turn("read the notes", &provider, &store, &policy(), &ApproveAll)
.await
.unwrap();
assert_eq!(provider.calls.load(Ordering::SeqCst), 1);
assert_eq!(
provider.streamed.load(Ordering::SeqCst),
0,
"a turn with no observer still asked the provider to stream"
);
}