magi-code 0.96.2

Repository-aware CLI coding agent for terminal work
Documentation
use crate::{
    cancellation::AgentCancellation,
    config::{McpServerConfig, McpStdioServerConfig},
    mcp::{
        McpError, McpResult,
        http::HttpConnection,
        jsonrpc::RequestId,
        protocol::{
            CallToolParams, CallToolResult, InitializeRequestParams, InitializeResult,
            ListToolsParams, ListToolsResult, METHOD_INITIALIZE, METHOD_INITIALIZED,
            METHOD_TOOLS_CALL, METHOD_TOOLS_LIST, PROTOCOL_VERSION, Tool,
        },
        stdio::StdioConnection,
    },
};
use serde::de::DeserializeOwned;
use serde_json::Value;
use std::{
    collections::HashSet,
    path::Path,
    sync::{Arc, Mutex, atomic::AtomicU64},
};

const MAX_TOOLS_LIST_PAGES: usize = 100;

pub(crate) struct McpClient {
    connection: McpConnection,
    next_id: AtomicU64,
    server_info: Arc<Mutex<Option<crate::mcp::protocol::Implementation>>>,
}

impl std::fmt::Debug for McpClient {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        f.debug_struct("McpClient")
            .field("connection", &self.connection)
            .finish_non_exhaustive()
    }
}

impl std::fmt::Debug for McpConnection {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        match self {
            Self::Stdio(_) => f.debug_tuple("Stdio").field(&"<stdio connection>").finish(),
            Self::Http(connection) => f.debug_tuple("Http").field(connection).finish(),
        }
    }
}

enum McpConnection {
    Stdio(StdioConnection),
    Http(HttpConnection),
}

impl McpConnection {
    fn send_request(
        &self,
        id: RequestId,
        method: &str,
        params: Option<Value>,
        cancellation: Option<&AgentCancellation>,
    ) -> McpResult<Value> {
        match self {
            Self::Stdio(connection) => connection.send_request(id, method, params, cancellation),
            Self::Http(connection) => connection.send_request(id, method, params, cancellation),
        }
    }

    fn send_notification(&self, method: &str, params: Option<Value>) -> McpResult<()> {
        match self {
            Self::Stdio(connection) => connection.send_notification(method, params),
            Self::Http(connection) => connection.send_notification(method, params),
        }
    }

    fn request_shutdown(&self) {
        match self {
            Self::Stdio(connection) => connection.request_shutdown(),
            Self::Http(connection) => connection.request_shutdown(),
        }
    }

    fn shutdown(&self) {
        match self {
            Self::Stdio(connection) => connection.shutdown(),
            Self::Http(connection) => connection.shutdown(),
        }
    }
}

impl McpClient {
    pub(crate) fn connect_stdio(config: &McpStdioServerConfig) -> McpResult<Self> {
        Ok(Self {
            connection: McpConnection::Stdio(StdioConnection::connect(config)?),
            next_id: AtomicU64::new(1),
            server_info: Arc::new(Mutex::new(None)),
        })
    }

    pub(crate) fn connect_named(
        server_name: Option<&str>,
        config: &McpServerConfig,
        mc_home: Option<&Path>,
    ) -> McpResult<Self> {
        if let Some(server_name) = server_name {
            crate::config::validate_mcp_server_name(server_name)
                .map_err(|error| McpError::Config(error.to_string()))?;
        }
        match config {
            McpServerConfig::Stdio(config) => Self::connect_stdio(config),
            McpServerConfig::Http(config) => Ok(Self {
                connection: McpConnection::Http(HttpConnection::connect_named(
                    server_name,
                    config,
                    mc_home,
                )?),
                next_id: AtomicU64::new(1),
                server_info: Arc::new(Mutex::new(None)),
            }),
        }
    }

    fn next_id(&self) -> RequestId {
        let id = self
            .next_id
            .fetch_add(1, std::sync::atomic::Ordering::SeqCst);
        RequestId::Number(i64::try_from(id).unwrap_or(i64::MAX))
    }

    fn send_request_inner(
        &self,
        method: &str,
        params: Option<Value>,
        cancellation: Option<&AgentCancellation>,
    ) -> McpResult<Value> {
        let id = self.next_id();
        self.connection
            .send_request(id, method, params, cancellation)
    }

    pub(crate) fn send_request(&self, method: &str, params: Option<Value>) -> McpResult<Value> {
        self.send_request_inner(method, params, None)
    }

    pub(crate) fn send_request_cancellable(
        &self,
        method: &str,
        params: Option<Value>,
        cancellation: &AgentCancellation,
    ) -> McpResult<Value> {
        self.send_request_inner(method, params, Some(cancellation))
    }

    pub(crate) fn send_notification(&self, method: &str, params: Option<Value>) -> McpResult<()> {
        self.connection.send_notification(method, params)
    }

    pub(crate) fn initialize(&self) -> McpResult<InitializeResult> {
        self.initialize_inner(None)
    }

    pub(crate) fn initialize_cancellable(
        &self,
        cancellation: &AgentCancellation,
    ) -> McpResult<InitializeResult> {
        self.initialize_inner(Some(cancellation))
    }

    fn initialize_inner(
        &self,
        cancellation: Option<&AgentCancellation>,
    ) -> McpResult<InitializeResult> {
        if let Some(cancellation) = cancellation {
            cancellation.check().map_err(McpError::transport)?;
        }
        let params = serde_json::to_value(InitializeRequestParams::default())
            .map_err(McpError::transport)?;
        let response = match cancellation {
            Some(cancellation) => {
                self.send_request_cancellable(METHOD_INITIALIZE, Some(params), cancellation)
            }
            None => self.send_request(METHOD_INITIALIZE, Some(params)),
        }?;
        if let Some(cancellation) = cancellation {
            cancellation.check().map_err(McpError::transport)?;
        }
        let result: InitializeResult = decode(response)?;
        if result.protocol_version != PROTOCOL_VERSION {
            return Err(McpError::Protocol {
                code: -32602,
                message: format!(
                    "unsupported MCP protocol version {}; expected {PROTOCOL_VERSION}",
                    result.protocol_version
                ),
            });
        }
        if let Ok(mut server_info) = self.server_info.lock() {
            *server_info = Some(result.server_info.clone());
        }
        if let Some(cancellation) = cancellation {
            cancellation.check().map_err(McpError::transport)?;
        }
        self.send_notification(METHOD_INITIALIZED, None)?;
        if let Some(cancellation) = cancellation {
            cancellation.check().map_err(McpError::transport)?;
        }
        Ok(result)
    }

    pub(crate) fn list_tools(&self) -> McpResult<Vec<Tool>> {
        self.list_tools_inner(None)
    }

    pub(crate) fn list_tools_cancellable(
        &self,
        cancellation: &AgentCancellation,
    ) -> McpResult<Vec<Tool>> {
        self.list_tools_inner(Some(cancellation))
    }

    fn list_tools_inner(&self, cancellation: Option<&AgentCancellation>) -> McpResult<Vec<Tool>> {
        let mut tools = Vec::new();
        let mut cursor = None;
        let mut seen_cursors = HashSet::new();
        for page in 0..MAX_TOOLS_LIST_PAGES {
            if let Some(cancellation) = cancellation {
                cancellation.check().map_err(McpError::transport)?;
            }
            let params = serde_json::to_value(ListToolsParams {
                cursor: cursor.clone(),
            })
            .map_err(McpError::transport)?;
            let response = match cancellation {
                Some(cancellation) => {
                    self.send_request_cancellable(METHOD_TOOLS_LIST, Some(params), cancellation)
                }
                None => self.send_request(METHOD_TOOLS_LIST, Some(params)),
            }?;
            if let Some(cancellation) = cancellation {
                cancellation.check().map_err(McpError::transport)?;
            }
            let result: ListToolsResult = decode(response)?;
            tools.extend(result.tools);
            if let Some(cancellation) = cancellation {
                cancellation.check().map_err(McpError::transport)?;
            }
            cursor = result.next_cursor;
            if let Some(next_cursor) = cursor.as_ref() {
                if !seen_cursors.insert(next_cursor.clone()) {
                    return Err(McpError::Protocol {
                        code: -32603,
                        message: "MCP tools/list returned repeated pagination cursor".to_string(),
                    });
                }
            } else {
                return Ok(tools);
            }
            if page + 1 == MAX_TOOLS_LIST_PAGES {
                return Err(McpError::Protocol {
                    code: -32603,
                    message: format!("MCP tools/list exceeded {MAX_TOOLS_LIST_PAGES} pages"),
                });
            }
        }
        unreachable!("tools/list pagination loop returns from inside bounded range")
    }

    pub(crate) fn call_tool_cancellable(
        &self,
        name: &str,
        arguments: Option<Value>,
        cancellation: &AgentCancellation,
    ) -> McpResult<CallToolResult> {
        let params = serde_json::to_value(CallToolParams {
            name: name.to_string(),
            arguments,
        })
        .map_err(McpError::transport)?;
        decode(self.send_request_cancellable(METHOD_TOOLS_CALL, Some(params), cancellation)?)
    }

    pub(crate) fn request_shutdown(&self) {
        self.connection.request_shutdown();
    }

    pub(crate) fn shutdown(&mut self) {
        self.shutdown_shared();
    }

    // Request handles share the client, but only transport teardown takes lifecycle locks.
    // Never wait for active requests before closing: closing must release those requests.
    pub(crate) fn shutdown_shared(&self) {
        self.connection.shutdown();
    }
}

impl Drop for McpClient {
    fn drop(&mut self) {
        self.shutdown();
    }
}

fn decode<T: DeserializeOwned>(value: Value) -> McpResult<T> {
    serde_json::from_value(value).map_err(McpError::transport)
}