#![allow(dead_code)]
use std::collections::VecDeque;
use std::io::{Read, Write};
use std::net::{TcpListener, TcpStream};
use std::sync::{Arc, Mutex};
use std::thread;
pub enum Chat {
Sse(Vec<String>),
Drop,
Status(u16),
StatusBody(u16, String),
Hang(Vec<String>),
}
pub struct MockServer {
port: u16,
stop: Arc<Mutex<bool>>,
chat_bodies: Arc<Mutex<Vec<String>>>,
}
impl MockServer {
pub fn start(chats: Vec<Chat>) -> Self {
let listener = TcpListener::bind("127.0.0.1:0").expect("bind mock");
let port = listener.local_addr().unwrap().port();
let queue: Arc<Mutex<VecDeque<Chat>>> = Arc::new(Mutex::new(chats.into_iter().collect()));
let stop = Arc::new(Mutex::new(false));
let stop_thread = Arc::clone(&stop);
let chat_bodies: Arc<Mutex<Vec<String>>> = Arc::new(Mutex::new(Vec::new()));
let bodies_thread = Arc::clone(&chat_bodies);
thread::spawn(move || {
for conn in listener.incoming() {
if *stop_thread.lock().unwrap() {
break;
}
let Ok(stream) = conn else { break };
let queue = Arc::clone(&queue);
let bodies = Arc::clone(&bodies_thread);
thread::spawn(move || handle(stream, &queue, &bodies));
}
});
MockServer {
port,
stop,
chat_bodies,
}
}
pub fn base_url(&self) -> String {
format!("http://127.0.0.1:{}/v1", self.port)
}
pub fn chat_bodies(&self) -> Vec<serde_json::Value> {
self.chat_bodies
.lock()
.unwrap()
.iter()
.map(|body| serde_json::from_str(body).expect("chat request body is JSON"))
.collect()
}
}
impl Drop for MockServer {
fn drop(&mut self) {
*self.stop.lock().unwrap() = true;
}
}
fn handle(mut stream: TcpStream, queue: &Mutex<VecDeque<Chat>>, chat_bodies: &Mutex<Vec<String>>) {
let mut buf = Vec::new();
let mut tmp = [0u8; 4096];
let headers_end = loop {
match stream.read(&mut tmp) {
Ok(0) | Err(_) => return,
Ok(n) => {
buf.extend_from_slice(&tmp[..n]);
if let Some(p) = buf.windows(4).position(|w| w == b"\r\n\r\n") {
break p + 4;
}
}
}
};
let head = String::from_utf8_lossy(&buf[..headers_end]).to_string();
let request_line = head.lines().next().unwrap_or("");
let path = request_line.split_whitespace().nth(1).unwrap_or("");
let content_len: usize = head
.lines()
.find_map(|l| {
l.to_ascii_lowercase()
.strip_prefix("content-length:")
.map(|v| v.trim().parse::<usize>().unwrap_or(0))
})
.unwrap_or(0);
let have = buf.len().saturating_sub(headers_end);
let mut remaining = content_len.saturating_sub(have);
while remaining > 0 {
match stream.read(&mut tmp) {
Ok(0) | Err(_) => break,
Ok(n) => {
buf.extend_from_slice(&tmp[..n]);
remaining = remaining.saturating_sub(n);
}
}
}
if path.contains("/chat/completions") {
let end = (headers_end + content_len).min(buf.len());
chat_bodies
.lock()
.unwrap()
.push(String::from_utf8_lossy(&buf[headers_end..end]).into_owned());
let next = queue.lock().unwrap().pop_front();
match next {
Some(Chat::Sse(lines)) => write_sse(&mut stream, &lines),
Some(Chat::Drop) => { }
Some(Chat::Status(code)) => write_status(&mut stream, code),
Some(Chat::StatusBody(code, body)) => {
let _ = write!(
stream,
"HTTP/1.1 {code} Error\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
body.len()
);
}
Some(Chat::Hang(lines)) => {
let mut body = String::from(
"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nConnection: close\r\n\r\n",
);
for line in &lines {
body.push_str("data: ");
body.push_str(line);
body.push_str("\n\n");
}
let _ = stream.write_all(body.as_bytes());
let _ = stream.flush();
thread::sleep(std::time::Duration::from_secs(30));
}
None => write_sse(
&mut stream,
&[text_chunk("x", ""), stop_chunk("x"), "[DONE]".to_string()],
),
}
} else if path.contains("/models") {
let body = r#"{"object":"list","data":[{"id":"mock-model","object":"model"},{"id":"other","object":"model"}]}"#;
let _ = write!(
stream,
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
body.len(),
body
);
} else {
let _ = write!(
stream,
"HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
);
}
let _ = stream.flush();
}
fn write_sse(stream: &mut TcpStream, lines: &[String]) {
let mut body = String::new();
for line in lines {
body.push_str("data: ");
body.push_str(line);
body.push_str("\n\n");
}
let _ = write!(
stream,
"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nConnection: close\r\n\r\n{body}"
);
}
fn write_status(stream: &mut TcpStream, code: u16) {
let _ = write!(
stream,
"HTTP/1.1 {code} Error\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
);
}
pub fn text_chunk(id: &str, text: &str) -> String {
serde_json::json!({
"id": id,
"choices": [{"index": 0, "delta": {"role": "assistant", "content": text}, "finish_reason": null}]
})
.to_string()
}
pub fn stop_chunk(id: &str) -> String {
serde_json::json!({
"id": id,
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]
})
.to_string()
}
pub fn tool_start_chunk(id: &str, call_id: &str, name: &str) -> String {
serde_json::json!({
"id": id,
"choices": [{"index": 0, "delta": {
"role": "assistant",
"content": null,
"tool_calls": [{"index": 0, "id": call_id, "type": "function",
"function": {"name": name, "arguments": ""}}]
}, "finish_reason": null}]
})
.to_string()
}
pub fn tool_args_chunk(id: &str, args_json: &str) -> String {
serde_json::json!({
"id": id,
"choices": [{"index": 0, "delta": {
"tool_calls": [{"index": 0, "function": {"arguments": args_json}}]
}, "finish_reason": null}]
})
.to_string()
}
pub fn tool_calls_stop_chunk(id: &str) -> String {
serde_json::json!({
"id": id,
"choices": [{"index": 0, "delta": {}, "finish_reason": "tool_calls"}]
})
.to_string()
}
pub fn write_config(config_home: &std::path::Path, base_url: &str) {
write_config_with(config_home, base_url, "");
}
pub fn write_config_with(config_home: &std::path::Path, base_url: &str, extra: &str) {
let dir = config_home.join("hrdr");
std::fs::create_dir_all(&dir).expect("config dir");
std::fs::write(
dir.join("config.toml"),
format!(
"model = \"mock://mock-model\"\n\
{extra}\n\
[providers.mock]\n\
base_url = \"{base_url}\"\n\
context_window = 200000\n"
),
)
.expect("write config.toml");
}
pub fn drain_pty(
mut reader: Box<dyn Read + Send>,
writer: Arc<Mutex<Box<dyn Write + Send>>>,
) -> Arc<Mutex<String>> {
let seen = Arc::new(Mutex::new(String::new()));
let sink = Arc::clone(&seen);
thread::spawn(move || {
let mut buf = [0u8; 8192];
loop {
match reader.read(&mut buf) {
Ok(0) => break,
Ok(n) => {
if buf[..n].windows(4).any(|w| w == b"\x1b[6n") {
let mut w = writer.lock().unwrap_or_else(|e| e.into_inner());
let _ = w.write_all(b"\x1b[1;1R");
let _ = w.flush();
}
let mut s = sink.lock().unwrap_or_else(|e| e.into_inner());
s.push_str(&String::from_utf8_lossy(&buf[..n]));
}
Err(e)
if matches!(
e.kind(),
std::io::ErrorKind::Interrupted | std::io::ErrorKind::WouldBlock
) =>
{
thread::sleep(std::time::Duration::from_millis(10));
}
Err(_) => break,
}
}
});
seen
}
pub fn pty_text(seen: &Arc<Mutex<String>>) -> String {
seen.lock().unwrap_or_else(|e| e.into_inner()).clone()
}