use std::net::{Ipv4Addr, SocketAddr, TcpStream};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use serde_json::{Value, json};
use tokio_tungstenite::tungstenite::{self, Message};
use crate::debug::protocol::{WatchTarget, reply_ok, validate_payload};
const CONNECT_TIMEOUT: Duration = Duration::from_secs(5);
const IO_TIMEOUT: Duration = Duration::from_secs(5);
const EXIT_TRANSPORT: i32 = 3;
type WsStream = tungstenite::WebSocket<TcpStream>;
fn connect(port: u16) -> Result<WsStream, String> {
let addr = SocketAddr::from((Ipv4Addr::LOCALHOST, port));
let stream = TcpStream::connect_timeout(&addr, CONNECT_TIMEOUT).map_err(|e| {
format!("cannot connect to ws://127.0.0.1:{port}: {e} (is `cn debug` running?)")
})?;
stream
.set_read_timeout(Some(IO_TIMEOUT))
.and_then(|()| stream.set_write_timeout(Some(IO_TIMEOUT)))
.map_err(|e| format!("cannot configure socket timeouts: {e}"))?;
let url = format!("ws://127.0.0.1:{port}/");
match tungstenite::client::client(url.as_str(), stream) {
Ok((ws, _resp)) => Ok(ws),
Err(tungstenite::HandshakeError::Failure(e)) => {
Err(format!("websocket handshake failed on port {port}: {e}"))
}
Err(tungstenite::HandshakeError::Interrupted(_)) => Err(format!(
"websocket handshake timed out on port {port} (>{}s)",
IO_TIMEOUT.as_secs()
)),
}
}
fn read_text(ws: &mut WsStream) -> Result<String, String> {
loop {
match ws.read() {
Ok(Message::Text(text)) => return Ok(text.to_string()),
Ok(Message::Ping(payload)) => {
let _ = ws.send(Message::Pong(payload));
}
Ok(Message::Close(_)) => return Err("server closed the connection".to_string()),
Ok(_) => {}
Err(e) => return Err(map_read_error(e)),
}
}
}
fn map_read_error(e: tungstenite::Error) -> String {
if let tungstenite::Error::Io(io) = &e
&& matches!(
io.kind(),
std::io::ErrorKind::WouldBlock | std::io::ErrorKind::TimedOut
)
{
return format!("timed out waiting for reply (>{}s)", IO_TIMEOUT.as_secs());
}
format!("read failed: {e}")
}
fn request(port: u16, payload: &str) -> Result<Value, String> {
let mut ws = connect(port)?;
ws.send(Message::text(payload))
.map_err(|e| format!("failed to send request: {e}"))?;
let reply = read_text(&mut ws)?;
let _ = ws.close(None);
let _ = ws.flush();
serde_json::from_str(&reply).map_err(|e| format!("malformed reply from server: {e}"))
}
fn request_cmd(port: u16, cmd: &str) -> Result<Value, String> {
request(port, &json!({ "cmd": cmd }).to_string())
}
fn print_reply(reply: &Value) {
println!(
"{}",
serde_json::to_string_pretty(reply).unwrap_or_else(|_| reply.to_string())
);
}
fn fail_transport(msg: &str) -> ! {
eprintln!("cn debug: {msg}");
std::process::exit(EXIT_TRANSPORT);
}
pub fn send(port: u16, json: &str) -> std::io::Result<()> {
let payload = match validate_payload(json) {
Ok(p) => p,
Err(msg) => {
eprintln!("cn debug send: {msg}");
std::process::exit(EXIT_TRANSPORT);
}
};
let reply = request(port, &payload).unwrap_or_else(|msg| fail_transport(&msg));
print_reply(&reply);
if reply_ok(&reply) {
Ok(())
} else {
std::process::exit(1);
}
}
pub fn screenshot(port: u16, path: &str) -> std::io::Result<()> {
let abs = std::path::absolute(path).unwrap_or_else(|_| std::path::PathBuf::from(path));
let abs = abs.to_string_lossy().to_string();
let reply = request(
port,
&json!({ "cmd": "screenshot", "path": abs }).to_string(),
)
.unwrap_or_else(|msg| fail_transport(&msg));
if reply_ok(&reply) {
let saved = reply.get("path").and_then(Value::as_str).unwrap_or(&abs);
println!("[screenshot] saved: {saved}");
Ok(())
} else {
let err = reply
.get("error")
.and_then(Value::as_str)
.unwrap_or("unknown error");
eprintln!("[screenshot] failed: {err}");
std::process::exit(1);
}
}
fn sleep_interruptible(total_ms: u64, running: &AtomicBool) {
let mut left = total_ms;
while left > 0 && running.load(Ordering::SeqCst) {
let chunk = left.min(100);
std::thread::sleep(Duration::from_millis(chunk));
left -= chunk;
}
}
pub fn watch(port: u16, target: WatchTarget, interval_ms: u64) -> std::io::Result<()> {
let running = Arc::new(AtomicBool::new(true));
let flag = Arc::clone(&running);
let _ = ctrlc::set_handler(move || flag.store(false, Ordering::SeqCst));
let cmd = target.cmd();
println!(
"[watch] polling {} on ws://127.0.0.1:{port} every {interval_ms}ms (Ctrl-C to stop)",
target.label()
);
let mut ever_ok = false;
while running.load(Ordering::SeqCst) {
match request_cmd(port, cmd) {
Ok(reply) => {
ever_ok = true;
print_reply(&reply);
}
Err(msg) => {
if !ever_ok {
fail_transport(&msg);
}
eprintln!("[watch] {msg}");
}
}
sleep_interruptible(interval_ms, &running);
}
println!();
Ok(())
}
#[cfg(test)]
mod tests {
use std::time::Instant;
use super::*;
const DEAD_PORT_CANDIDATES: [u16; 4] = [28474, 28475, 28476, 28477];
fn find_dead_port() -> Option<u16> {
DEAD_PORT_CANDIDATES
.into_iter()
.find(|&port| connect_refused(port))
}
fn connect_refused(port: u16) -> bool {
let addr = SocketAddr::from((Ipv4Addr::LOCALHOST, port));
matches!(
TcpStream::connect_timeout(&addr, Duration::from_millis(250)),
Err(e) if e.kind() == std::io::ErrorKind::ConnectionRefused
)
}
fn skip_no_dead_port(test: &str) {
eprintln!("[skip] {test}: no localhost port refuses connections on this host");
}
#[test]
fn connect_to_dead_port_errors_fast() {
let Some(port) = find_dead_port() else {
return skip_no_dead_port("connect_to_dead_port_errors_fast");
};
let start = Instant::now();
let err = connect(port).expect_err("connect to a dead port must fail");
assert!(
start.elapsed() < CONNECT_TIMEOUT,
"a refused connection should fail well before the connect timeout"
);
assert!(err.contains("cannot connect"), "unexpected error: {err}");
}
#[test]
fn request_to_dead_port_errors() {
let Some(port) = find_dead_port() else {
return skip_no_dead_port("request_to_dead_port_errors");
};
assert!(request_cmd(port, "state").is_err());
}
}