xz-mcp-engine 0.1.0

Engine implementations for xz-mcp-core: stdio and HTTP MCP clients, connection manager
Documentation
//! Stdio transport for MCP (`JSON-RPC` over stdin/stdout).
//!
//! Supports both:
//! - **Content-Length framing** (MCP / LSP standard)
//! - **Newline-delimited JSON** (legacy simple servers)

use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;

use async_trait::async_trait;
use serde_json::Value;
use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader};
use tokio::process::{Child, Command};
use tokio::sync::Mutex;
use xz_mcp_core::{McpClient, McpError, McpTool, McpToolResult};

static NEXT_ID: AtomicU64 = AtomicU64::new(1);

type Writer = tokio::io::BufWriter<tokio::process::ChildStdin>;
type Reader = BufReader<tokio::process::ChildStdout>;

struct Connection {
    _child: Child,
    writer: Mutex<Writer>,
    reader: Mutex<Reader>,
}

/// MCP client that spawns a subprocess and speaks JSON-RPC on stdio.
pub struct StdioMcpClient {
    command: String,
    args: Vec<String>,
    env: HashMap<String, String>,
    /// Prefer Content-Length framing when writing (MCP default).
    use_content_length: bool,
    conn: Mutex<Option<Arc<Connection>>>,
}

impl StdioMcpClient {
    /// Create a client for `command` with `args` and environment overrides.
    pub fn new(
        command: impl Into<String>,
        args: Vec<String>,
        env: HashMap<String, String>,
    ) -> Self {
        Self {
            command: command.into(),
            args,
            env,
            use_content_length: true,
            conn: Mutex::new(None),
        }
    }

    /// Prefer newline-delimited JSON writes (legacy servers).
    pub fn with_newline_framing(mut self) -> Self {
        self.use_content_length = false;
        self
    }

    fn next_id() -> u64 {
        NEXT_ID.fetch_add(1, Ordering::SeqCst)
    }
}

#[async_trait]
impl McpClient for StdioMcpClient {
    async fn connect(&mut self) -> Result<(), McpError> {
        let mut cmd = Command::new(&self.command);
        cmd.args(&self.args);
        for (k, v) in &self.env {
            cmd.env(k, v);
        }
        cmd.stdin(std::process::Stdio::piped());
        cmd.stdout(std::process::Stdio::piped());
        cmd.stderr(std::process::Stdio::piped());
        cmd.kill_on_drop(true);

        let mut child = cmd
            .spawn()
            .map_err(|e| McpError::Connection(format!("spawn: {e}")))?;
        let stdin = child
            .stdin
            .take()
            .ok_or_else(|| McpError::Connection("no stdin".into()))?;
        let stdout = child
            .stdout
            .take()
            .ok_or_else(|| McpError::Connection("no stdout".into()))?;

        let conn = Arc::new(Connection {
            _child: child,
            writer: Mutex::new(tokio::io::BufWriter::new(stdin)),
            reader: Mutex::new(BufReader::new(stdout)),
        });

        let init_req = rpc(
            "initialize",
            serde_json::json!({
                "protocolVersion": "2024-11-05",
                "capabilities": {},
                "clientInfo": { "name": "xz-mcp", "version": "0.1" }
            }),
        );
        send(&conn, &init_req, self.use_content_length).await?;

        let notif = serde_json::json!({"jsonrpc":"2.0","method":"notifications/initialized"});
        let _ = send_raw(&conn, &notif, self.use_content_length).await;

        *self.conn.lock().await = Some(conn);
        Ok(())
    }

    async fn list_tools(&self) -> Result<Vec<McpTool>, McpError> {
        let conn = self
            .conn
            .lock()
            .await
            .clone()
            .ok_or_else(|| McpError::Connection("not connected".into()))?;
        let req = rpc("tools/list", serde_json::json!({}));
        let resp = send(&conn, &req, self.use_content_length).await?;
        let tools = resp["result"]["tools"]
            .as_array()
            .ok_or_else(|| McpError::Protocol("missing tools array".into()))?;
        tools
            .iter()
            .map(|t| {
                Ok(McpTool {
                    name: t["name"].as_str().unwrap_or("").into(),
                    description: t["description"].as_str().unwrap_or("").into(),
                    input_schema: t
                        .get("inputSchema")
                        .cloned()
                        .unwrap_or_else(|| serde_json::json!({})),
                })
            })
            .collect()
    }

    async fn call_tool(&self, name: &str, args: Value) -> Result<McpToolResult, McpError> {
        let conn = self
            .conn
            .lock()
            .await
            .clone()
            .ok_or_else(|| McpError::Connection("not connected".into()))?;
        let req = rpc(
            "tools/call",
            serde_json::json!({"name":name,"arguments":args}),
        );
        let resp = send(&conn, &req, self.use_content_length).await?;
        let result = &resp["result"];
        let content: Vec<Value> = result["content"].as_array().cloned().unwrap_or_default();
        Ok(McpToolResult {
            content: serde_json::from_value(Value::Array(content)).unwrap_or_default(),
            is_error: result["isError"].as_bool().unwrap_or(false),
        })
    }

    async fn is_alive(&self) -> bool {
        self.conn.lock().await.is_some()
    }
}

fn rpc(method: &str, params: Value) -> Value {
    serde_json::json!({
        "jsonrpc": "2.0",
        "id": StdioMcpClient::next_id(),
        "method": method,
        "params": params,
    })
}

fn encode_content_length(body: &str) -> Vec<u8> {
    let header = format!("Content-Length: {}\r\n\r\n", body.len());
    let mut out = Vec::with_capacity(header.len() + body.len());
    out.extend_from_slice(header.as_bytes());
    out.extend_from_slice(body.as_bytes());
    out
}

async fn send_raw(conn: &Connection, msg: &Value, content_length: bool) -> Result<(), McpError> {
    let body = serde_json::to_string(msg)?;
    let bytes = if content_length {
        encode_content_length(&body)
    } else {
        let mut b = body.into_bytes();
        b.push(b'\n');
        b
    };
    let mut w = conn.writer.lock().await;
    w.write_all(&bytes).await.map_err(McpError::Io)?;
    w.flush().await.map_err(McpError::Io)
}

async fn send(conn: &Connection, req: &Value, content_length: bool) -> Result<Value, McpError> {
    let req_id = req["id"].clone();
    send_raw(conn, req, content_length).await?;

    let mut r = conn.reader.lock().await;
    loop {
        let msg = read_json_message(&mut r).await?;
        if msg.get("id").is_none() {
            continue; // notification
        }
        if msg["id"] != req_id {
            continue;
        }
        if let Some(err) = msg.get("error") {
            return Err(McpError::Server(err.to_string()));
        }
        return Ok(msg);
    }
}

/// Read one JSON-RPC message: Content-Length framed or newline-delimited.
async fn read_json_message(reader: &mut Reader) -> Result<Value, McpError> {
    let mut first = String::new();
    let n = reader
        .read_line(&mut first)
        .await
        .map_err(McpError::Io)?;
    if n == 0 {
        return Err(McpError::Connection(
            "server closed stdout while waiting for response".into(),
        ));
    }

    // Skip blank lines.
    let mut header_line = first;
    loop {
        let trimmed = header_line.trim_end_matches(['\r', '\n']);
        if !trimmed.is_empty() {
            break;
        }
        header_line.clear();
        let n = reader.read_line(&mut header_line).await.map_err(McpError::Io)?;
        if n == 0 {
            return Err(McpError::Connection(
                "server closed stdout while waiting for response".into(),
            ));
        }
    }

    let trimmed = header_line.trim_end_matches(['\r', '\n']);

    if let Some(rest) = trimmed
        .strip_prefix("Content-Length:")
        .or_else(|| trimmed.strip_prefix("content-length:"))
    {
        let content_length: usize = rest.trim().parse().map_err(|e| {
            McpError::Protocol(format!("Invalid Content-Length '{rest}': {e}"))
        })?;

        // Consume remaining headers until blank line.
        loop {
            let mut line = String::new();
            let n = reader.read_line(&mut line).await.map_err(McpError::Io)?;
            if n == 0 {
                return Err(McpError::Connection(
                    "server closed stdout mid-headers".into(),
                ));
            }
            if line == "\n" || line == "\r\n" || line.trim().is_empty() {
                break;
            }
        }

        let mut buf = vec![0u8; content_length];
        reader.read_exact(&mut buf).await.map_err(McpError::Io)?;
        let response: Value =
            serde_json::from_slice(&buf).map_err(|e| McpError::Protocol(e.to_string()))?;
        return Ok(response);
    }

    // Newline-delimited JSON body (first non-empty line).
    serde_json::from_str(trimmed).map_err(|e| McpError::Protocol(e.to_string()))
}

#[cfg(test)]
mod tests {
    use super::*;

    #[test]
    fn content_length_encode_roundtrip_shape() {
        let body = r#"{"jsonrpc":"2.0","id":1}"#;
        let enc = encode_content_length(body);
        let s = String::from_utf8_lossy(&enc);
        assert!(s.starts_with("Content-Length: "));
        assert!(s.contains("\r\n\r\n"));
        assert!(s.ends_with(body));
    }
}