Skip to main content

mcp_utils/tool_gateway/
transport.rs

1use rmcp::transport::async_rw::AsyncRwTransport;
2use rmcp::{RoleClient, RoleServer, ServiceExt};
3use std::env::{temp_dir, var_os};
4use std::fs::{Permissions, create_dir_all, remove_dir};
5use std::fs::{remove_file, set_permissions};
6use std::{
7    io,
8    os::unix::fs::PermissionsExt,
9    os::unix::net::UnixListener as StdUnixListener,
10    path::{Path, PathBuf},
11};
12use tokio::io::{ReadHalf, WriteHalf};
13use tokio::net::{UnixListener, UnixStream};
14use tokio::task::JoinHandle;
15use tokio_util::sync::CancellationToken;
16use uuid::Uuid;
17
18#[derive(Debug, thiserror::Error)]
19pub enum UnixSocketTransportError {
20    #[error("failed to create MCP socket directory: {0}")]
21    CreateDirectory(#[source] io::Error),
22    #[error("failed to bind MCP socket: {0}")]
23    Bind(#[source] io::Error),
24    #[error("MCP socket path must be absolute")]
25    NotAbsolute,
26    #[error("MCP socket path is not valid UTF-8")]
27    InvalidPath,
28}
29
30/// A session endpoint allocated by this process, or inherited from the session environment.
31#[derive(Debug)]
32pub struct UnixSocketPath {
33    directory: PathBuf,
34    socket: PathBuf,
35    /// Only the allocating process owns endpoint removal; inherited paths never remove.
36    remove_on_drop: bool,
37}
38
39impl UnixSocketPath {
40    pub fn new() -> Result<Self, UnixSocketTransportError> {
41        let short_id = &Uuid::new_v4().simple().to_string()[..8];
42        let runtime_dir =
43            var_os("XDG_RUNTIME_DIR").filter(|path| Path::new(path).is_absolute()).map_or_else(temp_dir, PathBuf::from);
44        let socket_dir = runtime_dir.join("aether").join(format!("aether-{short_id}"));
45        let socket = socket_dir.join("ipc.sock");
46        create_dir_all(&socket_dir).map_err(UnixSocketTransportError::CreateDirectory)?;
47        set_permissions(&socket_dir, Permissions::from_mode(0o700))
48            .map_err(UnixSocketTransportError::CreateDirectory)?;
49        Ok(Self { directory: socket_dir, socket, remove_on_drop: true })
50    }
51
52    pub fn from_path(path: impl Into<PathBuf>) -> Result<Self, UnixSocketTransportError> {
53        let socket = path.into();
54        if !socket.is_absolute() {
55            return Err(UnixSocketTransportError::NotAbsolute);
56        }
57        let directory = socket.parent().ok_or(UnixSocketTransportError::InvalidPath)?.to_path_buf();
58        Ok(Self { directory, socket, remove_on_drop: false })
59    }
60
61    pub fn path(&self) -> &Path {
62        &self.socket
63    }
64
65    pub fn directory(&self) -> &Path {
66        &self.directory
67    }
68}
69
70impl Drop for UnixSocketPath {
71    fn drop(&mut self) {
72        if self.remove_on_drop {
73            let _ = remove_file(&self.socket);
74            let _ = remove_dir(&self.directory);
75        }
76    }
77}
78
79pub struct UnixSocketMcpTransport {
80    path: UnixSocketPath,
81    listener: UnixListener,
82}
83
84impl UnixSocketMcpTransport {
85    pub fn bind(path: UnixSocketPath) -> Result<Self, UnixSocketTransportError> {
86        let _ = remove_file(path.path());
87        let listener = StdUnixListener::bind(path.path()).map_err(UnixSocketTransportError::Bind)?;
88        listener.set_nonblocking(true).map_err(UnixSocketTransportError::Bind)?;
89        let listener = UnixListener::from_std(listener).map_err(UnixSocketTransportError::Bind)?;
90        Ok(Self { path, listener })
91    }
92
93    pub fn path(&self) -> &Path {
94        self.path.path()
95    }
96
97    /// Serve connections until the returned server is dropped.
98    pub fn spawn<T>(self, server: T) -> UnixSocketServer
99    where
100        T: Clone + ServiceExt<RoleServer> + Send + 'static,
101    {
102        let Self { path, listener } = self;
103        let cancellation = CancellationToken::new();
104        let accept_cancellation = cancellation.clone();
105        let task = tokio::spawn(async move {
106            loop {
107                let accepted = tokio::select! {
108                    () = accept_cancellation.cancelled() => break,
109                    result = listener.accept() => result,
110                };
111                let Ok((stream, _)) = accepted else { break };
112                let server = server.clone();
113                let connection_cancellation = accept_cancellation.clone();
114                tokio::spawn(async move {
115                    match server.serve(stream).await {
116                        Ok(running) => {
117                            let service_cancellation = running.cancellation_token();
118                            let mut waiting = Box::pin(running.waiting());
119                            tokio::select! {
120                                () = connection_cancellation.cancelled() => service_cancellation.cancel(),
121                                _ = &mut waiting => {}
122                            }
123                        }
124                        Err(error) => tracing::debug!(%error, "MCP Unix socket client ended during initialization"),
125                    }
126                });
127            }
128        });
129        UnixSocketServer { path, cancellation, task }
130    }
131}
132
133/// Owns the accept task, connection cancellation, and the session endpoint.
134/// Dropping it cancels in-flight connections and removes the endpoint.
135pub struct UnixSocketServer {
136    path: UnixSocketPath,
137    cancellation: CancellationToken,
138    task: JoinHandle<()>,
139}
140
141impl UnixSocketServer {
142    pub fn path(&self) -> &Path {
143        self.path.path()
144    }
145}
146
147impl Drop for UnixSocketServer {
148    fn drop(&mut self) {
149        self.cancellation.cancel();
150        self.task.abort();
151    }
152}
153
154/// Connect an rmcp client to an inherited session endpoint.
155pub async fn connect(
156    path: impl AsRef<Path>,
157) -> io::Result<AsyncRwTransport<RoleClient, ReadHalf<UnixStream>, WriteHalf<UnixStream>>> {
158    let stream = UnixStream::connect(path).await?;
159    let (read, write) = tokio::io::split(stream);
160    Ok(AsyncRwTransport::new_client(read, write))
161}
162
163#[cfg(test)]
164mod tests {
165    use super::*;
166    use rmcp::ServerHandler;
167    use rmcp::handler::server::router::tool::ToolRouter;
168    use rmcp::model::{ServerCapabilities, ServerConfig};
169    use rmcp::{tool, tool_handler, tool_router};
170
171    #[derive(Clone)]
172    struct TestServer {
173        tool_router: ToolRouter<Self>,
174    }
175
176    #[tool_router(allow_empty)]
177    impl TestServer {}
178
179    #[allow(clippy::unused_async_trait_impl)]
180    #[tool_handler(router = self.tool_router)]
181    impl ServerHandler for TestServer {
182        fn get_info(&self) -> ServerConfig {
183            ServerConfig::new(ServerCapabilities::builder().enable_tools().build())
184        }
185    }
186
187    #[test]
188    fn allocated_endpoint_is_removed_on_drop_without_binding() {
189        let path = UnixSocketPath::new().unwrap();
190        let directory = path.directory().to_path_buf();
191        assert!(directory.exists());
192        drop(path);
193        assert!(!directory.exists());
194    }
195
196    #[test]
197    fn inherited_endpoint_is_not_removed_on_drop() {
198        let path = UnixSocketPath::new().unwrap();
199        let directory = path.directory().to_path_buf();
200        let inherited = UnixSocketPath::from_path(path.path()).unwrap();
201        drop(inherited);
202        assert!(directory.exists());
203        drop(path);
204        assert!(!directory.exists());
205    }
206
207    #[tokio::test]
208    async fn allocated_endpoint_is_private_and_removed_with_transport() {
209        let path = UnixSocketPath::new().unwrap();
210        let directory = path.directory().to_path_buf();
211        let transport = UnixSocketMcpTransport::bind(path).unwrap();
212        assert_eq!(std::fs::metadata(&directory).unwrap().permissions().mode() & 0o777, 0o700);
213        assert!(transport.path().exists());
214        drop(transport);
215        assert!(!directory.exists());
216    }
217
218    #[tokio::test]
219    async fn spawned_endpoint_is_removed_with_server() {
220        let path = UnixSocketPath::new().unwrap();
221        let directory = path.directory().to_path_buf();
222        let transport = UnixSocketMcpTransport::bind(path).unwrap();
223        let socket = transport.path().to_path_buf();
224        let server = transport.spawn(TestServer { tool_router: TestServer::tool_router() });
225        assert!(socket.exists());
226        drop(server);
227        assert!(!directory.exists());
228    }
229
230    #[tokio::test]
231    async fn connected_client_completes_initialization() {
232        let path = UnixSocketPath::new().unwrap();
233        let transport = UnixSocketMcpTransport::bind(path).unwrap();
234        let socket = transport.path().to_path_buf();
235        let _server = transport.spawn(TestServer { tool_router: TestServer::tool_router() });
236        let _client = ().serve(connect(&socket).await.unwrap()).await.unwrap();
237    }
238
239    #[tokio::test]
240    async fn connected_client_can_list_empty_tools() {
241        let path = UnixSocketPath::new().unwrap();
242        let transport = UnixSocketMcpTransport::bind(path).unwrap();
243        let socket = transport.path().to_path_buf();
244        let _server = transport.spawn(TestServer { tool_router: TestServer::tool_router() });
245        let client = ().serve(connect(&socket).await.unwrap()).await.unwrap();
246
247        assert!(client.list_all_tools().await.unwrap().is_empty());
248    }
249
250    #[tokio::test]
251    async fn dropping_server_cancels_in_flight_connections() {
252        use rmcp::model::CallToolRequestParams;
253        use tokio::sync::watch;
254
255        #[derive(Clone)]
256        struct SlowServer {
257            tool_router: ToolRouter<Self>,
258            started: watch::Sender<bool>,
259        }
260
261        #[tool_router]
262        impl SlowServer {
263            #[tool(description = "Blocks until the connection is aborted")]
264            async fn slow(&self) -> String {
265                let _ = self.started.send(true);
266                std::future::pending::<()>().await;
267                "done".to_string()
268            }
269        }
270
271        #[allow(clippy::unused_async_trait_impl)]
272        #[tool_handler(router = self.tool_router)]
273        impl ServerHandler for SlowServer {
274            fn get_info(&self) -> ServerConfig {
275                ServerConfig::new(ServerCapabilities::builder().enable_tools().build())
276            }
277        }
278
279        let path = UnixSocketPath::new().unwrap();
280        let transport = UnixSocketMcpTransport::bind(path).unwrap();
281        let socket = transport.path().to_path_buf();
282        let (started_tx, mut started_rx) = watch::channel(false);
283        let server = transport.spawn(SlowServer { tool_router: SlowServer::tool_router(), started: started_tx });
284
285        let client = ().serve(connect(&socket).await.unwrap()).await.unwrap();
286        let call = tokio::spawn(async move {
287            let _ = client.call_tool_once(CallToolRequestParams::new("slow")).await;
288        });
289        started_rx.changed().await.unwrap();
290
291        drop(server);
292        call.await.unwrap();
293    }
294}