sac-cli 0.1.0

Terminal-based AI coding agent — fork of NAC with extended backend support and context management
Documentation
use super::*;

#[derive(Clone)]
pub struct McpRegistry {
    tools: Arc<HashMap<String, Arc<McpToolBinding>>>,
}

#[derive(Clone)]
struct McpToolBinding {
    tool_name: String,
    definition: ToolDefinition,
    server: Arc<McpServer>,
}

struct McpServer {
    _service: Arc<McpService>,
    tool_call_timeout: Duration,
}

#[derive(Clone)]
pub(super) struct NacMcpClientHandler {
    root_uri: String,
    root_name: String,
}

impl McpRegistry {
    pub async fn load(cwd: &Path, sandbox: Option<&SandboxSession>) -> Result<Option<Arc<Self>>> {
        let Some(path) = default_config_path() else {
            return Ok(None);
        };
        if !path.exists() {
            return Ok(None);
        }

        let raw = match std::fs::read_to_string(&path) {
            Ok(raw) => raw,
            Err(error) => {
                eprintln!(
                    "MCP config at '{}' could not be read; MCP will be disabled: {:#}",
                    path.display(),
                    error
                );
                return Ok(None);
            }
        };
        let config: McpConfigFile = match toml::from_str(&raw) {
            Ok(config) => config,
            Err(error) => {
                eprintln!(
                    "MCP config at '{}' is invalid; MCP will be disabled: {:#}",
                    path.display(),
                    error
                );
                return Ok(None);
            }
        };

        let root_uri = if sandbox.is_some() {
            "file:///workspace".to_string()
        } else {
            Url::from_directory_path(cwd)
                .map_err(|_| anyhow!("failed to build file:// root for {}", cwd.display()))?
                .to_string()
        };
        let root_name = if sandbox.is_some() {
            "workspace".to_string()
        } else {
            cwd.file_name()
                .and_then(|value| value.to_str())
                .unwrap_or("workspace")
                .to_string()
        };

        let handler = NacMcpClientHandler {
            root_uri,
            root_name,
        };

        let mut tools = HashMap::new();
        let mut seen_names = HashMap::<String, usize>::new();

        for (server_name, server_config) in config.mcp_servers {
            if !server_config.enabled {
                continue;
            }

            let service = match timeout(
                MCP_CONNECT_TIMEOUT,
                connect_server(&server_name, &server_config, &handler, sandbox),
            )
            .await
            {
                Ok(Ok(service)) => Arc::new(service),
                Ok(Err(error)) => {
                    eprintln!(
                        "MCP server '{}' is unavailable and will be skipped: {:#}",
                        server_name, error
                    );
                    continue;
                }
                Err(_) => {
                    eprintln!(
                        "MCP server '{}' timed out during connect after {}s and will be skipped",
                        server_name,
                        MCP_CONNECT_TIMEOUT.as_secs()
                    );
                    continue;
                }
            };

            let listed_tools = match timeout(MCP_TOOL_INVENTORY_TIMEOUT, service.list_all_tools())
                .await
            {
                Ok(Ok(tools)) => tools,
                Ok(Err(error)) => {
                    eprintln!(
                        "MCP server '{}' could not list tools and will be skipped: {:#}",
                        server_name, error
                    );
                    continue;
                }
                Err(_) => {
                    eprintln!(
                        "MCP server '{}' timed out while listing tools after {}s and will be skipped",
                        server_name,
                        MCP_TOOL_INVENTORY_TIMEOUT.as_secs()
                    );
                    continue;
                }
            };

            let tool_call_timeout = server_config
                .tool_call_timeout_secs
                .map(Duration::from_secs)
                .unwrap_or(MCP_TOOL_CALL_TIMEOUT);
            let server = Arc::new(McpServer {
                _service: service.clone(),
                tool_call_timeout,
            });
            for tool in listed_tools {
                let qualified_name = allocate_tool_name(&server_name, &tool.name, &mut seen_names);
                let definition = tool_definition(&qualified_name, &server_name, &tool);
                tools.insert(
                    qualified_name,
                    Arc::new(McpToolBinding {
                        tool_name: tool.name.to_string(),
                        definition,
                        server: server.clone(),
                    }),
                );
            }
        }

        if tools.is_empty() {
            return Ok(None);
        }

        Ok(Some(Arc::new(Self {
            tools: Arc::new(tools),
        })))
    }

    pub fn tool_definitions(&self) -> Vec<ToolDefinition> {
        let mut definitions: Vec<ToolDefinition> = self
            .tools
            .values()
            .map(|binding| binding.definition.clone())
            .collect();
        definitions.sort_by(|left, right| left.function.name.cmp(&right.function.name));
        definitions
    }

    pub async fn call_tool(&self, name: &str, args: Value) -> ToolResult {
        let Some(binding) = self.tools.get(name) else {
            return ToolResult {
                content: format!("Error: unknown MCP tool '{}'", name),
                is_error: true,
            };
        };

        let arguments = match args {
            Value::Object(map) => Some(map),
            Value::Null => None,
            _ => {
                return ToolResult {
                    content: format!("Error: MCP tool '{}' requires object arguments", name),
                    is_error: true,
                }
            }
        };

        let mut params = CallToolRequestParams::new(binding.tool_name.clone());
        if let Some(arguments) = arguments {
            params = params.with_arguments(arguments);
        }
        let tool_timeout = binding.server.tool_call_timeout;
        match timeout(
            tool_timeout,
            binding.server._service.call_tool(params),
        )
        .await
        {
            Ok(Ok(result)) => flatten_tool_result(result),
            Ok(Err(error)) => ToolResult {
                content: format!("Error calling MCP tool '{}': {}", name, error),
                is_error: true,
            },
            Err(_) => ToolResult {
                content: format!(
                    "Error calling MCP tool '{}': timed out after {}s",
                    name,
                    tool_timeout.as_secs()
                ),
                is_error: true,
            },
        }
    }

    pub fn has_tool(&self, name: &str) -> bool {
        self.tools.contains_key(name)
    }
}

impl ClientHandler for NacMcpClientHandler {
    fn get_info(&self) -> ClientInfo {
        ClientInfo::new(
            serde_json::from_value(serde_json::json!({
                "roots": {
                    "listChanged": true
                }
            }))
            .expect("valid MCP client capabilities"),
            Implementation::new("sac", env!("CARGO_PKG_VERSION")),
        )
    }

    async fn list_roots(
        &self,
        _request_context: rmcp::service::RequestContext<RoleClient>,
    ) -> std::result::Result<ListRootsResult, rmcp::model::ErrorData> {
        Ok(ListRootsResult::new(vec![
            Root::new(self.root_uri.clone()).with_name(self.root_name.clone())
        ]))
    }
}

pub(super) fn tool_definition(full_name: &str, server_name: &str, tool: &Tool) -> ToolDefinition {
    let description = tool
        .description
        .as_ref()
        .map(|value| value.to_string())
        .unwrap_or_else(|| format!("MCP tool '{}' from server '{}'", tool.name, server_name));
    ToolDefinition {
        def_type: "function".to_string(),
        function: FunctionDef {
            name: full_name.to_string(),
            description,
            parameters: tool.schema_as_json_value(),
        },
    }
}