575 lines
18 KiB
Rust
575 lines
18 KiB
Rust
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<String>,
|
|
pub url: Option<String>,
|
|
pub registry: Option<String>,
|
|
pub ma_key: Option<String>,
|
|
pub ca_cert: Option<String>,
|
|
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<String, MemberQueue>,
|
|
challenges: HashMap<String, Challenge>,
|
|
seen: HashMap<String, Instant>,
|
|
seq: u64,
|
|
}
|
|
|
|
pub struct Relay {
|
|
inner: Mutex<Inner>,
|
|
options: RelayOptions,
|
|
registry: Option<Arc<Watcher>>,
|
|
client: reqwest::Client,
|
|
notify: Notify,
|
|
}
|
|
|
|
impl Relay {
|
|
pub fn new(options: RelayOptions) -> Result<Arc<Self>> {
|
|
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<usize, StatusCode> {
|
|
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<String> = 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<Relay>) -> 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<Arc<Relay>>,
|
|
Query(params): Query<ChallengeParams>,
|
|
) -> 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<Arc<Relay>>, Json(envelope): Json<Envelope>) -> 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<String>,
|
|
hops: Option<usize>,
|
|
}
|
|
|
|
async fn federation(
|
|
State(relay): State<Arc<Relay>>,
|
|
Query(params): Query<FederationParams>,
|
|
Json(envelope): Json<Envelope>,
|
|
) -> 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<Arc<Relay>>,
|
|
Query(params): Query<UnicastParams>,
|
|
Json(envelope): Json<Envelope>,
|
|
) -> 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<u64>,
|
|
}
|
|
|
|
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<Arc<Relay>>, Query(params): Query<PollParams>) -> 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<Envelope>, Option<u64>) = {
|
|
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, Infallible>(
|
|
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, Infallible>(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<Arc<Relay>>, Query(params): Query<PollParams>) -> 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<Envelope>, Option<u64>) = {
|
|
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());
|
|
}
|
|
}
|