Luigit
repositories / smith

smith

There are many coding harnesses - but this one is fast

owned by admin

smith-cli/tests/e2e_rpc.rs

Raw
//! End-to-end `smith rpc` workflows over piped standard input and output
//! (`SMH-SPEC-SPEC0001`, Interface modes: JSON-RPC).
//!
//! Every read waits on a line or end of output with a deadline; nothing
//! sleeps or polls.

use serde_json::{Value, json};
use std::io::{BufRead, BufReader, Read, Write};
use std::net::TcpListener;
use std::process::{Child, ChildStdin, Command, Stdio};
use std::sync::mpsc::{self, Receiver};
use std::time::Duration;

const DEADLINE: Duration = Duration::from_secs(10);
const MOCK_REPLY: &str = "rpc mock provider ready";

/// One running `smith rpc` process; killed on drop if a test fails early.
struct Rpc {
    child: Child,
    stdin: Option<ChildStdin>,
    /// Raw stdout lines; `None` marks end of output.
    lines: Receiver<Option<String>>,
    /// Every line sent (`> `) and received (`< `), in order.
    transcript: Vec<String>,
}

impl Rpc {
    fn spawn(name: &str, args: &[&str], envs: &[(&str, &str)]) -> Self {
        let home =
            std::env::temp_dir().join(format!("smith-rpc-test-{name}-{}", std::process::id()));
        let _ = std::fs::remove_dir_all(&home);
        let mut child = Command::new(assert_cmd::cargo::cargo_bin("smith"))
            .arg("rpc")
            .args(args)
            .env("XDG_DATA_HOME", &home)
            .env_remove("SMITH_PROVIDER")
            .env_remove("SMITH_MODEL")
            .env_remove("SMITH_BASE_URL")
            .env_remove("ANTHROPIC_API_KEY")
            .env_remove("OPENAI_API_KEY")
            .env_remove("GEMINI_API_KEY")
            .envs(envs.iter().copied())
            .stdin(Stdio::piped())
            .stdout(Stdio::piped())
            .stderr(Stdio::null())
            .spawn()
            .unwrap();
        let stdout = child.stdout.take().unwrap();
        let (sender, lines) = mpsc::channel();
        std::thread::spawn(move || {
            for line in BufReader::new(stdout).lines() {
                if sender.send(Some(line.unwrap())).is_err() {
                    return;
                }
            }
            let _ = sender.send(None);
        });
        let stdin = child.stdin.take();
        Self {
            child,
            stdin,
            lines,
            transcript: Vec::new(),
        }
    }

    fn mock(name: &str) -> Self {
        Self::spawn(name, &[], &[])
    }

    fn send(&mut self, line: &str) {
        self.transcript.push(format!("> {line}"));
        let stdin = self.stdin.as_mut().unwrap();
        stdin.write_all(line.as_bytes()).unwrap();
        stdin.write_all(b"\n").unwrap();
        stdin.flush().unwrap();
    }

    fn next(&mut self) -> Value {
        match self.lines.recv_timeout(DEADLINE) {
            Ok(Some(line)) => {
                let value = serde_json::from_str(&line)
                    .unwrap_or_else(|error| panic!("stdout line is not JSON ({error}): {line}"));
                self.transcript.push(format!("< {line}"));
                value
            }
            Ok(None) => panic!("rpc output ended early"),
            Err(error) => panic!("no rpc line within {DEADLINE:?}: {error}"),
        }
    }

    /// Close stdin, then expect end of output and the exit status.
    fn close(mut self) -> std::process::ExitStatus {
        drop(self.stdin.take());
        match self.lines.recv_timeout(DEADLINE) {
            Ok(None) => {}
            Ok(Some(line)) => panic!("unexpected rpc line after close: {line}"),
            Err(error) => panic!("rpc did not end output within {DEADLINE:?}: {error}"),
        }
        self.child.wait().unwrap()
    }
}

impl Drop for Rpc {
    fn drop(&mut self) {
        if matches!(self.child.try_wait(), Ok(None)) {
            let _ = self.child.kill();
            let _ = self.child.wait();
        }
    }
}

#[test]
fn rpc_ping_correlates_the_request_id() {
    let mut rpc = Rpc::mock("ping");
    rpc.send(r#"{"id": 7, "method": "ping"}"#);
    assert_eq!(rpc.next(), json!({"id": 7, "pong": true}));
    rpc.send(r#"{"id": "seven", "method": "ping"}"#);
    assert_eq!(rpc.next(), json!({"id": "seven", "pong": true}));
    assert!(rpc.close().success());
}

#[test]
fn rpc_prompt_streams_start_delta_end_then_wait_and_status_report_settled() {
    let mut rpc = Rpc::mock("prompt");
    rpc.send(r#"{"id": 1, "method": "prompt", "params": {"input": "hi"}}"#);
    rpc.send(r#"{"id": 2, "method": "wait"}"#);
    assert_eq!(rpc.next(), json!({"id": 1, "event": "start"}));
    assert_eq!(
        rpc.next(),
        json!({"id": 1, "event": "delta", "text": MOCK_REPLY})
    );
    assert_eq!(
        rpc.next(),
        json!({
            "id": 1,
            "event": "end",
            "text": MOCK_REPLY,
            "reason": "EndTurn",
            "cost": {"input_tokens": 0, "output_tokens": 0},
        })
    );
    // The worker answers `wait` only after the prompt has settled.
    assert_eq!(rpc.next(), json!({"id": 2, "settled": true}));
    // `wait` answers once the aborted turn has settled; `status` then agrees.
    rpc.send(r#"{"id": 3, "method": "wait"}"#);
    assert_eq!(rpc.next(), json!({"id": 3, "settled": true}));
    rpc.send(r#"{"id": 4, "method": "status"}"#);
    assert_eq!(rpc.next(), json!({"id": 4, "settled": true}));
    assert!(rpc.close().success());
}

#[test]
fn rpc_malformed_line_reports_parse_error_with_null_id() {
    let mut rpc = Rpc::mock("malformed");
    rpc.send("not json");
    let record = rpc.next();
    assert_eq!(record["id"], Value::Null);
    assert_eq!(record["error"]["code"], "PARSE");
    // The server keeps serving after a bad line.
    rpc.send(r#"{"id": 1, "method": "ping"}"#);
    assert_eq!(rpc.next(), json!({"id": 1, "pong": true}));
    assert!(rpc.close().success());
}

#[test]
fn rpc_unknown_method_reports_method_unknown() {
    let mut rpc = Rpc::mock("unknown");
    rpc.send(r#"{"id": "x", "method": "explode"}"#);
    assert_eq!(
        rpc.next(),
        json!({
            "id": "x",
            "error": {"code": "METHOD_UNKNOWN", "message": "unknown method explode"},
        })
    );
    assert!(rpc.close().success());
}

#[test]
fn rpc_closing_stdin_after_settle_exits_zero() {
    let mut rpc = Rpc::mock("shutdown");
    rpc.send(r#"{"id": 1, "method": "prompt", "params": {"input": "hi"}}"#);
    rpc.send(r#"{"id": 2, "method": "wait"}"#);
    let mut records = Vec::new();
    loop {
        let record = rpc.next();
        let done = record["id"] == 2;
        records.push(record);
        if done {
            break;
        }
    }
    assert!(records.iter().any(|r| r["id"] == 1 && r["event"] == "end"));
    let status = rpc.close();
    assert_eq!(status.code(), Some(0));
}

/// A loopback provider endpoint that holds its response: it reports each
/// received request, then answers only when released.
struct StalledProvider {
    url: String,
    received: Receiver<()>,
    release: mpsc::Sender<()>,
}

impl StalledProvider {
    fn start() -> Self {
        let listener = TcpListener::bind("127.0.0.1:0").unwrap();
        let url = format!("http://{}", listener.local_addr().unwrap());
        let (received_tx, received) = mpsc::channel();
        let (release, release_rx) = mpsc::channel::<()>();
        std::thread::spawn(move || {
            let Ok((stream, _)) = listener.accept() else {
                return;
            };
            let mut reader = BufReader::new(stream);
            let mut length = 0;
            loop {
                let mut header = String::new();
                if reader.read_line(&mut header).unwrap_or(0) == 0 {
                    return;
                }
                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")
                {
                    length = value.trim().parse().unwrap();
                }
            }
            let mut body = vec![0; length];
            reader.read_exact(&mut body).unwrap();
            let _ = received_tx.send(());
            if release_rx.recv().is_err() {
                return;
            }
            let chunk = json!({"choices": [{"index": 0, "delta": {"content": "late"}}]});
            let response = format!(
                "HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\nconnection: close\r\n\r\ndata: {chunk}\n\n"
            );
            let _ = reader.get_mut().write_all(response.as_bytes());
        });
        Self {
            url,
            received,
            release,
        }
    }
}

#[test]
fn rpc_abort_ends_the_running_prompt_with_turn_failed() {
    let provider = StalledProvider::start();
    let mut rpc = Rpc::spawn(
        "abort",
        &[
            "--provider",
            "openai",
            "--model",
            "e2e-model",
            "--base-url",
            &provider.url,
        ],
        &[("OPENAI_API_KEY", "sk-e2e-fixture")],
    );
    // Each step waits for its reply before the next; the provider holds its
    // response until the abort is acknowledged, so the order is fixed.
    rpc.send(r#"{"id": 1, "method": "prompt", "params": {"input": "block"}}"#);
    rpc.next();
    provider
        .received
        .recv_timeout(DEADLINE)
        .expect("provider request within the deadline");
    rpc.send(r#"{"id": 2, "method": "abort"}"#);
    rpc.next();
    provider.release.send(()).unwrap();
    rpc.next();
    // `wait` answers once the aborted turn has settled; `status` then agrees.
    rpc.send(r#"{"id": 3, "method": "wait"}"#);
    rpc.next();
    rpc.send(r#"{"id": 4, "method": "status"}"#);
    rpc.next();
    insta::assert_snapshot!("rpc_prompt_then_abort_then_wait", rpc.transcript.join("\n"));
    assert!(rpc.close().success());
}