use crate::backend::{
AgentBackend, AgentEvent, AgentSession, PromptMode, SessionExit, SessionSpec,
};
use crate::error::{EngineError, Result};
use crate::stream_bounds::TailWindow;
use crate::types::TokenUsage;
use serde_json::{json, Value};
use std::collections::VecDeque;
use std::time::Duration;
const BODY_TAIL_CHARS: usize = 500;
const REQUEST_TIMEOUT: Duration = Duration::from_secs(600);
const RESPONSE_BODY_CAP: usize = 8 * 1024 * 1024;
const CHARS_PER_TOKEN: usize = 4;
fn last_chars(text: &str, max: usize) -> String {
let chars: Vec<char> = text.chars().collect();
let start = chars.len().saturating_sub(max);
chars[start..].iter().collect()
}
fn prompt_text(spec: &SessionSpec) -> &str {
match &spec.prompt {
PromptMode::SingleShot(text) => text.as_str(),
PromptMode::Streaming(text) => text.as_str(),
}
}
fn build_messages(spec: &SessionSpec) -> Vec<Value> {
let mut messages = Vec::new();
if let Some(system) = &spec.append_system_prompt {
if !system.is_empty() {
messages.push(json!({"role": "system", "content": system}));
}
}
messages.push(json!({"role": "user", "content": prompt_text(spec)}));
messages
}
fn estimate_tokens(messages: &[Value]) -> u32 {
let total_chars: usize = messages
.iter()
.filter_map(|m| m.get("content").and_then(Value::as_str))
.map(|s| s.chars().count())
.sum();
total_chars.div_ceil(CHARS_PER_TOKEN) as u32
}
fn extract_completion(body: &Value) -> std::result::Result<(String, TokenUsage), String> {
let content = body
.get("choices")
.and_then(Value::as_array)
.and_then(|choices| choices.first())
.and_then(|choice| choice.get("message"))
.and_then(|message| message.get("content"))
.and_then(Value::as_str)
.ok_or_else(|| {
format!(
"response missing choices[0].message.content; body tail: {}",
last_chars(&body.to_string(), BODY_TAIL_CHARS)
)
})?
.to_string();
let usage = body.get("usage");
let input = usage
.and_then(|u| u.get("prompt_tokens"))
.and_then(Value::as_u64)
.unwrap_or(0);
let output = usage
.and_then(|u| u.get("completion_tokens"))
.and_then(Value::as_u64)
.unwrap_or(0);
Ok((
content,
TokenUsage {
input,
output,
cache_read: 0,
cache_write: 0,
},
))
}
async fn read_body_tail(mut response: reqwest::Response, cap: usize) -> String {
let mut window = TailWindow::new(cap);
while let Ok(Some(chunk)) = response.chunk().await {
window.push(&chunk);
}
window.render()
}
#[derive(Debug, Clone)]
pub struct LocalBackend {
base_url: String,
temperature: Option<f64>,
context_budget: u32,
request_timeout: Duration,
body_cap: usize,
client: reqwest::Client,
}
impl LocalBackend {
pub fn new(base_url: String, temperature: Option<f64>, context_budget: u32) -> Self {
LocalBackend {
base_url,
temperature,
context_budget,
request_timeout: REQUEST_TIMEOUT,
body_cap: RESPONSE_BODY_CAP,
client: reqwest::Client::new(),
}
}
}
#[async_trait::async_trait]
impl AgentBackend for LocalBackend {
async fn start(&self, spec: SessionSpec) -> Result<Box<dyn AgentSession>> {
if spec.resume.is_some() {
return Err(EngineError::Backend(
"local backend is single-shot only; resume is unsupported".to_string(),
));
}
let session_id = spec.session_id.clone();
let model = spec.model.clone();
let messages = build_messages(&spec);
let estimated_tokens = estimate_tokens(&messages);
if estimated_tokens > self.context_budget {
return Ok(Box::new(LocalSession::context_budget_exceeded(
session_id,
model,
estimated_tokens,
self.context_budget,
)));
}
let mut request_body = json!({
"model": model,
"messages": messages,
});
if let Some(temperature) = self.temperature {
request_body["temperature"] = json!(temperature);
}
let url = format!(
"{}/v1/chat/completions",
self.base_url.trim_end_matches('/')
);
let outcome = match self
.client
.post(&url)
.json(&request_body)
.timeout(self.request_timeout)
.send()
.await
{
Ok(response) => {
let status = response.status();
let body_text = read_body_tail(response, self.body_cap).await;
if status.is_success() {
match serde_json::from_str::<Value>(&body_text) {
Ok(parsed) => Ok(parsed),
Err(e) => Err(format!(
"failed to parse local backend response as JSON: {e}; body tail: {}",
last_chars(&body_text, BODY_TAIL_CHARS)
)),
}
} else {
Err(format!(
"local backend request failed with HTTP {status}; body tail: {}",
last_chars(&body_text, BODY_TAIL_CHARS)
))
}
}
Err(e) if e.is_timeout() => Err(format!(
"local backend request timed out after {:?}",
self.request_timeout
)),
Err(e) => Err(format!("local backend request failed: {e}")),
};
Ok(Box::new(LocalSession::from_response(
session_id, model, outcome,
)))
}
}
pub struct LocalSession {
session_id: String,
queue: VecDeque<AgentEvent>,
pending_exit: Option<SessionExit>,
exit: Option<SessionExit>,
}
impl LocalSession {
fn from_response(
session_id: String,
model: String,
outcome: std::result::Result<Value, String>,
) -> Self {
match outcome {
Ok(body) => match extract_completion(&body) {
Ok((content, usage)) => {
let mut queue = VecDeque::new();
queue.push_back(AgentEvent::Init {
session_id: session_id.clone(),
model,
raw: body.clone(),
});
queue.push_back(AgentEvent::Text {
text: content.clone(),
raw: body.clone(),
});
queue.push_back(AgentEvent::Result {
text: content,
is_error: false,
usage,
cost_usd: Some(0.0),
num_turns: Some(1),
raw: body,
});
LocalSession {
session_id,
queue,
pending_exit: Some(SessionExit::Completed),
exit: None,
}
}
Err(message) => LocalSession {
session_id,
queue: VecDeque::new(),
pending_exit: Some(SessionExit::Failed(message)),
exit: None,
},
},
Err(message) => LocalSession {
session_id,
queue: VecDeque::new(),
pending_exit: Some(SessionExit::Failed(message)),
exit: None,
},
}
}
fn context_budget_exceeded(
session_id: String,
model: String,
estimated_tokens: u32,
context_budget: u32,
) -> Self {
let mut queue = VecDeque::new();
queue.push_back(AgentEvent::Init {
session_id: session_id.clone(),
model,
raw: json!({}),
});
LocalSession {
session_id,
queue,
pending_exit: Some(SessionExit::Failed(format!(
"prompt estimated at {estimated_tokens} tokens exceeds context budget of \
{context_budget} tokens; no request was sent"
))),
exit: None,
}
}
}
#[async_trait::async_trait]
impl AgentSession for LocalSession {
fn session_id(&self) -> String {
self.session_id.clone()
}
async fn next_event(&mut self) -> Result<Option<AgentEvent>> {
if let Some(event) = self.queue.pop_front() {
return Ok(Some(event));
}
if self.exit.is_none() {
self.exit = self.pending_exit.take();
}
Ok(None)
}
async fn send_user_message(&mut self, _text: &str) -> Result<()> {
Err(EngineError::Backend(
"local backend is single-shot only; send_user_message is unsupported".to_string(),
))
}
async fn abort(&mut self) -> Result<()> {
self.queue.clear();
if self.exit.is_none() {
self.exit = Some(SessionExit::Aborted);
}
self.pending_exit = None;
Ok(())
}
fn exit_status(&self) -> Option<SessionExit> {
self.exit.clone()
}
}
#[cfg(test)]
pub(crate) mod tests {
use super::*;
use std::collections::HashMap;
use std::net::SocketAddr;
use std::path::PathBuf;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
const TEST_MODEL: &str = "local-test-model";
fn find_subslice(haystack: &[u8], needle: &[u8]) -> Option<usize> {
haystack
.windows(needle.len())
.position(|window| window == needle)
}
async fn read_http_request(socket: &mut TcpStream) -> Vec<u8> {
let mut buf = Vec::new();
let mut chunk = [0u8; 4096];
loop {
let header_end = find_subslice(&buf, b"\r\n\r\n");
if let Some(header_end) = header_end {
let headers = String::from_utf8_lossy(&buf[..header_end]).to_string();
let content_length: usize = headers
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
if name.eq_ignore_ascii_case("content-length") {
value.trim().parse().ok()
} else {
None
}
})
.unwrap_or(0);
let body_start = header_end + 4;
if buf.len() >= body_start + content_length {
break;
}
}
match socket.read(&mut chunk).await {
Ok(0) => break,
Ok(n) => buf.extend_from_slice(&chunk[..n]),
Err(_) => break,
}
}
buf
}
pub(crate) async fn spawn_stub(
status_line: &'static str,
body: String,
) -> (String, Arc<AtomicUsize>, Arc<Mutex<Vec<u8>>>) {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind stub");
let addr: SocketAddr = listener.local_addr().expect("stub addr");
let count = Arc::new(AtomicUsize::new(0));
let count_for_task = Arc::clone(&count);
let received = Arc::new(Mutex::new(Vec::new()));
let received_for_task = Arc::clone(&received);
tokio::spawn(async move {
loop {
let (mut socket, _) = match listener.accept().await {
Ok(v) => v,
Err(_) => break,
};
count_for_task.fetch_add(1, Ordering::SeqCst);
let request_bytes = read_http_request(&mut socket).await;
*received_for_task.lock().expect("stub request lock") = request_bytes;
let response = format!(
"{status_line}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
body.len(),
body
);
let _ = socket.write_all(response.as_bytes()).await;
let _ = socket.shutdown().await;
}
});
(format!("http://{addr}"), count, received)
}
fn base_spec(session_id: &str, prompt: &str, context_budget_prompt: bool) -> SessionSpec {
let _ = context_budget_prompt;
SessionSpec {
cwd: PathBuf::from("."),
prompt: PromptMode::SingleShot(prompt.to_string()),
append_system_prompt: Some("be terse".to_string()),
model: TEST_MODEL.to_string(),
effort: String::new(),
session_id: session_id.to_string(),
resume: None,
permission_mode: None,
allowed_tools: vec![],
disallowed_tools: vec![],
tools: vec![],
writable: true,
settings_json: None,
json_schema: None,
max_budget_usd: None,
max_turns: None,
env: HashMap::new(),
sandbox: None,
hook_status: None,
}
}
async fn drain(session: &mut dyn AgentSession) -> Vec<AgentEvent> {
let mut events = Vec::new();
while let Some(event) = session.next_event().await.expect("next_event") {
events.push(event);
}
events
}
#[tokio::test]
async fn local_http_roundtrip_yields_init_text_result_with_usage_and_zero_cost() {
let stub_body = json!({
"id": "chatcmpl-1",
"choices": [{"message": {"role": "assistant", "content": "hello from stub"}}],
"usage": {"prompt_tokens": 12, "completion_tokens": 34, "total_tokens": 46}
})
.to_string();
let (base_url, requests, received) = spawn_stub("HTTP/1.1 200 OK", stub_body).await;
let backend = LocalBackend::new(base_url, Some(0.2), 100_000);
let spec = base_spec("sess-1", "do the thing", false);
let mut session = backend.start(spec).await.expect("start");
let events = drain(session.as_mut()).await;
assert_eq!(requests.load(Ordering::SeqCst), 1);
let raw_request = received.lock().expect("stub request lock").clone();
let request_text = String::from_utf8_lossy(&raw_request).to_string();
let request_line = request_text.lines().next().expect("request line");
assert!(
request_line.starts_with("POST "),
"expected a POST request, got: {request_line}"
);
assert!(
request_line
.split_whitespace()
.nth(1)
.expect("request target")
.ends_with("/v1/chat/completions"),
"expected the request target to end with /v1/chat/completions, got: {request_line}"
);
let header_end = find_subslice(&raw_request, b"\r\n\r\n").expect("request headers");
let request_body: Value = serde_json::from_slice(&raw_request[header_end + 4..])
.expect("request body should be JSON");
assert_eq!(request_body["model"], json!(TEST_MODEL));
assert_eq!(request_body["temperature"], json!(0.2));
let messages = request_body["messages"].as_array().expect("messages array");
assert!(
messages
.iter()
.any(|m| m["role"] == "system" && m["content"] == "be terse"),
"expected the system message from append_system_prompt, got: {messages:?}"
);
assert_eq!(
messages.last().expect("at least one message"),
&json!({"role": "user", "content": "do the thing"})
);
assert!(
matches!(&events[0], AgentEvent::Init { session_id, model, .. }
if session_id == "sess-1" && model == TEST_MODEL)
);
assert!(matches!(&events[1], AgentEvent::Text { text, .. } if text == "hello from stub"));
match &events[2] {
AgentEvent::Result {
text,
is_error,
usage,
cost_usd,
num_turns,
..
} => {
assert_eq!(text, "hello from stub");
assert!(!is_error);
assert_eq!(usage.input, 12);
assert_eq!(usage.output, 34);
assert_eq!(*cost_usd, Some(0.0));
assert_eq!(*num_turns, Some(1));
}
other => panic!("expected terminal Result, got {other:?}"),
}
assert_eq!(events.len(), 3);
assert_eq!(session.exit_status(), Some(SessionExit::Completed));
}
#[tokio::test]
async fn local_http_rejects_resumed_spec() {
let backend = LocalBackend::new("http://127.0.0.1:1".to_string(), None, 100_000);
let mut spec = base_spec("sess-1", "do the thing", false);
spec.resume = Some("sess-0".to_string());
let result = backend.start(spec).await;
assert!(result.is_err(), "expected resume to be rejected");
}
#[tokio::test]
async fn local_http_send_user_message_errors() {
let stub_body = json!({
"choices": [{"message": {"role": "assistant", "content": "hi"}}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1}
})
.to_string();
let (base_url, _requests, _received) = spawn_stub("HTTP/1.1 200 OK", stub_body).await;
let backend = LocalBackend::new(base_url, None, 100_000);
let spec = base_spec("sess-1", "do the thing", false);
let mut session = backend.start(spec).await.expect("start");
let result = session.send_user_message("nope").await;
assert!(result.is_err(), "expected send_user_message to be rejected");
}
#[tokio::test]
async fn local_http_500_fails_cleanly() {
let (base_url, requests, _received) =
spawn_stub("HTTP/1.1 500 Internal Server Error", "boom".to_string()).await;
let backend = LocalBackend::new(base_url, None, 100_000);
let spec = base_spec("sess-1", "do the thing", false);
let mut session = backend.start(spec).await.expect("start");
let events = drain(session.as_mut()).await;
assert!(events.is_empty(), "expected no events on a failed response");
assert_eq!(requests.load(Ordering::SeqCst), 1);
match session.exit_status() {
Some(SessionExit::Failed(message)) => {
assert!(
message.contains("500"),
"expected the failure message to include the HTTP status, got: {message}"
);
}
other => panic!("expected SessionExit::Failed, got {other:?}"),
}
}
#[tokio::test]
async fn local_http_context_budget_exceeds_fails_cleanly() {
let (base_url, requests, _received) = spawn_stub("HTTP/1.1 200 OK", "{}".to_string()).await;
let backend = LocalBackend::new(base_url, None, 1);
let spec = base_spec(
"sess-1",
"this prompt is far too long for a one-token budget",
false,
);
let mut session = backend.start(spec).await.expect("start");
let events = drain(session.as_mut()).await;
assert_eq!(
requests.load(Ordering::SeqCst),
0,
"must not send an HTTP request when the context budget is exceeded"
);
assert_eq!(events.len(), 1, "expected only a synthesized Init event");
assert!(matches!(&events[0], AgentEvent::Init { .. }));
match session.exit_status() {
Some(SessionExit::Failed(message)) => {
assert!(
message.contains("context budget"),
"expected the failure message to name the context budget, got: {message}"
);
}
other => panic!("expected SessionExit::Failed, got {other:?}"),
}
}
async fn spawn_hung_stub() -> String {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("bind hung stub");
let addr: SocketAddr = listener.local_addr().expect("hung stub addr");
tokio::spawn(async move {
let mut held = Vec::new();
while let Ok((socket, _)) = listener.accept().await {
held.push(socket);
}
});
format!("http://{addr}")
}
#[tokio::test]
async fn local_http_hung_endpoint_times_out_instead_of_stalling() {
let base_url = spawn_hung_stub().await;
let mut backend = LocalBackend::new(base_url, None, 100_000);
backend.request_timeout = Duration::from_millis(200);
let start = std::time::Instant::now();
let spec = base_spec("sess-hung", "do the thing", false);
let mut session = backend.start(spec).await.expect("start");
let events = drain(session.as_mut()).await;
assert!(events.is_empty(), "a timed-out request yields no events");
assert!(
start.elapsed() < Duration::from_secs(10),
"the request returned near the 200ms timeout, not after a stall"
);
match session.exit_status() {
Some(SessionExit::Failed(message)) => {
assert!(
message.contains("timed out"),
"expected the failure to name the timeout, got: {message}"
);
}
other => panic!("expected SessionExit::Failed, got {other:?}"),
}
}
#[tokio::test]
async fn local_http_error_body_over_the_cap_is_tailed_with_a_marker() {
let body = format!("{}{}", "x".repeat(4096), "BODY-END");
let (base_url, _requests, _received) =
spawn_stub("HTTP/1.1 500 Internal Server Error", body).await;
let mut backend = LocalBackend::new(base_url, None, 100_000);
backend.body_cap = 128;
let spec = base_spec("sess-cap", "do the thing", false);
let mut session = backend.start(spec).await.expect("start");
let events = drain(session.as_mut()).await;
assert!(events.is_empty(), "expected no events on a failed response");
match session.exit_status() {
Some(SessionExit::Failed(message)) => {
assert!(
message.contains(crate::stream_bounds::TRUNCATION_MARKER),
"expected the truncation marker, got: {message}"
);
assert!(
message.contains("BODY-END"),
"expected the END of the body to be kept, got: {message}"
);
assert!(
message.len() < 1024,
"the surfaced body stayed bounded, got {} bytes",
message.len()
);
}
other => panic!("expected SessionExit::Failed, got {other:?}"),
}
}
#[tokio::test]
async fn local_http_success_body_over_the_cap_fails_honestly_with_a_marker() {
let body = format!(
"{}{}",
json!({"choices": [{"message": {"role": "assistant", "content": "hi"}}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1}}),
" ".repeat(4096)
);
let (base_url, _requests, _received) = spawn_stub("HTTP/1.1 200 OK", body).await;
let mut backend = LocalBackend::new(base_url, None, 100_000);
backend.body_cap = 128;
let spec = base_spec("sess-cap200", "do the thing", false);
let mut session = backend.start(spec).await.expect("start");
let events = drain(session.as_mut()).await;
assert!(
events.is_empty(),
"an unparseable over-cap body yields no events"
);
match session.exit_status() {
Some(SessionExit::Failed(message)) => {
assert!(
message.contains("failed to parse"),
"expected an honest parse failure, got: {message}"
);
assert!(
message.contains(crate::stream_bounds::TRUNCATION_MARKER),
"expected the truncation marker, got: {message}"
);
}
other => panic!("expected SessionExit::Failed, got {other:?}"),
}
}
}