Luigit
repositories / smith

smith

There are many coding harnesses - but this one is fast

owned by admin

smith-rpc/src/lib.rs

Raw
//! LF-delimited JSON-RPC over standard input and output (`SMH-SPEC-SPEC0001`,
//! Interface modes).
//!
//! Standard output carries protocol records only. Request IDs are preserved
//! on every response. Agent work runs on one worker, leaving control requests
//! able to inspect, abort, and wait for the active turn.

use smith::error::Result;
use smith::tool::CancelHandle;
use smith_core::agent::{Agent, AgentEvent};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering::SeqCst};
use std::sync::mpsc::{self, Receiver, SyncSender};

/// Commands queued for the worker before it must answer.
const COMMAND_QUEUE: usize = 64;

/// Response lines the worker may produce before the reader drains them.
const RESPONSE_QUEUE: usize = 256;

/// The RPC session state; owned by the reading thread.
///
/// Shared state with the worker is two atomics and channels: the worker
/// owns the agent, the reader owns the responses.
pub struct RpcState {
    commands: SyncSender<WorkerCommand>,
    responses: Receiver<RpcResponse>,
    /// True when no prompt runs; cleared by the reader on `prompt`, set by
    /// the worker after the turn.
    settled: Arc<AtomicBool>,
    /// Abort requested before the worker exposed the turn's cancel handle;
    /// the worker consumes it when the turn starts.
    abort_requested: Arc<AtomicBool>,
    /// The worker hands over each turn's cancel handle here.
    cancels: Receiver<CancelHandle>,
    active_cancel: Option<CancelHandle>,
}

impl RpcState {
    /// Start one worker owning `agent`.
    #[must_use]
    pub fn new(agent: Agent) -> Self {
        let (command_tx, command_rx) = mpsc::sync_channel(COMMAND_QUEUE);
        let (response_tx, response_rx) = mpsc::sync_channel(RESPONSE_QUEUE);
        let (cancel_tx, cancel_rx) = mpsc::sync_channel(1);
        let settled = Arc::new(AtomicBool::new(true));
        let abort_requested = Arc::new(AtomicBool::new(false));
        let worker = Worker {
            agent,
            responses: response_tx,
            cancels: cancel_tx,
            settled: Arc::clone(&settled),
            abort_requested: Arc::clone(&abort_requested),
        };
        std::thread::spawn(move || worker.run(&command_rx));
        Self {
            commands: command_tx,
            responses: response_rx,
            settled,
            abort_requested,
            cancels: cancel_rx,
            active_cancel: None,
        }
    }

    /// Remove all completed asynchronous response lines without blocking.
    #[must_use]
    pub fn drain_lines(&self) -> Vec<String> {
        self.responses
            .try_iter()
            .flat_map(|response| line_or_empty(&response))
            .collect()
    }

    /// Block up to `timeout` for the next response line.
    ///
    /// The deterministic settle signal for callers that must wait: pair a
    /// `wait` request with this receive instead of polling [`Self::is_settled`].
    #[must_use]
    pub fn recv_line(&self, timeout: std::time::Duration) -> Option<String> {
        let response = self.responses.recv_timeout(timeout).ok()?;
        response.to_line().ok()
    }

    /// Whether no prompt is currently running.
    #[must_use]
    pub fn is_settled(&self) -> bool {
        self.settled.load(SeqCst)
    }

    /// The newest cancel handle the worker exposed, if any.
    fn refresh_active_cancel(&mut self) {
        if let Some(cancel) = self.cancels.try_iter().last() {
            self.active_cancel = Some(cancel);
        }
    }
}

/// One incoming request line.
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RpcRequest {
    /// Correlation ID; echoed on every response for this request.
    pub id: serde_json::Value,
    /// Command name.
    pub method: String,
    /// Command parameters, when present.
    pub params: serde_json::Value,
}

/// One outgoing response line.
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RpcResponse {
    /// Correlation ID of the owning request.
    pub id: serde_json::Value,
    /// Record payload: result or error.
    pub record: serde_json::Value,
}

impl RpcResponse {
    /// Serialize as one LF-terminated JSON line.
    ///
    /// # Errors
    ///
    /// Returns a serialization fault when the record cannot encode.
    pub fn to_line(&self) -> Result<String> {
        let mut value = serde_json::json!({"id": self.id});
        merge(&mut value, &self.record);
        let mut line =
            serde_json::to_string(&value).map_err(|error| smith::error::SmithError::Session {
                code: "RPC_ENCODE".to_string(),
                message: error.to_string(),
            })?;
        line.push('\n');
        Ok(line)
    }
}

enum WorkerCommand {
    Prompt { request: RpcRequest, input: String },
    Wait { request: RpcRequest },
}

struct Worker {
    agent: Agent,
    responses: SyncSender<RpcResponse>,
    cancels: SyncSender<CancelHandle>,
    settled: Arc<AtomicBool>,
    abort_requested: Arc<AtomicBool>,
}

impl Worker {
    /// Serve commands in order; a `wait` queued behind a prompt answers
    /// once that prompt has finished.
    fn run(mut self, commands: &Receiver<WorkerCommand>) {
        while let Ok(command) = commands.recv() {
            match command {
                WorkerCommand::Prompt { request, input } => self.prompt(&request, &input),
                WorkerCommand::Wait { request } => {
                    let _ = self
                        .responses
                        .send(result(&request, serde_json::json!({"settled": true})));
                }
            }
        }
    }

    fn prompt(&mut self, request: &RpcRequest, input: &str) {
        let cancel = self.agent.renew_cancel_handle();
        // Publish the handle first, then look for an abort that raced it:
        // either the reader sees the handle or the worker sees the flag.
        let _ = self.cancels.send(cancel.clone());
        if self.abort_requested.swap(false, SeqCst) {
            cancel.cancel();
        }
        emit_turn(&mut self.agent, request, input, &self.responses);
        self.settled.store(true, SeqCst);
    }
}

fn emit_turn(
    agent: &mut Agent,
    request: &RpcRequest,
    input: &str,
    responses: &SyncSender<RpcResponse>,
) {
    match agent.run_turn(input) {
        Ok(outcome) => {
            for event in &outcome.events {
                if let AgentEvent::Streamed {
                    event: smith::stream::StreamEvent::TextDelta { delta, .. },
                } = event
                {
                    let _ = responses.send(result(
                        request,
                        serde_json::json!({"event": "delta", "text": delta}),
                    ));
                }
            }
            let _ = responses.send(result(
                request,
                serde_json::json!({
                    "event": "end",
                    "text": outcome.text,
                    "reason": format!("{:?}", outcome.reason),
                    "cost": {
                        "input_tokens": outcome.cost.input_tokens,
                        "output_tokens": outcome.cost.output_tokens,
                    },
                }),
            ));
        }
        Err(error) => {
            let _ = responses.send(RpcResponse {
                id: request.id.clone(),
                record: serde_json::json!({
                    "event": "end",
                    "error": {"code": "TURN_FAILED", "message": error.to_string()},
                }),
            });
        }
    }
}

fn merge(target: &mut serde_json::Value, patch: &serde_json::Value) {
    if let (Some(target), Some(patch)) = (target.as_object_mut(), patch.as_object()) {
        for (key, value) in patch {
            target.insert(key.clone(), value.clone());
        }
    }
}

/// Handle one input line, returning immediate records.
///
/// Prompt completion and deferred wait records are available through
/// [`RpcState::drain_lines`]. Blank lines produce no records.
#[must_use]
pub fn handle_line(state: &mut RpcState, line: &str) -> Vec<String> {
    let trimmed = line.trim();
    if trimmed.is_empty() {
        return Vec::new();
    }
    let parsed: serde_json::Value = match serde_json::from_str(trimmed) {
        Ok(value) => value,
        Err(error) => {
            return line_or_empty(&RpcResponse {
                id: serde_json::Value::Null,
                record: serde_json::json!({
                    "error": {"code": "PARSE", "message": error.to_string()},
                }),
            });
        }
    };
    let request = RpcRequest {
        id: parsed.get("id").cloned().unwrap_or(serde_json::Value::Null),
        method: parsed
            .get("method")
            .and_then(serde_json::Value::as_str)
            .unwrap_or_default()
            .to_string(),
        params: parsed
            .get("params")
            .cloned()
            .unwrap_or(serde_json::json!({})),
    };
    dispatch(state, request)
        .into_iter()
        .flat_map(|response| line_or_empty(&response))
        .collect()
}

fn line_or_empty(response: &RpcResponse) -> Vec<String> {
    response
        .to_line()
        .map_or_else(|_| Vec::new(), |line| vec![line])
}

fn dispatch(state: &mut RpcState, request: RpcRequest) -> Vec<RpcResponse> {
    match request.method.as_str() {
        "ping" => vec![result(&request, serde_json::json!({"pong": true}))],
        "prompt" => start_prompt(state, &request),
        "abort" => {
            let active = !state.is_settled();
            if active {
                // Flag first, then look for the handle: mirrors the worker.
                state.abort_requested.store(true, SeqCst);
                state.refresh_active_cancel();
                if let Some(cancel) = &state.active_cancel {
                    cancel.cancel();
                }
            }
            vec![result(&request, serde_json::json!({"aborting": active}))]
        }
        "status" => vec![result(
            &request,
            serde_json::json!({"settled": state.is_settled()}),
        )],
        "wait" => match state.commands.try_send(WorkerCommand::Wait {
            request: request.clone(),
        }) {
            Ok(()) => Vec::new(),
            Err(_) => vec![error(&request, "STATE", "agent worker unavailable")],
        },
        other => vec![RpcResponse {
            id: request.id,
            record: serde_json::json!({
                "error": {"code": "METHOD_UNKNOWN", "message": format!("unknown method {other}")},
            }),
        }],
    }
}

fn start_prompt(state: &mut RpcState, request: &RpcRequest) -> Vec<RpcResponse> {
    if !state.is_settled() {
        return vec![error(request, "BUSY", "a prompt is already running")];
    }
    // Drop the finished turn's handle so the worker's next send never waits.
    state.refresh_active_cancel();
    state.active_cancel = None;
    state.abort_requested.store(false, SeqCst);
    state.settled.store(false, SeqCst);
    let input = request
        .params
        .get("input")
        .and_then(serde_json::Value::as_str)
        .unwrap_or_default()
        .to_string();
    let queued = state.commands.try_send(WorkerCommand::Prompt {
        request: request.clone(),
        input,
    });
    if queued.is_err() {
        state.settled.store(true, SeqCst);
        return vec![error(request, "STATE", "agent worker unavailable")];
    }
    vec![result(request, serde_json::json!({"event": "start"}))]
}

fn result(request: &RpcRequest, record: serde_json::Value) -> RpcResponse {
    RpcResponse {
        id: request.id.clone(),
        record,
    }
}

fn error(request: &RpcRequest, code: &str, message: &str) -> RpcResponse {
    RpcResponse {
        id: request.id.clone(),
        record: serde_json::json!({"error": {"code": code, "message": message}}),
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use futures::StreamExt;
    use smith::provider::StreamFn;
    use smith_core::tools::ToolSession;
    use smith_harness::mock_text_reply;
    use smith_harness::runtime::{RuntimeConfig, build_agent};
    use std::sync::atomic::Ordering;
    use std::time::Duration;

    /// Upper bound on any single wait for the worker; generous for loaded CI.
    const DEADLINE: Duration = Duration::from_secs(10);

    fn state() -> RpcState {
        static NEXT_DIR: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
        let sequence = NEXT_DIR.fetch_add(1, Ordering::Relaxed);
        let dir =
            std::env::temp_dir().join(format!("smith_rpc_{}_{}", std::process::id(), sequence));
        let _ = std::fs::remove_dir_all(&dir);
        std::fs::create_dir_all(&dir).unwrap();
        let config = RuntimeConfig::mock(dir, vec![mock_text_reply("hello world")]);
        RpcState::new(build_agent(&config).unwrap())
    }

    fn values(lines: &[String]) -> Vec<serde_json::Value> {
        lines
            .iter()
            .map(|line| serde_json::from_str(line.trim()).unwrap())
            .collect()
    }

    /// The next response record, or a failure once `DEADLINE` passes.
    fn next(state: &RpcState) -> serde_json::Value {
        let line = state.recv_line(DEADLINE);
        assert!(
            line.is_some(),
            "worker answered nothing within {DEADLINE:?}"
        );
        serde_json::from_str(line.unwrap_or_default().trim()).unwrap()
    }

    /// Every record before the answer to the `wait` with `id`.
    fn until_waited(state: &RpcState, id: &serde_json::Value) -> Vec<serde_json::Value> {
        let mut records = Vec::new();
        loop {
            let record = next(state);
            if record["id"] == *id {
                assert_eq!(record["settled"], true);
                return records;
            }
            records.push(record);
        }
    }

    fn complete(state: &mut RpcState, input: &str) -> Vec<serde_json::Value> {
        let mut records = values(&handle_line(state, input));
        // synchronize: a `wait` queued behind the prompt answers only after the turn settled.
        assert!(handle_line(state, r#"{"id": "settle", "method": "wait"}"#).is_empty());
        records.extend(until_waited(state, &serde_json::json!("settle")));
        assert!(state.is_settled());
        records
    }

    #[test]
    fn ping_correlates_and_unknown_methods_error_with_the_same_id() {
        let mut state = state();
        let pong = &values(&handle_line(&mut state, r#"{"id": 7, "method": "ping"}"#))[0];
        assert_eq!(pong["id"], 7);
        assert_eq!(pong["pong"], true);
        let unknown = &values(&handle_line(
            &mut state,
            r#"{"id": "abc", "method": "explode"}"#,
        ))[0];
        assert_eq!(unknown["id"], "abc");
        assert_eq!(unknown["error"]["code"], "METHOD_UNKNOWN");
        let malformed = &values(&handle_line(&mut state, "not json"))[0];
        assert_eq!(malformed["error"]["code"], "PARSE");
    }

    #[test]
    fn prompt_streams_deltas_and_assembles_the_final_message() {
        let mut state = state();
        let records = complete(
            &mut state,
            r#"{"id": 1, "method": "prompt", "params": {"input": "hi"}}"#,
        );
        assert_eq!(records[0]["event"], "start");
        assert_eq!(records[1]["event"], "delta");
        assert_eq!(records[1]["text"], "hello world");
        let end = records.last().unwrap();
        assert_eq!(end["event"], "end");
        assert_eq!(end["text"], "hello world");
        let assembled: String = records
            .iter()
            .filter(|record| record["event"] == "delta")
            .filter_map(|record| record["text"].as_str().map(str::to_string))
            .collect();
        assert_eq!(assembled, end["text"].as_str().unwrap_or_default());
        assert!(records.iter().all(|record| record["id"] == 1));
    }

    #[test]
    fn active_prompt_reports_busy_aborts_waits_and_can_restart() {
        let stream_calls = std::sync::atomic::AtomicUsize::new(0);
        let (started_tx, started) = mpsc::sync_channel::<()>(1);
        let (release, release_rx) = mpsc::sync_channel::<()>(1);
        let release_rx = Arc::new(std::sync::Mutex::new(release_rx));
        let stream: StreamFn = Arc::new(move |_request, cancel| {
            if stream_calls.fetch_add(1, Ordering::Relaxed) > 0 {
                return futures::stream::iter(mock_text_reply("restarted").into_iter().map(Ok))
                    .boxed();
            }
            let _ = started_tx.send(());
            let release_rx = Arc::clone(&release_rx);
            let items = std::iter::once_with(move || {
                // synchronize: the turn holds until the test has aborted it.
                let released = release_rx
                    .lock()
                    .is_ok_and(|rx| rx.recv_timeout(DEADLINE).is_ok());
                if released && cancel.is_cancelled() {
                    vec![Err(smith::error::SmithError::Cancelled)]
                } else {
                    mock_text_reply("not aborted").into_iter().map(Ok).collect()
                }
            })
            .flatten();
            futures::stream::iter(items).boxed()
        });
        let tools = ToolSession::new(std::env::temp_dir());
        let mut state = RpcState::new(Agent::new(stream, tools, "mock-smith"));

        let start = values(&handle_line(
            &mut state,
            r#"{"id": 1, "method": "prompt", "params": {"input": "block"}}"#,
        ));
        assert_eq!(start[0]["event"], "start");
        // synchronize: the provider call is the signal that the turn runs.
        assert_eq!(started.recv_timeout(DEADLINE), Ok(()));
        let status = values(&handle_line(&mut state, r#"{"id": 2, "method": "status"}"#));
        assert_eq!(status[0]["settled"], false);
        let busy = values(&handle_line(
            &mut state,
            r#"{"id": 3, "method": "prompt", "params": {"input": "too soon"}}"#,
        ));
        assert_eq!(busy[0]["error"]["code"], "BUSY");
        assert!(handle_line(&mut state, r#"{"id": 4, "method": "wait"}"#).is_empty());
        let abort = values(&handle_line(&mut state, r#"{"id": 5, "method": "abort"}"#));
        assert_eq!(abort[0]["aborting"], true);
        release.send(()).unwrap();

        let deferred = until_waited(&state, &serde_json::json!(4));
        assert!(deferred.iter().any(|record| {
            record["id"] == 1
                && record["event"] == "end"
                && record["error"]["code"] == "TURN_FAILED"
        }));
        assert!(state.is_settled());

        let restarted = complete(
            &mut state,
            r#"{"id": 6, "method": "prompt", "params": {"input": "again"}}"#,
        );
        assert_eq!(restarted.last().unwrap()["text"], "restarted");
    }

    #[test]
    fn blank_lines_are_ignored() {
        let mut state = state();
        assert!(handle_line(&mut state, "   ").is_empty());
    }
}