use std::{
io::{Read, Write},
net::{TcpListener, TcpStream},
sync::{
Arc, Mutex,
atomic::{AtomicUsize, Ordering},
},
thread,
};
use basis::{AllowAll, CollectingSink, RunOutcome, Runtime, runtime::Wire};
use super::{offline, offline_runtime, write};
#[tokio::test]
async fn an_ephemeral_workspace_runs_a_turn_and_resumes_its_own_conversation() {
let dir = tempfile::tempdir().expect("tempdir");
write(&dir.path().join("AGENTS.md"), "house rules");
let endpoint = ScriptedEndpoint::start();
let workspace = offline(dir.path())
.with_runtime_builder(offline_runtime().with_base_url(&endpoint.base_url))
.open()
.await
.expect("opens");
let agent_id = {
let mut run = workspace.prepare("go").expect("mints");
let agent_id = run.agent_id().to_string();
let report = run
.execute_with_approver(CollectingSink::default(), AllowAll)
.await
.expect("the run completes");
assert!(matches!(report.outcome, RunOutcome::Ok));
agent_id
};
assert_eq!(
workspace
.resume(&agent_id, "again")
.expect("the store is alive as long as the workspace is")
.agent_id(),
agent_id,
"inside its workspace an ephemeral conversation behaves like any other"
);
}
#[tokio::test]
async fn two_runs_from_one_workspace_are_driven_concurrently() {
let dir = tempfile::tempdir().expect("tempdir");
write(&dir.path().join("AGENTS.md"), "house rules");
let endpoint = ScriptedEndpoint::start();
let workspace = offline(dir.path())
.with_runtime_builder(offline_runtime().with_base_url(&endpoint.base_url))
.open()
.await
.expect("opens");
let mut first = workspace.prepare("one").expect("mints");
let mut second = workspace.prepare("two").expect("mints");
let (left, right) = tokio::join!(
first.execute_with_approver(CollectingSink::default(), AllowAll),
second.execute_with_approver(CollectingSink::default(), AllowAll),
);
let left = left.expect("the first run completes");
let right = right.expect("the second run completes");
assert!(matches!(left.outcome, RunOutcome::Ok));
assert!(matches!(right.outcome, RunOutcome::Ok));
assert_eq!(
endpoint.served(),
2,
"each run makes its own request rather than sharing one"
);
let mut answers = [
left.final_message.expect("a final message"),
right.final_message.expect("a final message"),
];
answers.sort();
assert_eq!(answers, ["reply-1".to_string(), "reply-2".to_string()]);
}
#[tokio::test]
async fn a_custom_endpoint_is_addressed_on_the_chat_completions_wire() {
let dir = tempfile::tempdir().expect("tempdir");
write(&dir.path().join("AGENTS.md"), "house rules");
let endpoint = ScriptedEndpoint::start();
let workspace = offline(dir.path())
.with_runtime_builder(offline_runtime().with_base_url(&endpoint.base_url))
.open()
.await
.expect("opens");
let report = workspace
.prepare("go")
.expect("mints")
.execute_with_approver(CollectingSink::default(), AllowAll)
.await
.expect("the run completes");
assert!(matches!(report.outcome, RunOutcome::Ok));
assert_eq!(report.final_message.as_deref(), Some("reply-1"));
assert_eq!(endpoint.paths(), ["/v1/chat/completions"]);
}
#[tokio::test]
async fn a_responses_speaking_endpoint_is_reached_by_asking_for_it() {
let dir = tempfile::tempdir().expect("tempdir");
write(&dir.path().join("AGENTS.md"), "house rules");
let endpoint = ScriptedEndpoint::start_with(responses_sse_body);
let workspace = offline(dir.path())
.with_runtime_builder(
offline_runtime()
.with_base_url(&endpoint.base_url)
.with_wire(Wire::Responses),
)
.open()
.await
.expect("opens");
let report = workspace
.prepare("go")
.expect("mints")
.execute_with_approver(CollectingSink::default(), AllowAll)
.await
.expect("the run completes");
assert!(matches!(report.outcome, RunOutcome::Ok));
assert_eq!(report.final_message.as_deref(), Some("reply-1"));
assert_eq!(endpoint.paths(), ["/v1/responses"]);
}
#[tokio::test]
async fn a_published_url_ending_in_v1_is_not_doubled() {
let dir = tempfile::tempdir().expect("tempdir");
write(&dir.path().join("AGENTS.md"), "house rules");
let endpoint = ScriptedEndpoint::start();
let workspace = offline(dir.path())
.with_runtime_builder(offline_runtime().with_base_url(format!("{}v1", endpoint.base_url)))
.open()
.await
.expect("opens");
workspace
.prepare("go")
.expect("mints")
.execute_with_approver(CollectingSink::default(), AllowAll)
.await
.expect("the run completes");
assert_eq!(endpoint.paths(), ["/v1/chat/completions"]);
}
#[tokio::test]
async fn a_base_url_is_asked_with_the_key_resolution_found_or_no_header_at_all() {
let dir = tempfile::tempdir().expect("tempdir");
write(&dir.path().join("AGENTS.md"), "house rules");
let endpoint = ScriptedEndpoint::start();
let workspace = offline(dir.path())
.with_runtime_builder(
Runtime::builder()
.with_base_url(&endpoint.base_url)
.with_ephemeral_history(),
)
.open()
.await
.expect("opens");
workspace
.prepare("go")
.expect("mints")
.execute_with_approver(CollectingSink::default(), AllowAll)
.await
.expect("the run completes");
let exported = ["BASIS_API_KEY", "OPENAI_API_KEY"]
.into_iter()
.find_map(|var| std::env::var(var).ok().filter(|key| !key.trim().is_empty()));
assert_eq!(endpoint.bearers(), [exported]);
}
pub(crate) struct ScriptedEndpoint {
pub(crate) base_url: String,
served: Arc<AtomicUsize>,
seen: Arc<Mutex<Vec<Seen>>>,
}
#[derive(Clone, Debug)]
struct Seen {
path: String,
bearer: Option<String>,
}
impl ScriptedEndpoint {
pub(crate) fn start() -> Self {
Self::start_with(sse_body)
}
pub(crate) fn start_with(script: fn(usize) -> String) -> Self {
let listener = TcpListener::bind("127.0.0.1:0").expect("bind test endpoint");
let address = listener.local_addr().expect("read endpoint address");
let served = Arc::new(AtomicUsize::new(0));
let seen = Arc::new(Mutex::new(Vec::new()));
let counted = Arc::clone(&served);
let recorded = Arc::clone(&seen);
thread::spawn(move || {
while let Ok((stream, _)) = listener.accept() {
let counted = Arc::clone(&counted);
let recorded = Arc::clone(&recorded);
thread::spawn(move || answer(stream, script, &counted, &recorded));
}
});
Self {
base_url: format!("http://{address}/"),
served,
seen,
}
}
pub(crate) fn served(&self) -> usize {
self.served.load(Ordering::SeqCst)
}
pub(crate) fn paths(&self) -> Vec<String> {
self.seen().into_iter().map(|seen| seen.path).collect()
}
pub(crate) fn bearers(&self) -> Vec<Option<String>> {
self.seen().into_iter().map(|seen| seen.bearer).collect()
}
fn seen(&self) -> Vec<Seen> {
self.seen.lock().expect("seen").clone()
}
}
fn model_listing(request: &str) -> Option<String> {
let line = request.lines().next()?;
let target = line.split_whitespace().nth(1)?;
(line.starts_with("GET ") && target.ends_with("/models")).then(|| {
let body = r#"{"object":"list","data":[{"id":"test-model","object":"model"}]}"#;
format!(
"HTTP/1.1 200 OK\r\nconnection: close\r\ncontent-type: application/json\r\ncontent-length: {}\r\n\r\n{body}",
body.len()
)
})
}
fn answer(
mut stream: TcpStream,
script: fn(usize) -> String,
turns: &AtomicUsize,
recorded: &Mutex<Vec<Seen>>,
) {
let request = read_http_request(&mut stream);
if let Some(listing) = model_listing(&request) {
let _ = stream.write_all(listing.as_bytes());
return;
}
let body = script(turns.fetch_add(1, Ordering::SeqCst) + 1);
recorded.lock().expect("seen").push(Seen {
path: request_path(&request).to_string(),
bearer: request_bearer(&request),
});
let response = format!(
"HTTP/1.1 200 OK\r\nconnection: close\r\ncontent-type: text/event-stream\r\ncontent-length: {}\r\n\r\n{body}",
body.len()
);
let _ = stream.write_all(response.as_bytes());
}
fn request_bearer(request: &str) -> Option<String> {
request.lines().find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("authorization")
.then(|| value.trim().strip_prefix("Bearer ").map(str::to_string))
.flatten()
})
}
fn request_path(request: &str) -> &str {
request
.lines()
.next()
.and_then(|line| line.split_whitespace().nth(1))
.unwrap_or_default()
}
fn sse_body(index: usize) -> String {
[
format!(
r#"{{"id":"chatcmpl_{index}","model":"test-model","choices":[{{"index":0,"delta":{{"role":"assistant","content":"reply-{index}"}}}}]}}"#
),
format!(
r#"{{"id":"chatcmpl_{index}","choices":[{{"index":0,"delta":{{}},"finish_reason":"stop"}}]}}"#
),
"[DONE]".to_string(),
]
.iter()
.map(|event| format!("data: {event}\n\n"))
.collect()
}
fn responses_sse_body(index: usize) -> String {
[
format!(
r#"{{"type":"response.created","response":{{"id":"resp_{index}","model":"test-model","status":"in_progress"}}}}"#
),
r#"{"type":"response.output_item.added","output_index":0,"item":{"type":"message","content":[]}}"#.to_string(),
format!(
r#"{{"type":"response.output_item.done","output_index":0,"item":{{"type":"message","content":[{{"type":"output_text","text":"reply-{index}"}}]}}}}"#
),
format!(
r#"{{"type":"response.completed","response":{{"id":"resp_{index}","model":"test-model","status":"completed"}}}}"#
),
]
.iter()
.map(|event| format!("data: {event}\n\n"))
.collect()
}
fn read_http_request(stream: &mut TcpStream) -> String {
let mut bytes = Vec::new();
let mut buffer = [0_u8; 4096];
let mut header_end = None;
let mut content_length = 0_usize;
while let Ok(read) = stream.read(&mut buffer) {
if read == 0 {
break;
}
bytes.extend_from_slice(&buffer[..read]);
if header_end.is_none()
&& let Some(index) = bytes.windows(4).position(|window| window == b"\r\n\r\n")
{
let end = index + 4;
header_end = Some(end);
content_length = String::from_utf8_lossy(&bytes[..end])
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("content-length")
.then(|| value.trim().parse::<usize>().unwrap_or_default())
})
.unwrap_or_default();
}
if header_end.is_some_and(|end| bytes.len() >= end + content_length) {
break;
}
}
String::from_utf8_lossy(&bytes).into_owned()
}