use std::collections::{HashMap, VecDeque}; use std::convert::Infallible; use std::net::SocketAddr; use std::sync::{Arc, Mutex}; use std::time::{Duration, Instant}; use anyhow::{Result, anyhow}; use axum::extract::{Query, State}; use axum::http::{StatusCode, header}; use axum::response::sse::{Event, KeepAlive, Sse}; use axum::response::{IntoResponse, Response}; use axum::routing::{get, post}; use axum::{Json, Router}; use serde::Deserialize; use serde_json::json; use tokio::net::TcpListener; use tokio::sync::Notify; use crate::crypto::{now_ts, poll_signing_bytes, random_nonce, verify_signature}; use crate::message::{Envelope, TYPE_AGGREGATE, TYPE_QUERY, TYPE_RESPONSE, timestamp_is_fresh}; use crate::registry::Watcher; pub const DEFAULT_CAPACITY: usize = 256; const CHALLENGE_TTL: Duration = Duration::from_secs(120); const SEEN_TTL: Duration = Duration::from_secs(30); #[derive(Clone, Debug)] pub struct RelayOptions { pub capacity: usize, pub peers: Vec, pub url: Option, pub registry: Option, pub ma_key: Option, pub ca_cert: Option, pub allow_insecure: bool, pub max_hops: usize, } impl Default for RelayOptions { fn default() -> Self { Self { capacity: DEFAULT_CAPACITY, peers: Vec::new(), url: None, registry: None, ma_key: None, ca_cert: None, allow_insecure: false, max_hops: 3, } } } #[derive(Default)] struct MemberQueue { items: VecDeque<(u64, Envelope)>, lagging: bool, missed: u64, } struct Challenge { member: String, created: Instant, } struct Inner { members: HashMap, challenges: HashMap, seen: HashMap, seq: u64, } pub struct Relay { inner: Mutex, options: RelayOptions, registry: Option>, client: reqwest::Client, notify: Notify, } impl Relay { pub fn new(options: RelayOptions) -> Result> { let registry = match (options.registry.as_deref(), options.ma_key.as_deref()) { (Some(source), Some(ma_key)) => Some(Watcher::new(source, ma_key, None)), (Some(_), None) => return Err(anyhow!("registry configured without ma_key")), _ => None, }; if !options.peers.is_empty() && options.url.is_none() { return Err(anyhow!("--url is required when --peer is set")); } if !options.allow_insecure { let mut urls = options.peers.clone(); if let Some(registry) = &options.registry { urls.push(registry.clone()); } let insecure = crate::net::insecure_http_urls(urls); if !insecure.is_empty() { return Err(anyhow!( "refusing plain http endpoints (use https, configure --ca-cert, or set --allow-insecure for private networks): {}", insecure.join(", ") )); } } let client = crate::net::build_client( options.ca_cert.as_deref().map(std::path::Path::new), Duration::from_secs(5), )?; let relay = Arc::new(Self { inner: Mutex::new(Inner { members: HashMap::new(), challenges: HashMap::new(), seen: HashMap::new(), seq: 0, }), options, registry, client, notify: Notify::new(), }); if let Some(watcher) = &relay.registry { watcher.load_initial(); } Ok(relay) } fn admit(&self, envelope: &Envelope) -> bool { match &self.registry { Some(watcher) => { watcher.refresh_if_changed(); watcher.authorized(&envelope.key, now_ts()).is_some() } None => true, } } fn mark_seen(&self, signature: &str) -> bool { let mut inner = self.inner.lock().expect("relay lock"); let now = Instant::now(); inner .seen .retain(|_, at| now.duration_since(*at) < SEEN_TTL); if inner.seen.contains_key(signature) { return false; } inner.seen.insert(signature.to_string(), now); true } fn push_local(&self, target: Option<&str>, envelope: Envelope) -> Result { let outcome = { let mut inner = self.inner.lock().expect("relay lock"); inner.seq += 1; let seq = inner.seq; match target { Some(member) => match inner.members.get_mut(member) { None => Err(StatusCode::NOT_FOUND), Some(queue) => { if queue.items.len() >= self.options.capacity { queue.lagging = true; queue.missed += 1; Err(StatusCode::TOO_MANY_REQUESTS) } else { queue.items.push_back((seq, envelope)); Ok(1) } } }, None => { let publisher = envelope.key.clone(); inner.members.entry(publisher).or_default(); let mut delivered = 0; for queue in inner.members.values_mut() { if queue.items.len() >= self.options.capacity { queue.lagging = true; queue.missed += 1; } else { queue.items.push_back((seq, envelope.clone())); delivered += 1; } } Ok(delivered) } } }; if outcome.is_ok() { self.notify.notify_waiters(); } outcome } fn forward(&self, envelope: Envelope, origin: Option<&str>, hops: usize) { if self.options.peers.is_empty() { return; } let Some(url) = self.options.url.clone() else { return; }; let peers: Vec = self .options .peers .iter() .filter(|peer| Some(peer.as_str()) != origin) .cloned() .collect(); if peers.is_empty() { return; } let client = self.client.clone(); tokio::spawn(async move { for peer in peers { let target = format!( "{}/v1/federation?origin={}&hops={}", peer.trim_end_matches('/'), url, hops ); if let Err(error) = client.post(&target).json(&envelope).send().await { eprintln!("federation forward to {peer} failed: {error}"); } } }); } } pub fn router(relay: Arc) -> Router { Router::new() .route("/health", get(health)) .route("/v1/challenge", get(challenge)) .route("/v1/publish", post(publish)) .route("/v1/federation", post(federation)) .route("/v1/unicast", post(unicast)) .route("/v1/poll", get(poll)) .route("/v1/stream", get(stream)) .with_state(relay) } pub async fn run(listener: TcpListener, options: RelayOptions) -> Result<()> { let relay = Relay::new(options)?; if let Some(watcher) = relay.registry.clone() { if watcher.is_url() { let client = relay.client.clone(); if let Err(error) = watcher.fetch(&client).await { eprintln!("initial registry fetch failed: {error}"); } tokio::spawn(async move { loop { tokio::time::sleep(Duration::from_secs(60)).await; if let Err(error) = watcher.fetch(&client).await { eprintln!("registry fetch failed: {error}"); } } }); } } axum::serve(listener, router(relay)).await?; Ok(()) } pub async fn bind(addr: &str) -> Result<(TcpListener, SocketAddr)> { let listener = TcpListener::bind(addr).await?; let local = listener.local_addr()?; Ok((listener, local)) } fn valid_key(member: &str) -> bool { matches!(hex::decode(member), Ok(bytes) if bytes.len() == 32) } async fn health() -> &'static str { "ok" } #[derive(Deserialize)] struct ChallengeParams { member: String, } async fn challenge( State(relay): State>, Query(params): Query, ) -> Response { if !valid_key(¶ms.member) { return bad_request("valid member key required"); } let nonce = random_nonce(); let mut inner = relay.inner.lock().expect("relay lock"); inner .challenges .retain(|_, challenge| challenge.created.elapsed() < CHALLENGE_TTL); inner.challenges.insert( nonce.clone(), Challenge { member: params.member, created: Instant::now(), }, ); (StatusCode::OK, Json(json!({ "nonce": nonce }))).into_response() } async fn publish(State(relay): State>, Json(envelope): Json) -> Response { if envelope.verify().is_err() { return bad_request("invalid signature"); } if !timestamp_is_fresh(envelope.ts, now_ts()) { return bad_request("stale timestamp"); } if envelope.msg_type != TYPE_QUERY { return bad_request("relay carries broadcast queries only"); } if !relay.admit(&envelope) { return bad_request("sender not admitted"); } relay.mark_seen(&envelope.sig); let delivered = relay.push_local(None, envelope.clone()).unwrap_or(0); relay.forward(envelope, None, 1); ( StatusCode::ACCEPTED, Json(json!({ "delivered": delivered })), ) .into_response() } #[derive(Deserialize)] struct FederationParams { origin: Option, hops: Option, } async fn federation( State(relay): State>, Query(params): Query, Json(envelope): Json, ) -> Response { let Some(origin) = params.origin.clone() else { return bad_request("origin required"); }; if !relay.options.peers.iter().any(|peer| peer == &origin) { return bad_request("unknown peer origin"); } if envelope.verify().is_err() { return bad_request("invalid signature"); } if !timestamp_is_fresh(envelope.ts, now_ts()) { return bad_request("stale timestamp"); } if envelope.msg_type != TYPE_QUERY { return bad_request("relay carries broadcast queries only"); } if !relay.admit(&envelope) { return bad_request("sender not admitted"); } if !relay.mark_seen(&envelope.sig) { return ( StatusCode::ACCEPTED, Json(json!({ "delivered": 0, "duplicate": true })), ) .into_response(); } let delivered = relay.push_local(None, envelope.clone()).unwrap_or(0); let hops = params.hops.unwrap_or(1); if hops < relay.options.max_hops { relay.forward(envelope, Some(&origin), hops + 1); } ( StatusCode::ACCEPTED, Json(json!({ "delivered": delivered })), ) .into_response() } #[derive(Deserialize)] struct UnicastParams { to: String, } async fn unicast( State(relay): State>, Query(params): Query, Json(envelope): Json, ) -> Response { if envelope.verify().is_err() { return bad_request("invalid signature"); } if !timestamp_is_fresh(envelope.ts, now_ts()) { return bad_request("stale timestamp"); } if envelope.msg_type != TYPE_RESPONSE && envelope.msg_type != TYPE_AGGREGATE { return bad_request("unicast carries responses and aggregates only"); } if !relay.admit(&envelope) { return bad_request("sender not admitted"); } match relay.push_local(Some(¶ms.to), envelope) { Ok(_) => (StatusCode::OK, Json(json!({ "delivered": true }))).into_response(), Err(StatusCode::NOT_FOUND) => ( StatusCode::NOT_FOUND, Json(json!({ "error": "member not connected" })), ) .into_response(), Err(status) => backpressure(status, 0), } } #[derive(Deserialize)] struct PollParams { member: String, nonce: String, sig: String, timeout_ms: Option, } fn authenticate(relay: &Relay, params: &PollParams) -> Result<(), Response> { if !valid_key(¶ms.member) { return Err(bad_request("valid member key required")); } if verify_signature( ¶ms.member, &poll_signing_bytes(¶ms.member, ¶ms.nonce), ¶ms.sig, ) .is_err() { return Err(unauthorized("invalid poll signature")); } let mut inner = relay.inner.lock().expect("relay lock"); match inner.challenges.remove(¶ms.nonce) { Some(challenge) if challenge.member == params.member && challenge.created.elapsed() < CHALLENGE_TTL => { Ok(()) } _ => Err(unauthorized("unknown, expired, or reused challenge")), } } async fn stream(State(relay): State>, Query(params): Query) -> Response { if let Err(response) = authenticate(&relay, ¶ms) { return response; } let relay_for_stream = relay.clone(); let member = params.member.clone(); let events = async_stream::stream! { loop { let (batch, lagged): (Vec, Option) = { let mut inner = relay_for_stream.inner.lock().expect("relay lock"); let queue = inner.members.entry(member.clone()).or_default(); if queue.lagging { let missed = queue.missed; queue.lagging = false; queue.missed = 0; queue.items.clear(); (Vec::new(), Some(missed)) } else { ( queue .items .drain(..) .map(|(_, envelope)| envelope) .collect(), None, ) } }; if let Some(missed) = lagged { yield Ok::( Event::default() .event("lag") .data(json!({ "missed": missed }).to_string()), ); continue; } for envelope in batch { if let Ok(data) = serde_json::to_string(&envelope) { yield Ok::(Event::default().event("envelope").data(data)); } } let _ = tokio::time::timeout(Duration::from_secs(15), relay_for_stream.notify.notified()) .await; } }; Sse::new(events) .keep_alive(KeepAlive::default()) .into_response() } async fn poll(State(relay): State>, Query(params): Query) -> Response { if let Err(response) = authenticate(&relay, ¶ms) { return response; } let timeout = Duration::from_millis(params.timeout_ms.unwrap_or(25_000).min(60_000)); let deadline = Instant::now() + timeout; loop { let (batch, lagged): (Vec, Option) = { let mut inner = relay.inner.lock().expect("relay lock"); let queue = inner.members.entry(params.member.clone()).or_default(); if queue.lagging { let missed = queue.missed; queue.lagging = false; queue.missed = 0; queue.items.clear(); (Vec::new(), Some(missed)) } else { ( queue .items .drain(..) .map(|(_, envelope)| envelope) .collect(), None, ) } }; if let Some(missed) = lagged { return backpressure(StatusCode::TOO_MANY_REQUESTS, missed); } if !batch.is_empty() { return (StatusCode::OK, Json(json!({ "messages": batch }))).into_response(); } if Instant::now() >= deadline { return StatusCode::NO_CONTENT.into_response(); } tokio::time::sleep(Duration::from_millis(25)).await; } } fn bad_request(message: &str) -> Response { (StatusCode::BAD_REQUEST, Json(json!({ "error": message }))).into_response() } fn unauthorized(message: &str) -> Response { (StatusCode::UNAUTHORIZED, Json(json!({ "error": message }))).into_response() } fn backpressure(status: StatusCode, missed: u64) -> Response { ( status, [(header::RETRY_AFTER, "1")], Json(json!({ "error": "lagging", "missed": missed })), ) .into_response() } #[cfg(test)] mod tests { use super::*; #[test] fn refuses_plain_http_peers_off_loopback() { let options = RelayOptions { peers: vec!["http://10.0.0.1:7700".to_string()], url: Some("http://10.0.0.1:7700".to_string()), ..Default::default() }; assert!(Relay::new(options).is_err()); let options = RelayOptions { peers: vec!["http://10.0.0.1:7700".to_string()], url: Some("http://10.0.0.1:7700".to_string()), allow_insecure: true, ..Default::default() }; assert!(Relay::new(options).is_ok()); } #[test] fn loopback_peers_need_no_opt_in() { let options = RelayOptions { peers: vec!["http://127.0.0.1:7701".to_string()], url: Some("http://127.0.0.1:7700".to_string()), ..Default::default() }; assert!(Relay::new(options).is_ok()); } }