#![allow(dead_code)]
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use futures::StreamExt;
use rai_sdk::provider::ProviderStreamEvent;
use rai_sdk::{ClientBuilder, Error, Result, RetryConfig};
use wiremock::{MockServer, Request, Respond, ResponseTemplate};
pub const OPENAI_TEST_KEY: &str = "test-openai-key";
pub const ANTHROPIC_TEST_KEY: &str = "test-anthropic-key";
pub const OPENROUTER_TEST_KEY: &str = "test-openrouter-key";
pub fn openai_builder(base_url: &str) -> ClientBuilder {
ClientBuilder::new()
.openai_key(OPENAI_TEST_KEY)
.openai_base_url(base_url.to_string())
.no_retry()
}
pub fn anthropic_builder(base_url: &str) -> ClientBuilder {
ClientBuilder::new()
.anthropic_key(ANTHROPIC_TEST_KEY)
.anthropic_base_url(base_url.to_string())
.no_retry()
}
pub fn openai_compatible_builder(base_url: &str) -> ClientBuilder {
ClientBuilder::new()
.openai_compatible_base_url(base_url.to_string())
.no_retry()
}
pub fn openrouter_builder(base_url: &str) -> ClientBuilder {
ClientBuilder::new()
.openrouter_key(OPENROUTER_TEST_KEY)
.openrouter_base_url(base_url.to_string())
.no_retry()
}
pub fn fast_retry(max_retries: u32) -> RetryConfig {
RetryConfig::new()
.with_max_retries(max_retries)
.with_initial_delay(Duration::from_millis(2))
.with_max_delay(Duration::from_millis(50))
.with_backoff_multiplier(2.0)
.with_jitter(false)
}
#[derive(Clone)]
pub struct Step {
status: u16,
content_type: &'static str,
body: String,
}
impl Step {
pub fn json(status: u16, body: serde_json::Value) -> Self {
Self {
status,
content_type: "application/json",
body: body.to_string(),
}
}
pub fn ok(body: serde_json::Value) -> Self {
Self::json(200, body)
}
pub fn raw(status: u16, content_type: &'static str, body: impl Into<String>) -> Self {
Self {
status,
content_type,
body: body.into(),
}
}
pub fn sse(body: impl Into<String>) -> Self {
Self::raw(200, "text/event-stream", body)
}
fn template(&self) -> ResponseTemplate {
ResponseTemplate::new(self.status)
.insert_header("content-type", self.content_type)
.set_body_string(self.body.clone())
}
}
pub struct Script {
steps: Vec<Step>,
calls: AtomicUsize,
}
impl Script {
pub fn new(steps: Vec<Step>) -> Self {
assert!(!steps.is_empty(), "a script needs at least one step");
Self {
steps,
calls: AtomicUsize::new(0),
}
}
}
impl Respond for Script {
fn respond(&self, _request: &Request) -> ResponseTemplate {
let index = self.calls.fetch_add(1, Ordering::SeqCst);
let step = self
.steps
.get(index)
.unwrap_or_else(|| self.steps.last().expect("script is non-empty"));
step.template()
}
}
pub async fn received_json_bodies(server: &MockServer) -> Vec<serde_json::Value> {
server
.received_requests()
.await
.expect("wiremock request recording should be enabled")
.iter()
.map(|request| {
serde_json::from_slice(&request.body).unwrap_or_else(|error| {
panic!(
"recorded request body should be JSON: {error}\nbody: {}",
String::from_utf8_lossy(&request.body)
)
})
})
.collect()
}
pub async fn received_header(server: &MockServer, index: usize, header: &str) -> Option<String> {
let requests = server
.received_requests()
.await
.expect("wiremock request recording should be enabled");
let request = requests
.get(index)
.unwrap_or_else(|| panic!("expected at least {} recorded request(s)", index + 1));
request
.headers
.get(header)
.map(|value| value.to_str().expect("header should be valid UTF-8").into())
}
pub async fn request_count(server: &MockServer) -> usize {
server
.received_requests()
.await
.expect("wiremock request recording should be enabled")
.len()
}
pub fn sse_body(events: &[&str]) -> String {
events
.iter()
.map(|event| format!("{event}\n\n"))
.collect::<String>()
}
pub fn data_event(payload: serde_json::Value) -> String {
format!("data: {payload}")
}
pub fn named_event(event: &str, payload: serde_json::Value) -> String {
format!("event: {event}\ndata: {payload}")
}
pub async fn collect_events<S>(stream: S) -> Vec<ProviderStreamEvent>
where
S: futures::Stream<Item = Result<ProviderStreamEvent>>,
{
let mut stream = Box::pin(stream);
let mut events = Vec::new();
while let Some(event) = stream.next().await {
events.push(event.expect("stream should not yield an error"));
}
events
}
pub async fn collect_results<S>(stream: S) -> Vec<Result<ProviderStreamEvent>>
where
S: futures::Stream<Item = Result<ProviderStreamEvent>>,
{
let mut stream = Box::pin(stream);
let mut events = Vec::new();
while let Some(event) = stream.next().await {
events.push(event);
}
events
}
pub fn describe_event(event: &ProviderStreamEvent) -> String {
match event {
ProviderStreamEvent::Text(text) => format!("text:{text}"),
ProviderStreamEvent::ToolCallStart { id, name } => format!("tool_start:{id}:{name}"),
ProviderStreamEvent::ToolCallChunk { id, arguments } => {
format!("tool_chunk:{id}:{arguments}")
}
ProviderStreamEvent::Done {
finish_reason,
usage,
} => format!(
"done:{}:{}",
finish_reason.as_deref().unwrap_or("-"),
usage
.as_ref()
.map(|usage| format!(
"{}/{}/{}",
opt(usage.prompt_tokens),
opt(usage.completion_tokens),
opt(usage.total_tokens)
))
.unwrap_or_else(|| "-".to_string())
),
}
}
fn opt(value: Option<i32>) -> String {
value.map_or_else(|| "-".to_string(), |value| value.to_string())
}
pub fn stream_text(events: &[ProviderStreamEvent]) -> String {
events
.iter()
.filter_map(|event| match event {
ProviderStreamEvent::Text(text) => Some(text.as_str()),
_ => None,
})
.collect()
}
pub struct SplitSseServer {
base_url: String,
handle: tokio::task::JoinHandle<()>,
}
impl SplitSseServer {
pub fn base_url(&self) -> &str {
&self.base_url
}
}
impl Drop for SplitSseServer {
fn drop(&mut self) {
self.handle.abort();
}
}
pub async fn split_sse(parts: Vec<String>) -> SplitSseServer {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("bind loopback listener");
let addr = listener.local_addr().expect("read listener address");
let handle = tokio::spawn(async move {
let Ok((mut socket, _)) = listener.accept().await else {
return;
};
let mut scratch = [0_u8; 8192];
let _ = socket.read(&mut scratch).await;
let head = "HTTP/1.1 200 OK\r\n\
Content-Type: text/event-stream\r\n\
Cache-Control: no-cache\r\n\
Connection: close\r\n\r\n";
if socket.write_all(head.as_bytes()).await.is_err() {
return;
}
if socket.flush().await.is_err() {
return;
}
for part in parts {
if socket.write_all(part.as_bytes()).await.is_err() {
return;
}
if socket.flush().await.is_err() {
return;
}
tokio::time::sleep(Duration::from_millis(40)).await;
}
let _ = socket.shutdown().await;
});
SplitSseServer {
base_url: format!("http://{addr}"),
handle,
}
}
pub struct DisconnectSseServer {
base_url: String,
streaming: Option<tokio::sync::oneshot::Receiver<()>>,
disconnected: Option<tokio::sync::oneshot::Receiver<()>>,
handle: tokio::task::JoinHandle<()>,
}
impl DisconnectSseServer {
pub fn base_url(&self) -> &str {
&self.base_url
}
pub async fn wait_until_streaming(&mut self) {
let receiver = self
.streaming
.take()
.expect("wait_until_streaming should only be called once");
tokio::time::timeout(Duration::from_secs(10), receiver)
.await
.expect("the client never opened the upstream connection")
.expect("the server task ended before it started streaming");
}
pub async fn wait_for_disconnect(&mut self) {
let receiver = self
.disconnected
.take()
.expect("wait_for_disconnect should only be called once");
tokio::time::timeout(Duration::from_secs(10), receiver)
.await
.expect("the client never closed the upstream connection")
.expect("the server task ended without reporting a disconnect");
}
}
impl Drop for DisconnectSseServer {
fn drop(&mut self) {
self.handle.abort();
}
}
pub async fn sse_until_disconnect(parts: Vec<String>) -> DisconnectSseServer {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("bind loopback listener");
let addr = listener.local_addr().expect("read listener address");
let (notify, disconnected) = tokio::sync::oneshot::channel();
let (notify_streaming, streaming) = tokio::sync::oneshot::channel();
let handle = tokio::spawn(async move {
let Ok((mut socket, _)) = listener.accept().await else {
return;
};
let mut scratch = [0_u8; 8192];
let _ = socket.read(&mut scratch).await;
let head = "HTTP/1.1 200 OK\r\n\
Content-Type: text/event-stream\r\n\
Cache-Control: no-cache\r\n\
Connection: close\r\n\r\n";
if socket.write_all(head.as_bytes()).await.is_err() {
return;
}
for part in parts {
if socket.write_all(part.as_bytes()).await.is_err() {
return;
}
if socket.flush().await.is_err() {
return;
}
}
let _ = notify_streaming.send(());
loop {
match socket.read(&mut scratch).await {
Ok(0) | Err(_) => break,
Ok(_) => continue,
}
}
let _ = notify.send(());
});
DisconnectSseServer {
base_url: format!("http://{addr}"),
streaming: Some(streaming),
disconnected: Some(disconnected),
handle,
}
}
pub async fn truncated_chunked_sse(parts: Vec<String>) -> SplitSseServer {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("bind loopback listener");
let addr = listener.local_addr().expect("read listener address");
let handle = tokio::spawn(async move {
let Ok((mut socket, _)) = listener.accept().await else {
return;
};
let mut scratch = [0_u8; 8192];
let _ = socket.read(&mut scratch).await;
let head = "HTTP/1.1 200 OK\r\n\
Content-Type: text/event-stream\r\n\
Cache-Control: no-cache\r\n\
Transfer-Encoding: chunked\r\n\r\n";
if socket.write_all(head.as_bytes()).await.is_err() {
return;
}
for part in parts {
let chunk = format!("{:x}\r\n{part}\r\n", part.len());
if socket.write_all(chunk.as_bytes()).await.is_err() {
return;
}
if socket.flush().await.is_err() {
return;
}
}
drop(socket);
});
SplitSseServer {
base_url: format!("http://{addr}"),
handle,
}
}
pub fn expect_error<T: std::fmt::Debug>(result: Result<T>) -> Error {
match result {
Ok(value) => panic!("expected an error, got: {value:?}"),
Err(error) => error,
}
}
const ENV_CHILD_MARKER: &str = "RAI_SDK_ENV_TEST_CHILD";
pub const SDK_ENV_VARS: &[&str] = &[
"OPENAI_API_KEY",
"OPENAI_BASE_URL",
"ANTHROPIC_API_KEY",
"ANTHROPIC_BASE_URL",
"OPENROUTER_API_KEY",
"OPENROUTER_BASE_URL",
"OPENROUTER_HTTP_REFERER",
"OPENROUTER_APP_URL",
"OPENROUTER_TITLE",
"OPENROUTER_APP_TITLE",
"OPENROUTER_CATEGORIES",
"AI_TIMEOUT_SECONDS",
"AI_MAX_RETRIES",
"AI_RETRY_INITIAL_DELAY_MS",
"AI_RETRY_MAX_DELAY_MS",
"AI_RETRY_BACKOFF_MULTIPLIER",
"AI_RETRY_JITTER",
];
pub fn in_env_child() -> bool {
std::env::var(ENV_CHILD_MARKER).is_ok()
}
pub fn run_in_clean_env(test_name: &str, vars: &[(&str, &str)]) {
let exe = std::env::current_exe().expect("locate the current test binary");
let mut command = std::process::Command::new(exe);
command
.arg("--exact")
.arg(test_name)
.arg("--nocapture")
.arg("--test-threads=1")
.env(ENV_CHILD_MARKER, "1");
for name in SDK_ENV_VARS {
command.env_remove(name);
}
for (name, value) in vars {
command.env(name, value);
}
let output = command.output().expect("spawn the child test process");
let stdout = String::from_utf8_lossy(&output.stdout);
let stderr = String::from_utf8_lossy(&output.stderr);
assert!(
output.status.success(),
"child run of `{test_name}` failed\n--- stdout ---\n{stdout}\n--- stderr ---\n{stderr}"
);
assert!(
stdout.contains("1 passed"),
"child run of `{test_name}` did not execute exactly one test \
(is the name spelled correctly?)\n--- stdout ---\n{stdout}"
);
}