use std::sync::Arc;
use mcpmesh_net::framing::write_frame;
use mcpmesh_net::{ALPN_MCP, ALPN_PAIR, ALPN_PING, Services, run_mesh_connection};
use tokio::task::JoinHandle;
use crate::pairing;
use super::{MeshState, STACK_VERSION};
fn gate_and_register(
mesh: &Arc<MeshState>,
conn: &iroh::endpoint::Connection,
blob_conn_limit: bool,
) -> Option<mcpmesh_net::registry::Registration> {
let remote = mcpmesh_net::EndpointId::from(conn.remote_id());
if mesh.gate.resolve(&remote).is_none() {
conn.close(mcpmesh_net::CLOSE_UNAUTHORIZED.into(), b"unauthorized");
return None;
}
if blob_conn_limit && !mesh.limits().admit_blob_conn(&remote) {
conn.close(0u32.into(), b"blob rate limited");
return None;
}
let roster_user = mesh.gate.roster_user(&remote);
let registration = mesh
.conn_registry
.register_checked(conn, roster_user.clone(), |eid| {
mesh.gate.should_sever_now(eid, roster_user.as_deref())
});
if registration.is_none() {
conn.close(mcpmesh_net::CLOSE_UNAUTHORIZED.into(), b"unauthorized");
}
registration
}
pub fn spawn_accept_loop(mesh: Arc<MeshState>, services: Arc<Services>) -> JoinHandle<()> {
tokio::spawn(async move {
while let Some(incoming) = mesh.endpoint.accept().await {
let (mesh, services) = (mesh.clone(), services.clone());
tokio::spawn(async move {
let conn = match incoming.await {
Ok(conn) => conn,
Err(e) => {
tracing::debug!(%e, "inbound handshake failed");
return;
}
};
let alpn = conn.alpn().to_vec();
match alpn.as_slice() {
a if a == ALPN_MCP => {
run_mesh_connection(
conn,
mesh.gate.clone(),
services,
mesh.conn_registry.clone(),
)
.await;
}
a if a == ALPN_PAIR => {
if mesh.invites.count() == 0 {
conn.close(0u32.into(), b"no pairing in progress");
return;
}
if !mesh.limits().admit_pair_accept() {
conn.close(0u32.into(), b"pair rate limited");
return;
}
if let Err(e) =
pairing::rendezvous::handle_inviter_side(conn, mesh.inviter_ctx()).await
{
tracing::debug!(%e, "pair rendezvous error");
}
}
a if a == ALPN_PING => {
let remote = mcpmesh_net::EndpointId::from(conn.remote_id());
if mesh.gate.resolve(&remote).is_none() {
conn.close(mcpmesh_net::CLOSE_UNAUTHORIZED.into(), b"unauthorized");
return;
}
if let Ok((mut send, _recv)) = conn.accept_bi().await {
let meta = mesh.app_metadata();
let pong = if meta.is_empty() {
serde_json::json!({ "stack_version": STACK_VERSION })
} else {
serde_json::json!({ "stack_version": STACK_VERSION, "meta": meta })
};
if write_frame(&mut send, &pong).await.is_ok() {
let _ = send.finish();
let _ = send.stopped().await;
}
}
}
a if a == crate::roster::transport::GOSSIP_ALPN => {
let Some(gossip) = mesh.gossip.clone() else {
conn.close(0u32.into(), b"gossip not enabled");
return;
};
let Some(_registration) = gate_and_register(&mesh, &conn, false) else {
return;
};
if let Err(e) = iroh::protocol::ProtocolHandler::accept(&gossip, conn).await
{
tracing::debug!(%e, "gossip accept error");
}
}
a if a == crate::roster::transport::BLOB_ALPN => {
let Some(blobs) = mesh.blobs.clone() else {
conn.close(0u32.into(), b"blobs not enabled");
return;
};
let Some(_registration) = gate_and_register(&mesh, &conn, false) else {
return;
};
let blob_proto = blobs.protocol();
if let Err(e) =
iroh::protocol::ProtocolHandler::accept(&blob_proto, conn).await
{
tracing::debug!(%e, "blob accept error");
}
}
a if a == crate::blobs::APP_BLOB_ALPN => {
let Some(app_blobs) = mesh.app_blobs().await else {
conn.close(0u32.into(), b"app blobs not enabled");
return;
};
let Some(_registration) = gate_and_register(&mesh, &conn, true) else {
return;
};
let blob_proto = app_blobs.protocol();
if let Err(e) =
iroh::protocol::ProtocolHandler::accept(&blob_proto, conn).await
{
tracing::debug!(%e, "app-blob accept error");
}
}
_ => conn.close(0u32.into(), b"unknown alpn"),
}
});
}
})
}
pub(crate) async fn reload_accept_loop(mesh: &Arc<MeshState>, services: Services) {
let mut guard = mesh.accept_task.lock().await;
if let Some(old) = guard.take() {
old.abort();
}
*guard = Some(spawn_accept_loop(mesh.clone(), Arc::new(services)));
}