mcp_utils/tool_gateway/
transport.rs1use 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#[derive(Debug)]
32pub struct UnixSocketPath {
33 directory: PathBuf,
34 socket: PathBuf,
35 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 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
133pub 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
154pub 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}