use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader};
use tokio::net::{TcpListener, TcpStream};
pub struct CaptureUpstream {
pub port: u16,
pub received_body: Arc<Mutex<Option<Vec<u8>>>>,
pub received_headers: Arc<Mutex<Option<String>>>,
pub received_request_line: Arc<Mutex<Option<String>>>,
}
impl CaptureUpstream {
pub fn header(&self, name: &str) -> Option<String> {
let headers = self.received_headers.lock().unwrap().clone()?;
headers.lines().find_map(|l| {
l.split_once(':')
.filter(|(k, _)| k.eq_ignore_ascii_case(name))
.map(|(_, v)| v.trim().to_string())
})
}
}
async fn read_request(stream: &mut TcpStream) -> std::io::Result<(String, Vec<u8>)> {
let mut reader = BufReader::new(stream);
let mut header_block = String::new();
loop {
let mut line = String::new();
let n = reader.read_line(&mut line).await?;
if n == 0 {
break;
}
let blank = line == "\r\n" || line == "\n";
header_block.push_str(&line);
if blank {
break;
}
}
let body = if let Some(cl) = parse_content_length(&header_block) {
let mut buf = vec![0u8; cl];
let _ = reader.read_exact(&mut buf).await;
buf
} else if is_chunked(&header_block) {
read_chunked(&mut reader).await?
} else {
Vec::new()
};
Ok((header_block, body))
}
async fn read_chunked(reader: &mut BufReader<&mut TcpStream>) -> std::io::Result<Vec<u8>> {
let mut body = Vec::new();
loop {
let mut size_line = String::new();
if reader.read_line(&mut size_line).await? == 0 {
break;
}
let size = usize::from_str_radix(size_line.trim(), 16).unwrap_or(0);
if size == 0 {
let mut trailer = String::new();
let _ = reader.read_line(&mut trailer).await;
break;
}
let mut chunk = vec![0u8; size];
reader.read_exact(&mut chunk).await?;
body.extend_from_slice(&chunk);
let mut crlf = [0u8; 2];
let _ = reader.read_exact(&mut crlf).await;
}
Ok(body)
}
fn parse_content_length(headers: &str) -> Option<usize> {
headers
.lines()
.find_map(|l| {
l.split_once(':')
.filter(|(k, _)| k.eq_ignore_ascii_case("content-length"))
})
.and_then(|(_, v)| v.trim().parse().ok())
}
fn is_chunked(headers: &str) -> bool {
headers.lines().any(|l| {
l.split_once(':').is_some_and(|(k, v)| {
k.eq_ignore_ascii_case("transfer-encoding")
&& v.to_ascii_lowercase().contains("chunked")
})
})
}
pub async fn spawn_capture_200() -> CaptureUpstream {
let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
let port = listener.local_addr().unwrap().port();
let received_body = Arc::new(Mutex::new(None));
let received_headers = Arc::new(Mutex::new(None));
let received_request_line = Arc::new(Mutex::new(None));
let body_sink = received_body.clone();
let hdr_sink = received_headers.clone();
let line_sink = received_request_line.clone();
tokio::spawn(async move {
if let Ok((mut stream, _)) = listener.accept().await {
if let Ok((headers, body)) = read_request(&mut stream).await {
*line_sink.lock().unwrap() = headers.lines().next().map(str::to_string);
*hdr_sink.lock().unwrap() = Some(headers);
*body_sink.lock().unwrap() = Some(body);
}
let _ = stream
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok")
.await;
let _ = stream.flush().await;
}
});
CaptureUpstream {
port,
received_body,
received_headers,
received_request_line,
}
}
pub struct TrickleUpstream {
pub port: u16,
pub final_written_at: Arc<Mutex<Option<Instant>>>,
}
pub async fn spawn_trickle_sse(chunk_count: usize, gap: Duration) -> TrickleUpstream {
let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
let port = listener.local_addr().unwrap().port();
let final_written_at = Arc::new(Mutex::new(None));
let stamp = final_written_at.clone();
tokio::spawn(async move {
if let Ok((mut stream, _)) = listener.accept().await {
let _ = read_request(&mut stream).await;
let head = "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\
Transfer-Encoding: chunked\r\n\r\n";
let _ = stream.write_all(head.as_bytes()).await;
let _ = stream.flush().await;
for i in 0..chunk_count {
let payload = format!("data: chunk-{i}\n\n");
let framed = format!("{:x}\r\n{}\r\n", payload.len(), payload);
let _ = stream.write_all(framed.as_bytes()).await;
let _ = stream.flush().await;
if i + 1 < chunk_count {
tokio::time::sleep(gap).await;
}
}
let _ = stream.write_all(b"0\r\n\r\n").await;
let _ = stream.flush().await;
*stamp.lock().unwrap() = Some(Instant::now());
}
});
TrickleUpstream {
port,
final_written_at,
}
}
#[allow(clippy::too_many_arguments)]
pub async fn spawn_capture_usage_sse(
input_tokens: u64,
cache_read: u64,
cache_write: u64,
eph_5m: u64,
eph_1h: u64,
output_tokens: u64,
) -> CaptureUpstream {
let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
let port = listener.local_addr().unwrap().port();
let received_body = Arc::new(Mutex::new(None));
let received_headers = Arc::new(Mutex::new(None));
let received_request_line = Arc::new(Mutex::new(None));
let body_sink = received_body.clone();
let hdr_sink = received_headers.clone();
let line_sink = received_request_line.clone();
tokio::spawn(async move {
if let Ok((mut stream, _)) = listener.accept().await {
if let Ok((headers, body)) = read_request(&mut stream).await {
*line_sink.lock().unwrap() = headers.lines().next().map(str::to_string);
*hdr_sink.lock().unwrap() = Some(headers);
*body_sink.lock().unwrap() = Some(body);
}
let sse = format!(
"event: message_start\n\
data: {{\"type\":\"message_start\",\"message\":{{\"usage\":{{\"input_tokens\":{input_tokens},\"cache_read_input_tokens\":{cache_read},\"cache_creation_input_tokens\":{cache_write},\"cache_creation\":{{\"ephemeral_5m_input_tokens\":{eph_5m},\"ephemeral_1h_input_tokens\":{eph_1h}}},\"output_tokens\":1}}}}}}\n\n\
event: content_block_delta\n\
data: {{\"type\":\"content_block_delta\",\"delta\":{{\"text\":\"hello\"}}}}\n\n\
event: message_delta\n\
data: {{\"type\":\"message_delta\",\"usage\":{{\"output_tokens\":{output_tokens}}}}}\n\n"
);
let head = format!(
"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\
Content-Length: {}\r\nConnection: close\r\n\r\n",
sse.len()
);
let _ = stream.write_all(head.as_bytes()).await;
let _ = stream.write_all(sse.as_bytes()).await;
let _ = stream.flush().await;
}
});
CaptureUpstream {
port,
received_body,
received_headers,
received_request_line,
}
}
pub async fn spawn_message_start_then_hangup() -> CaptureUpstream {
let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
let port = listener.local_addr().unwrap().port();
let received_body = Arc::new(Mutex::new(None));
let received_headers = Arc::new(Mutex::new(None));
let received_request_line = Arc::new(Mutex::new(None));
let body_sink = received_body.clone();
let hdr_sink = received_headers.clone();
let line_sink = received_request_line.clone();
tokio::spawn(async move {
if let Ok((mut stream, _)) = listener.accept().await {
if let Ok((headers, body)) = read_request(&mut stream).await {
*line_sink.lock().unwrap() = headers.lines().next().map(str::to_string);
*hdr_sink.lock().unwrap() = Some(headers);
*body_sink.lock().unwrap() = Some(body);
}
let start = "event: message_start\n\
data: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":40,\"cache_read_input_tokens\":0,\"cache_creation_input_tokens\":0,\"output_tokens\":1}}}\n\n";
let head =
"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nConnection: close\r\n\r\n";
let _ = stream.write_all(head.as_bytes()).await;
let _ = stream.write_all(start.as_bytes()).await;
let _ = stream.flush().await;
drop(stream);
}
});
CaptureUpstream {
port,
received_body,
received_headers,
received_request_line,
}
}
pub async fn closed_port() -> u16 {
let l = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
let p = l.local_addr().unwrap().port();
drop(l);
p
}
pub async fn spawn_hang_after_accept() -> u16 {
let listener = TcpListener::bind(("127.0.0.1", 0)).await.unwrap();
let port = listener.local_addr().unwrap().port();
tokio::spawn(async move {
if let Ok((mut stream, _)) = listener.accept().await {
let _ = read_request(&mut stream).await;
tokio::time::sleep(Duration::from_secs(30)).await;
drop(stream);
}
});
port
}