use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use salvor_cli::demo_script;
use serde_json::{Value, json};
use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader};
use tokio::net::{TcpListener, TcpStream};
const DEFAULT_PORT: u16 = 8899;
const DEFAULT_DELAY_MS: u64 = 300;
struct Settings {
port: u16,
delay: Duration,
}
fn resolve_u64(args: &[String], flag: &str, env: &str, default: u64) -> u64 {
if let Some(index) = args.iter().position(|arg| arg == flag) {
let raw = args
.get(index + 1)
.unwrap_or_else(|| panic!("{flag} needs a value"));
return raw
.parse()
.unwrap_or_else(|_| panic!("{flag} value `{raw}` is not a number"));
}
if let Ok(raw) = std::env::var(env)
&& !raw.is_empty()
{
return raw
.parse()
.unwrap_or_else(|_| panic!("{env}=`{raw}` is not a number"));
}
default
}
fn settings() -> Settings {
let args: Vec<String> = std::env::args().skip(1).collect();
let port = resolve_u64(
&args,
"--port",
"SALVOR_DEMO_MODEL_PORT",
u64::from(DEFAULT_PORT),
);
let delay = resolve_u64(
&args,
"--delay-ms",
"SALVOR_DEMO_MODEL_DELAY_MS",
DEFAULT_DELAY_MS,
);
Settings {
port: u16::try_from(port).expect("port fits in u16"),
delay: Duration::from_millis(delay),
}
}
fn response_for(script: &[(usize, Value)], count: usize) -> (u16, Value) {
for (expected, response) in script {
if *expected == count {
return (200, response.clone());
}
}
(
500,
json!({
"error": {
"type": "demo_script",
"message": format!("no scripted response for {count} messages")
}
}),
)
}
async fn serve_connection(
stream: TcpStream,
script: Arc<Vec<(usize, Value)>>,
delay: Duration,
requests: Arc<AtomicUsize>,
) -> std::io::Result<()> {
let (read_half, mut write_half) = stream.into_split();
let mut reader = BufReader::new(read_half);
loop {
let mut request_line = String::new();
if reader.read_line(&mut request_line).await? == 0 {
return Ok(());
}
let path = request_line
.split_whitespace()
.nth(1)
.unwrap_or("")
.to_owned();
let mut content_length = 0usize;
loop {
let mut header = String::new();
if reader.read_line(&mut header).await? == 0 {
return Ok(());
}
let header = header.trim_end();
if header.is_empty() {
break;
}
if let Some((name, value)) = header.split_once(':')
&& name.eq_ignore_ascii_case("content-length")
{
content_length = value.trim().parse().unwrap_or(0);
}
}
let mut body = vec![0u8; content_length];
reader.read_exact(&mut body).await?;
let (status, payload) = if path == "/v1/messages" {
let count = serde_json::from_slice::<Value>(&body)
.ok()
.and_then(|value| {
value
.get("messages")
.and_then(Value::as_array)
.map(Vec::len)
})
.unwrap_or(0);
let nth = requests.fetch_add(1, Ordering::SeqCst) + 1;
let (status, response) = response_for(&script, count);
eprintln!(
"[salvor-demo-model] request #{nth}: {count} messages -> {}",
if status == 200 {
"scripted"
} else {
"unscripted (500)"
}
);
(status, response)
} else {
(
404,
json!({ "error": { "type": "not_found", "message": path } }),
)
};
tokio::time::sleep(delay).await;
let serialized = serde_json::to_vec(&payload).expect("response serializes");
let reason = if status == 200 { "OK" } else { "Error" };
let head = format!(
"HTTP/1.1 {status} {reason}\r\n\
content-type: application/json\r\n\
content-length: {}\r\n\
connection: keep-alive\r\n\r\n",
serialized.len()
);
write_half.write_all(head.as_bytes()).await?;
write_half.write_all(&serialized).await?;
write_half.flush().await?;
}
}
#[tokio::main]
async fn main() -> std::io::Result<()> {
let settings = settings();
let script = Arc::new(demo_script::script());
let requests = Arc::new(AtomicUsize::new(0));
let listener = TcpListener::bind(("127.0.0.1", settings.port)).await?;
let port = listener.local_addr()?.port();
println!("salvor-demo-model listening on http://127.0.0.1:{port}");
eprintln!(
"[salvor-demo-model] serving {} scripted turns, {:?} per turn",
script.len(),
settings.delay
);
loop {
let (stream, _) = listener.accept().await?;
let script = script.clone();
let requests = requests.clone();
let delay = settings.delay;
tokio::spawn(async move {
if let Err(error) = serve_connection(stream, script, delay, requests).await {
eprintln!("[salvor-demo-model] connection ended: {error}");
}
});
}
}