use std::convert::Infallible;
use std::sync::Arc;
use axum::Router;
use axum::extract::State;
use axum::response::sse::{Event, KeepAlive, Sse};
use axum::routing::post;
use futures::{Stream, StreamExt};
use rai_sdk::client::ModelReady;
use rai_sdk::wire::{StreamAccumulator, WireStreamEvent};
use rai_sdk::{Client, ClientBuilder, Model};
async fn generate(
State(client): State<Arc<Client<ModelReady>>>,
prompt: String,
) -> Sse<impl Stream<Item = Result<Event, Infallible>>> {
let events = match client
.request()
.no_tools()
.prompt(prompt)
.stream_wire_events()
.await
{
Ok(events) => events.left_stream(),
Err(error) => {
futures::stream::once(async move { WireStreamEvent::error(&error) }).right_stream()
}
};
let events = events.inspect(|event| {
if let WireStreamEvent::Usage { usage } = event {
eprintln!("[server] usage: {usage:?}");
}
});
let sse = events.map(|event| {
Ok(Event::default()
.event(event.tag())
.json_data(&event)
.expect("a wire event always serializes"))
});
Sse::new(sse).keep_alive(KeepAlive::default())
}
async fn consume(url: &str, prompt: &str) -> Result<(), Box<dyn std::error::Error>> {
let response = reqwest::Client::new()
.post(url)
.body(prompt.to_string())
.send()
.await?
.error_for_status()?;
let mut accumulator = StreamAccumulator::new();
let mut body = response.bytes_stream();
let mut buffer = String::new();
'outer: while let Some(chunk) = body.next().await {
buffer.push_str(&String::from_utf8_lossy(&chunk?));
while let Some(end) = buffer.find("\n\n") {
let block: String = buffer.drain(..end + 2).collect();
for line in block.lines() {
let Some(payload) = line.strip_prefix("data:") else {
continue;
};
let Ok(event) = serde_json::from_str::<WireStreamEvent>(payload.trim()) else {
eprintln!("[client] skipping unrecognized event: {}", payload.trim());
continue;
};
if let WireStreamEvent::TextDelta { text } = &event {
print!("{text}");
}
if let Err(error) = accumulator.push(event) {
eprintln!("\n[client] provider error ({}): {error}", error.kind);
break 'outer;
}
}
}
}
println!();
let response = accumulator.finish()?;
println!("[client] model: {}", response.model);
println!("[client] provider: {}", response.provider);
println!(
"[client] finish reason: {}",
response.finish_reason.as_deref().unwrap_or("-")
);
if let Some(usage) = &response.usage {
println!("[client] usage: {usage:?}");
}
println!("[client] reassembled {} characters", response.text().len());
Ok(())
}
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
let client = Arc::new(
ClientBuilder::new()
.from_env()
.model(Model::gpt4o_mini())
.build()?,
);
let app = Router::new()
.route("/generate", post(generate))
.with_state(client);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?;
let url = format!("http://{}/generate", listener.local_addr()?);
println!("[server] listening on {url}");
let server = tokio::spawn(async move { axum::serve(listener, app).await });
consume(&url, "Explain server-sent events in three sentences.").await?;
server.abort();
Ok(())
}