use std::future::Future;
use std::path::PathBuf;
use std::sync::Arc;
use anyhow::{Context, Result};
use iroh_blobs::ticket::BlobTicket;
use iroh_gossip::api::{GossipReceiver, GossipSender};
use crate::roster::gate::RosterGate;
use crate::roster::transport::{self, RosterAddrBook, RosterAnnounce, RosterBlobs};
use std::time::{Duration, Instant};
pub trait DistributionHost: Send + Sync + 'static {
fn endpoint(&self) -> &iroh::Endpoint;
fn roster(&self) -> &RosterGate;
fn blobs(&self) -> Option<&RosterBlobs>;
fn gossip_active(&self) -> bool;
fn installed_roster_path(&self) -> PathBuf;
fn pinned_org_root_pk(&self) -> Result<Option<String>>;
fn addr_book(&self) -> Option<Arc<RosterAddrBook>>;
fn roster_topic_sender(&self) -> impl Future<Output = Option<GossipSender>> + Send;
fn take_roster_topic_receiver(&self) -> impl Future<Output = Option<GossipReceiver>> + Send;
fn confirm_roster_current(&self, now: i64) -> impl Future<Output = ()> + Send;
fn install_roster_bytes(
&self,
bytes: &[u8],
serial: u64,
channel: &'static str,
) -> impl Future<Output = Result<bool>> + Send;
}
const GOSSIP_FETCH_TIMEOUT: Duration = Duration::from_secs(30);
const GOSSIP_FETCH_CONCURRENCY: usize = 4;
const GOSSIP_ANNOUNCE_PER_MIN: u32 = 60;
const POLL_TIMEOUT: Duration = Duration::from_secs(10);
const MAX_ROSTER_BYTES: usize = 4 * 1024 * 1024;
pub async fn announce_roster<H: DistributionHost>(mesh: &Arc<H>) -> Result<()> {
if !mesh.gossip_active() {
return Ok(()); }
let Some(blobs) = mesh.blobs() else {
return Ok(()); };
let Some(view) = mesh.roster().view() else {
return Ok(()); };
let serial = view.serial();
let path = mesh.installed_roster_path();
let bytes = crate::util::blocking("join roster read", move || std::fs::read(path))
.await?
.context("read installed roster for announce")?;
let (ticket, roster_hash) = blobs.publish(&bytes, mesh.endpoint()).await?;
let announce = RosterAnnounce {
serial,
roster_hash,
blob_ticket: ticket,
};
if let Some(sender) = mesh.roster_topic_sender().await {
transport::broadcast(&sender, announce.to_bytes()).await?;
}
Ok(())
}
pub async fn on_announce<H: DistributionHost>(
mesh: &Arc<H>,
announce: RosterAnnounce,
) -> Result<()> {
if announce.serial <= mesh.roster().view().map(|v| v.serial()).unwrap_or(0) {
return Ok(()); }
let Some(blobs) = mesh.blobs() else {
return Ok(()); };
if let Ok(ticket) = announce.blob_ticket.parse::<BlobTicket>() {
if let Some(book) = mesh.addr_book() {
book.note(ticket.addr().clone());
} else {
let mem = iroh::address_lookup::MemoryLookup::new();
mem.add_endpoint_info(ticket.addr().clone());
if let Ok(lookup) = mesh.endpoint().address_lookup() {
lookup.add(mem);
}
}
}
let bytes = match tokio::time::timeout(
GOSSIP_FETCH_TIMEOUT,
blobs.fetch(
&announce.blob_ticket,
&announce.roster_hash,
mesh.endpoint(),
),
)
.await
{
Ok(r) => r.context("fetch announced roster blob")?,
Err(_) => {
tracing::debug!("gossip roster fetch timed out; dropping (will re-converge)");
return Ok(()); }
};
if mesh
.install_roster_bytes(&bytes, announce.serial, "gossip")
.await?
{
return announce_roster(mesh).await;
}
Ok(())
}
pub(crate) async fn fetch_capped(url: &str, max: usize, timeout: Duration) -> Result<Vec<u8>> {
let client = reqwest::Client::builder()
.timeout(timeout)
.build()
.context("build roster poll client")?;
let mut resp = client
.get(url)
.send()
.await
.context("GET roster url")?
.error_for_status()
.context("roster url status")?;
if let Some(len) = resp.content_length() {
anyhow::ensure!(
len as usize <= max,
"roster body exceeds {max} bytes (content-length {len})"
);
}
let mut body = Vec::new();
while let Some(chunk) = resp.chunk().await.context("read roster url body chunk")? {
anyhow::ensure!(
body.len() + chunk.len() <= max,
"roster body exceeds {max} bytes (streamed)"
);
body.extend_from_slice(&chunk);
}
Ok(body)
}
pub async fn poll_roster_url_once<H: DistributionHost>(mesh: &Arc<H>, url: &str) -> Result<()> {
let body = fetch_capped(url, MAX_ROSTER_BYTES, POLL_TIMEOUT)
.await
.context("poll roster url")?;
let now = crate::util::epoch_now_i64();
let installed = mesh.roster().view().map(|v| v.serial()).unwrap_or(0);
let parsed = serde_json::from_slice::<mcpmesh_trust::roster::Roster>(&body).ok();
let parsed_serial = parsed.as_ref().map(|r| r.serial);
if let Some(s) = parsed_serial.filter(|s| *s > installed) {
if mesh.install_roster_bytes(&body, s, "url-poll").await? {
let _ = announce_roster(mesh).await;
}
} else if let Some(s) = parsed_serial {
if s == installed {
if equal_serial_body_is_authentic(&**mesh, parsed.as_ref()) {
mesh.confirm_roster_current(now).await;
} else {
tracing::warn!(
serial = s,
"roster URL served an unauthenticated/mismatched body at the installed serial; \
not confirming currency (org-root sig is the sole trust input)"
);
}
} else {
tracing::debug!(
serial = s,
installed,
"roster URL served a stale (older) serial; ignoring"
);
}
} else {
tracing::warn!("roster URL body did not parse as a signed roster; check [roster].url");
}
Ok(())
}
fn equal_serial_body_is_authentic<H: DistributionHost>(
mesh: &H,
parsed: Option<&mcpmesh_trust::roster::Roster>,
) -> bool {
let Some(roster) = parsed else {
return false;
};
let Ok(Some(pk_b64)) = mesh.pinned_org_root_pk() else {
return false;
};
let Ok(pk) = crate::roster::parse_org_root_pk(&pk_b64) else {
return false;
};
mcpmesh_trust::roster::sign::verify(roster, &pk).is_ok()
}
pub fn spawn_poll_loop<H: DistributionHost>(
mesh: Arc<H>,
url: String,
interval_secs: i64,
) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move {
let period = std::time::Duration::from_secs(interval_secs.max(1) as u64);
loop {
if let Err(e) = poll_roster_url_once(&mesh, &url).await {
tracing::debug!(%e, "roster URL poll failed; will retry next interval");
}
tokio::time::sleep(period).await;
}
})
}
pub fn spawn_receive_loop<H: DistributionHost>(mesh: Arc<H>) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move {
let Some(mut receiver) = mesh.take_roster_topic_receiver().await else {
return;
};
let fetch_slots =
std::sync::Arc::new(tokio::sync::Semaphore::new(GOSSIP_FETCH_CONCURRENCY));
let mut announce_bucket = crate::limits::TokenBucket::new(
f64::from(GOSSIP_ANNOUNCE_PER_MIN),
f64::from(GOSSIP_ANNOUNCE_PER_MIN) / 60.0,
Instant::now(),
);
while let Some(content) = transport::next_message(&mut receiver).await {
let announce = match RosterAnnounce::from_bytes(&content) {
Ok(a) => a,
Err(e) => {
tracing::debug!(%e, "malformed roster announce dropped");
continue;
}
};
if announce_bucket.try_take(Instant::now()).is_err() {
tracing::debug!("gossip announce rate limit engaged; dropping announce");
continue;
}
let Ok(permit) = fetch_slots.clone().try_acquire_owned() else {
tracing::debug!("gossip fetch pool full; dropping announce (will re-converge)");
continue;
};
let mesh2 = mesh.clone();
tokio::spawn(async move {
let _permit = permit; if let Err(e) = on_announce(&mesh2, announce).await {
tracing::debug!(%e, "gossip roster announce handling failed");
}
});
}
})
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::{Read, Write};
use std::net::TcpListener;
fn serve_once(body: Vec<u8>, sleep_ms: u64) -> String {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
std::thread::spawn(move || {
if let Ok((mut stream, _)) = listener.accept() {
let mut buf = [0u8; 1024];
let _ = stream.read(&mut buf);
if sleep_ms > 0 {
std::thread::sleep(std::time::Duration::from_millis(sleep_ms));
}
let header = format!(
"HTTP/1.1 200 OK\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
body.len()
);
let _ = stream.write_all(header.as_bytes());
let _ = stream.write_all(&body);
let _ = stream.flush();
}
});
format!("http://{addr}/roster.json")
}
fn install_ring() {
let _ = rustls::crypto::ring::default_provider().install_default();
}
#[tokio::test]
async fn fetch_capped_reads_a_small_body() {
install_ring();
let url = serve_once(b"{\"format\":\"mcpmesh-roster/1\"}".to_vec(), 0);
let got = fetch_capped(&url, 1024, Duration::from_secs(5))
.await
.unwrap();
assert!(got.starts_with(b"{\"format\""));
}
#[tokio::test]
async fn fetch_capped_rejects_an_oversized_body_without_oom() {
install_ring();
let url = serve_once(vec![b'x'; 2 * 1024 * 1024], 0);
let err = fetch_capped(&url, 64 * 1024, Duration::from_secs(5))
.await
.unwrap_err();
assert!(
format!("{err:#}").contains("exceeds"),
"size cap rejects: {err:#}"
);
}
#[tokio::test]
async fn fetch_capped_times_out_a_hung_host() {
install_ring();
let url = serve_once(b"late".to_vec(), 2000);
let err = fetch_capped(&url, 1024, Duration::from_millis(200))
.await
.unwrap_err();
let _ = err;
}
}