use anyhow::{Context, Result};
use mcpmesh_local_api::transport::{LocalListener, LocalStream};
use std::path::Path;
pub const MAX_FRAME_BYTES: usize = 16 * 1024 * 1024;
#[cfg(unix)]
pub fn ensure_runtime_dir(dir: &Path) -> Result<()> {
mcpmesh_local_api::service::ensure_private_dir(dir)
.with_context(|| format!("secure runtime dir {}", dir.display()))
}
pub async fn bind_control_socket(path: &Path) -> Result<LocalListener> {
mcpmesh_local_api::transport::bind_local(path)
.with_context(|| format!("bind control socket {}", path.display()))
}
pub fn check_peer(stream: &LocalStream) -> Result<()> {
anyhow::ensure!(
mcpmesh_local_api::transport::authorize_local_peer(stream),
"refusing local connection: peer not authorized (same-user gate refused the connection)"
);
Ok(())
}
#[cfg(all(test, unix))]
mod tests {
use super::*;
use mcpmesh_local_api::{API_NAME, API_VERSION, Hello};
use mcpmesh_net::framing::{FrameReader, Inbound, write_frame};
use std::time::Duration;
use tokio::io::BufReader;
#[tokio::test]
async fn hello_frame_roundtrips_over_real_uds_and_peer_uid_passes() {
tokio::time::timeout(Duration::from_secs(10), async {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("mcpmesh").join("mcpmesh.sock");
let mut listener = bind_control_socket(&path).await.unwrap();
{
use std::os::unix::fs::PermissionsExt;
let dir_mode = std::fs::metadata(path.parent().unwrap())
.unwrap()
.permissions()
.mode();
assert_eq!(dir_mode & 0o777, 0o700, "runtime dir must be 0700");
let sock_mode = std::fs::metadata(&path).unwrap().permissions().mode();
assert_eq!(sock_mode & 0o777, 0o600, "control socket must be 0600");
}
let expected = Hello {
api: API_NAME.into(),
api_version: API_VERSION.into(),
api_minor: 0,
stack_version: "0.1.0".into(),
};
let server_hello = expected.clone();
let accept = tokio::spawn(async move {
let stream = listener.accept().await.unwrap();
check_peer(&stream).unwrap();
let (_r, mut w) = stream.into_split();
let frame = serde_json::to_value(&server_hello).unwrap();
write_frame(&mut w, &frame).await.unwrap();
});
let client = LocalStream::connect(&path).await.unwrap();
let (r, _w) = client.into_split();
let mut reader = FrameReader::new(BufReader::new(r), 16 * 1024 * 1024);
let got: Hello = match reader.next().await.unwrap().unwrap() {
Inbound::Frame(v) => serde_json::from_value(v).unwrap(),
other => panic!("expected a hello frame, got {other:?}"),
};
assert_eq!(got.api, "mcpmesh-local/1");
assert_eq!(got.stack_version, "0.1.0");
assert_eq!(got, expected);
accept.await.unwrap();
})
.await
.expect("hello round-trip timed out");
}
}