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());
}
}