#![allow(dead_code)]
use std::fmt::Write as _;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
pub(crate) struct MockBackend {
pub(crate) url: String,
accepts: Arc<AtomicUsize>,
received: Arc<tokio::sync::Mutex<Vec<Vec<u8>>>>,
}
impl MockBackend {
pub(crate) async fn always(response: &'static [u8]) -> Self {
Self::spawn(move |_| response.to_vec()).await
}
pub(crate) async fn json(status: u16, body: &'static str) -> Self {
Self::spawn(move |_| {
format!(
"HTTP/1.1 {status} OK\r\nContent-Type: application/json\r\n\
Content-Length: {}\r\nConnection: close\r\n\r\n{body}",
body.len()
)
.into_bytes()
})
.await
}
pub(crate) async fn streaming(frames: &'static [&'static str]) -> Self {
Self::spawn_with(move |mut socket| async move {
let mut seen = Vec::new();
let _ = read_head(&mut socket, &mut seen).await;
let _ = socket
.write_all(
b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nConnection: close\r\n\r\n",
)
.await;
let _ = socket.flush().await;
for frame in frames {
tokio::time::sleep(std::time::Duration::from_millis(20)).await;
let _ = socket.write_all(frame.as_bytes()).await;
let _ = socket.flush().await;
}
seen
})
.await
}
pub(crate) async fn endless_stream() -> (Self, Arc<AtomicUsize>) {
let written = Arc::new(AtomicUsize::new(0));
let counter = written.clone();
let backend = Self::spawn_with(move |mut socket| {
let counter = counter.clone();
async move {
let mut seen = Vec::new();
let _ = read_head(&mut socket, &mut seen).await;
if socket
.write_all(b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\r\n")
.await
.is_err()
{
return seen;
}
let _ = socket.flush().await;
for i in 0..2000 {
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
if socket
.write_all(format!("data: {i}\n\n").as_bytes())
.await
.is_err()
|| socket.flush().await.is_err()
{
break;
}
counter.fetch_add(1, Ordering::SeqCst);
}
seen
}
})
.await;
(backend, written)
}
pub(crate) async fn reads_whole_body(response: &'static str) -> Self {
Self::spawn_with(move |mut socket| async move {
let mut seen = Vec::new();
let _ = read_head(&mut socket, &mut seen).await;
let head_end = seen
.windows(4)
.position(|w| w == b"\r\n\r\n")
.map_or(seen.len(), |at| at + 4);
let declared = declared_length(&seen[..head_end]);
let mut buf = [0u8; 4096];
while seen.len() - head_end < declared {
match socket.read(&mut buf).await {
Ok(0) | Err(_) => return seen,
Ok(n) => seen.extend_from_slice(&buf[..n]),
}
}
let _ = socket.write_all(response.as_bytes()).await;
let _ = socket.flush().await;
seen
})
.await
}
async fn spawn(respond: impl Fn(&[u8]) -> Vec<u8> + Send + Sync + 'static) -> Self {
let respond = Arc::new(respond);
Self::spawn_with(move |mut socket| {
let respond = respond.clone();
async move {
let mut seen = Vec::new();
let _ = read_head(&mut socket, &mut seen).await;
let _ = socket.write_all(&respond(&seen)).await;
let _ = socket.flush().await;
seen
}
})
.await
}
async fn spawn_with<F, Fut>(handle: F) -> Self
where
F: Fn(TcpStream) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Vec<u8>> + Send + 'static,
{
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let addr = listener.local_addr().expect("addr");
let accepts = Arc::new(AtomicUsize::new(0));
let received = Arc::new(tokio::sync::Mutex::new(Vec::new()));
let counter = accepts.clone();
let sink = received.clone();
let handle = Arc::new(handle);
tokio::spawn(async move {
while let Ok((socket, _)) = listener.accept().await {
counter.fetch_add(1, Ordering::SeqCst);
let sink = sink.clone();
let handle = handle.clone();
tokio::spawn(async move {
let seen = handle(socket).await;
sink.lock().await.push(seen);
});
}
});
Self {
url: format!("http://{addr}"),
accepts,
received,
}
}
pub(crate) fn accepts(&self) -> usize {
self.accepts.load(Ordering::SeqCst)
}
pub(crate) async fn received(&self) -> String {
let seen = self.received.lock().await;
String::from_utf8_lossy(&seen.concat()).into_owned()
}
}
fn declared_length(head: &[u8]) -> usize {
String::from_utf8_lossy(head)
.to_ascii_lowercase()
.split("content-length:")
.nth(1)
.and_then(|rest| rest.split("\r\n").next())
.and_then(|value| value.trim().parse().ok())
.unwrap_or(0)
}
async fn read_head(socket: &mut TcpStream, into: &mut Vec<u8>) -> std::io::Result<()> {
let mut buf = [0u8; 4096];
loop {
let n = socket.read(&mut buf).await?;
if n == 0 {
return Ok(());
}
into.extend_from_slice(&buf[..n]);
if into.windows(4).any(|w| w == b"\r\n\r\n") {
return Ok(());
}
}
}
pub(crate) async fn request(
base_url: &str,
path: &str,
auth: Option<&str>,
) -> std::io::Result<String> {
let authority = base_url
.trim_start_matches("http://")
.split('/')
.next()
.expect("authority");
let mut socket = TcpStream::connect(authority).await?;
let mut req = format!("GET {path} HTTP/1.1\r\nHost: {authority}\r\n");
if let Some(value) = auth {
let _ = write!(req, "Authorization: {value}\r\n");
}
req.push_str("\r\n");
socket.write_all(req.as_bytes()).await?;
socket.flush().await?;
let mut seen = Vec::new();
socket.read_to_end(&mut seen).await?;
Ok(String::from_utf8_lossy(&seen).into_owned())
}
pub(crate) struct Scratch(std::path::PathBuf);
impl Scratch {
pub(crate) fn new(name: &str) -> Self {
let dir = std::env::temp_dir().join(format!("modelpipe-it-{}-{name}", std::process::id()));
std::fs::create_dir_all(&dir).expect("a scratch directory");
Self(dir)
}
pub(crate) fn join(&self, file: &str) -> std::path::PathBuf {
self.0.join(file)
}
}
impl Drop for Scratch {
fn drop(&mut self) {
let _ = std::fs::remove_dir_all(&self.0);
}
}
pub(crate) async fn within<F: Future>(why: &str, future: F) -> F::Output {
tokio::time::timeout(std::time::Duration::from_secs(20), future)
.await
.unwrap_or_else(|_| panic!("{why}"))
}