use std::sync::Arc;
use std::time::Duration;
use anyhow::{Context, Result};
use mcpmesh_net::framing::{FrameReader, Inbound, write_frame};
use mcpmesh_net::{SessionTransport, connect};
use serde_json::Value;
use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt};
use super::MeshState;
use crate::allowlist::PeerEntry;
pub async fn dial_service(
mesh: &Arc<MeshState>,
peer: &str,
service: &str,
) -> Result<SessionTransport> {
if let Some(view) = mesh.roster.view() {
let devices = view.devices_for_user(peer);
if !devices.is_empty() {
let candidates = order_dial_candidates(&devices, &mesh.presence_table, peer);
return race_dial(&mesh.endpoint, candidates, service)
.await
.with_context(|| format!("dial {peer}/{service}"));
}
}
let peer_owned = peer.to_string();
let store = mesh.store.clone();
let (single, multi): (Option<PeerEntry>, Vec<[u8; 32]>) =
tokio::task::spawn_blocking(move || -> Result<_> {
if let Some(e) = store.entry_for(&peer_owned)? {
return Ok((Some(e), Vec::new()));
}
let mut by_user = store.entries_for_user(&peer_owned)?;
match by_user.len() {
0 => Ok((None, Vec::new())),
1 => Ok((by_user.pop(), Vec::new())),
_ => Ok((None, by_user.iter().map(|e| e.endpoint_id).collect())),
}
})
.await
.context("join peer resolve")??;
if !multi.is_empty() {
return race_dial(&mesh.endpoint, multi, service)
.await
.with_context(|| format!("dial {peer}/{service}"));
}
let entry = single.with_context(|| format!("peer '{peer}' is not in the allowlist"))?;
let endpoint_id = iroh::EndpointId::from_bytes(&entry.endpoint_id)
.map_err(|e| anyhow::anyhow!("stored endpoint id for '{peer}' is invalid: {e}"))?;
let addr = stored_dial_addr(entry.last_addr.as_deref(), endpoint_id);
connect_with_timeout(&mesh.endpoint, addr, service, DIAL_TIMEOUT)
.await
.with_context(|| format!("dial {peer}/{service}"))
}
pub(crate) fn stored_dial_addr(
last_addr: Option<&str>,
endpoint_id: iroh::EndpointId,
) -> iroh::EndpointAddr {
if let Some(json) = last_addr
&& let Ok(addr) = serde_json::from_str::<iroh::EndpointAddr>(json)
&& addr.id == endpoint_id
{
return addr;
}
iroh::EndpointAddr::from(endpoint_id)
}
const DIAL_STAGGER: Duration = Duration::from_millis(500);
const DIAL_TIMEOUT: Duration = Duration::from_secs(20);
pub(crate) async fn connect_with_timeout(
endpoint: &iroh::Endpoint,
addr: iroh::EndpointAddr,
service: &str,
timeout: Duration,
) -> Result<SessionTransport> {
match tokio::time::timeout(timeout, connect(endpoint, addr, service)).await {
Ok(r) => r.map_err(Into::into),
Err(_) => anyhow::bail!("dial timed out after {timeout:?}"),
}
}
fn order_dial_candidates(
devices: &[([u8; 32], String)],
presence: &crate::roster::presence::PresenceTable,
user_id: &str,
) -> Vec<[u8; 32]> {
let by_recency = presence.endpoints_for_user_by_recency(user_id);
let recency_rank = |eid: &[u8; 32]| -> usize {
by_recency
.iter()
.position(|e| e == eid)
.unwrap_or(usize::MAX)
};
let mut ordered: Vec<([u8; 32], String)> = devices.to_vec();
ordered.sort_by_key(|(eid, role)| (dial_role_rank(role), recency_rank(eid)));
ordered.into_iter().map(|(eid, _)| eid).collect()
}
pub(crate) fn dial_role_rank(role: &str) -> u8 {
match role {
"primary" => 0,
"mirror" => 1,
_ => 2,
}
}
pub async fn race_dial(
endpoint: &iroh::Endpoint,
candidates: Vec<[u8; 32]>,
service: &str,
) -> Result<SessionTransport> {
anyhow::ensure!(!candidates.is_empty(), "no dial candidates to race");
let mut set: tokio::task::JoinSet<Result<SessionTransport>> = tokio::task::JoinSet::new();
let spawn_dial = |set: &mut tokio::task::JoinSet<Result<SessionTransport>>, eid: [u8; 32]| {
let ep = endpoint.clone();
let svc = service.to_string();
set.spawn(async move { dial_one(&ep, eid, &svc).await });
};
let mut next = 0usize; spawn_dial(&mut set, candidates[next]); next += 1;
let mut last_err: Option<anyhow::Error> = None;
loop {
if next < candidates.len() {
tokio::select! {
biased;
joined = set.join_next() => match joined {
Some(Ok(Ok(t))) => return Ok(t), Some(Ok(Err(e))) => last_err = Some(e), Some(Err(e)) => last_err = Some(anyhow::anyhow!("dial task join error: {e}")),
None => {
spawn_dial(&mut set, candidates[next]);
next += 1;
}
},
() = tokio::time::sleep(DIAL_STAGGER) => {
spawn_dial(&mut set, candidates[next]);
next += 1;
}
}
} else {
match set.join_next().await {
Some(Ok(Ok(t))) => return Ok(t),
Some(Ok(Err(e))) => last_err = Some(e),
Some(Err(e)) => last_err = Some(anyhow::anyhow!("dial task join error: {e}")),
None => {
return Err(
last_err.unwrap_or_else(|| anyhow::anyhow!("all dial candidates failed"))
);
}
}
}
}
}
async fn dial_one(
endpoint: &iroh::Endpoint,
eid: [u8; 32],
service: &str,
) -> Result<SessionTransport> {
let endpoint_id = iroh::EndpointId::from_bytes(&eid)
.map_err(|e| anyhow::anyhow!("roster device endpoint id is invalid: {e}"))?;
let addr = iroh::EndpointAddr::from(endpoint_id);
connect_with_timeout(endpoint, addr, service, DIAL_TIMEOUT).await
}
pub async fn pipe_session<CR, CW>(
mut transport: SessionTransport,
service: &str,
mut control_reader: FrameReader<CR>,
mut control_writer: CW,
) -> Result<()>
where
CR: AsyncRead + Unpin + Send,
CW: AsyncWrite + Unpin + Send,
{
let init = match control_reader.next().await? {
Some(Inbound::Frame(v)) => inject_service(v, service),
Some(Inbound::Violation(_)) | None => return Ok(()),
};
transport
.send_value(init)
.await
.context("forward initialize to peer")?;
let transport_writer = transport.writer();
let to_mesh = async {
loop {
match control_reader.next().await {
Ok(Some(Inbound::Frame(frame))) => {
if transport_writer.send_value(frame).await.is_err() {
break; }
}
Ok(Some(Inbound::Violation(_))) => break,
Ok(None) | Err(_) => break, }
}
let _ = transport_writer.shutdown().await;
std::future::pending::<()>().await
};
let to_control = async {
while let Ok(Some(frame)) = transport.recv_value().await {
if write_frame(&mut control_writer, &frame).await.is_err() {
break; }
}
};
tokio::select! {
() = to_mesh => {}
() = to_control => {}
}
let _ = transport.shutdown().await;
let _ = control_writer.shutdown().await;
Ok(())
}
fn inject_service(mut frame: Value, service: &str) -> Value {
let Some(obj) = frame.as_object_mut() else {
return frame;
};
let params = obj
.entry("params")
.or_insert_with(|| Value::Object(Default::default()));
if !params.is_object() {
*params = Value::Object(Default::default());
}
let params = params.as_object_mut().expect("params set to object above");
let meta = params
.entry("_meta")
.or_insert_with(|| Value::Object(Default::default()));
if !meta.is_object() {
*meta = Value::Object(Default::default()); }
meta.as_object_mut()
.expect("meta set to object above")
.insert("mcpmesh/service".into(), Value::String(service.to_string()));
frame
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn inject_service_sets_meta_across_shapes() {
use serde_json::json;
let f = inject_service(json!({"method": "initialize"}), "kb");
assert_eq!(f["params"]["_meta"]["mcpmesh/service"], "kb");
assert_eq!(f["method"], "initialize");
let f = inject_service(json!({"params": {"x": 1}}), "loc");
assert_eq!(f["params"]["x"], 1);
assert_eq!(f["params"]["_meta"]["mcpmesh/service"], "loc");
let f = inject_service(json!({"params": 7}), "kb");
assert_eq!(f["params"]["_meta"]["mcpmesh/service"], "kb");
let f = inject_service(json!({"params": {"_meta": "nope"}}), "kb");
assert_eq!(f["params"]["_meta"]["mcpmesh/service"], "kb");
assert_eq!(inject_service(json!("scalar"), "kb"), json!("scalar"));
}
#[tokio::test(flavor = "multi_thread")]
async fn pipe_session_delivers_the_echo_after_control_eof() {
use mcpmesh_net::framing::{FrameReader, Inbound, write_frame};
use serde_json::json;
use tokio::io::duplex;
tokio::time::timeout(std::time::Duration::from_secs(20), async {
let server_ep = iroh::Endpoint::builder(iroh::endpoint::presets::Minimal)
.relay_mode(iroh::RelayMode::Disabled)
.alpns(vec![mcpmesh_net::ALPN_MCP.to_vec()])
.bind()
.await
.unwrap();
let server_addr = server_ep.addr();
let (done_tx, done_rx) = tokio::sync::oneshot::channel::<()>();
let peer = tokio::spawn(async move {
let incoming = server_ep.accept().await.expect("one inbound connection");
let conn = incoming.await.expect("handshake");
let (send, recv) = conn.accept_bi().await.expect("session bi-stream");
let mut t = mcpmesh_net::SessionTransport::new(recv, send, 1024 * 1024);
let mut seen = Vec::new();
while let Ok(Some(f)) = t.recv_value().await {
seen.push(f);
}
for f in &seen {
t.send_value(f.clone()).await.unwrap();
}
t.shutdown().await.unwrap(); let _ = done_rx.await; seen
});
let client_ep = iroh::Endpoint::builder(iroh::endpoint::presets::Minimal)
.relay_mode(iroh::RelayMode::Disabled)
.alpns(vec![mcpmesh_net::ALPN_MCP.to_vec()])
.bind()
.await
.unwrap();
let transport = connect(&client_ep, server_addr, "echo").await.unwrap();
let (mut ctl_in_w, ctl_in_r) = duplex(64 * 1024);
let (ctl_out_w, ctl_out_test_r) = duplex(64 * 1024);
let init = json!({"jsonrpc": "2.0", "id": 1, "method": "initialize", "params": {}});
write_frame(&mut ctl_in_w, &init).await.unwrap();
drop(ctl_in_w);
let session = tokio::spawn(pipe_session(
transport,
"echo",
FrameReader::new(ctl_in_r, 1024 * 1024),
ctl_out_w,
));
let mut ctl_out = FrameReader::new(ctl_out_test_r, 1024 * 1024);
match ctl_out.next().await.unwrap() {
Some(Inbound::Frame(f)) => {
assert_eq!(f["id"], 1, "the echoed initialize answers our id: {f}");
assert_eq!(
f["params"]["_meta"]["mcpmesh/service"], "echo",
"the peer saw the service-injected initialize (the one enumerated \
edit), echoed verbatim: {f}"
);
}
other => panic!("the echo must reach the control side, got {other:?}"),
}
assert!(
ctl_out.next().await.unwrap().is_none(),
"the peer closing ends the session cleanly (control-side EOF)"
);
session.await.unwrap().expect("pipe_session returns Ok");
let _ = done_tx.send(()); assert_eq!(
peer.await.unwrap(),
vec![inject_service(init, "echo")],
"the peer received exactly the injected initialize before the half-close"
);
})
.await
.expect("pipe_session drain test timed out");
}
#[test]
fn stored_dial_addr_attaches_validates_and_degrades() {
let id = iroh::SecretKey::from_bytes(&[7u8; 32]).public();
let other = iroh::SecretKey::from_bytes(&[8u8; 32]).public();
let sock: std::net::SocketAddr = "127.0.0.1:4444".parse().unwrap();
let stored = iroh::EndpointAddr::from_parts(id, [iroh::TransportAddr::Ip(sock)]);
let stored_json = serde_json::to_string(&stored).unwrap();
let addr = stored_dial_addr(Some(&stored_json), id);
assert_eq!(addr, stored, "a matching-id hint is dialed as stored");
assert_eq!(stored_dial_addr(None, id), iroh::EndpointAddr::from(id));
assert_eq!(
stored_dial_addr(Some("not json"), id),
iroh::EndpointAddr::from(id)
);
let mismatched = serde_json::to_string(&iroh::EndpointAddr::from_parts(
other,
[iroh::TransportAddr::Ip(sock)],
))
.unwrap();
assert_eq!(
stored_dial_addr(Some(&mismatched), id),
iroh::EndpointAddr::from(id)
);
}
#[tokio::test]
async fn connect_with_timeout_fails_fast_on_an_unreachable_peer() {
let ep = iroh::Endpoint::builder(iroh::endpoint::presets::Minimal)
.relay_mode(iroh::RelayMode::Disabled)
.alpns(vec![mcpmesh_net::ALPN_MCP.to_vec()])
.bind()
.await
.unwrap();
let dead = iroh::EndpointAddr::from(iroh::EndpointId::from_bytes(&[3u8; 32]).unwrap());
let start = std::time::Instant::now();
let r =
super::connect_with_timeout(&ep, dead, "svc", std::time::Duration::from_millis(300))
.await;
assert!(r.is_err(), "an unreachable dial times out to Err");
assert!(
start.elapsed() < std::time::Duration::from_secs(3),
"the explicit timeout fired fast"
);
}
}