SSE streaming with long-poll fallback; encrypted unicast profile; registry enc keys
This commit is contained in:
+107
-38
@@ -1,4 +1,5 @@
|
||||
use std::collections::{HashMap, VecDeque};
|
||||
use std::convert::Infallible;
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::{Duration, Instant};
|
||||
@@ -6,12 +7,14 @@ 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};
|
||||
@@ -68,6 +71,7 @@ pub struct Relay {
|
||||
options: RelayOptions,
|
||||
registry: Option<Arc<Watcher>>,
|
||||
client: reqwest::Client,
|
||||
notify: Notify,
|
||||
}
|
||||
|
||||
impl Relay {
|
||||
@@ -92,6 +96,7 @@ impl Relay {
|
||||
client: reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(5))
|
||||
.build()?,
|
||||
notify: Notify::new(),
|
||||
});
|
||||
if let Some(watcher) = &relay.registry {
|
||||
watcher.load_initial();
|
||||
@@ -123,38 +128,45 @@ impl Relay {
|
||||
}
|
||||
|
||||
fn push_local(&self, target: Option<&str>, envelope: Envelope) -> Result<usize, StatusCode> {
|
||||
let mut inner = self.inner.lock().expect("relay lock");
|
||||
inner.seq += 1;
|
||||
let seq = inner.seq;
|
||||
match target {
|
||||
Some(member) => {
|
||||
let Some(queue) = inner.members.get_mut(member) else {
|
||||
return Err(StatusCode::NOT_FOUND);
|
||||
};
|
||||
if queue.items.len() >= self.options.capacity {
|
||||
queue.lagging = true;
|
||||
queue.missed += 1;
|
||||
return Err(StatusCode::TOO_MANY_REQUESTS);
|
||||
}
|
||||
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;
|
||||
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)
|
||||
}
|
||||
Ok(delivered)
|
||||
}
|
||||
};
|
||||
if outcome.is_ok() {
|
||||
self.notify.notify_waiters();
|
||||
}
|
||||
outcome
|
||||
}
|
||||
|
||||
fn forward(&self, envelope: Envelope, origin: Option<&str>, hops: usize) {
|
||||
@@ -199,6 +211,7 @@ pub fn router(relay: Arc<Relay>) -> Router {
|
||||
.route("/v1/federation", post(federation))
|
||||
.route("/v1/unicast", post(unicast))
|
||||
.route("/v1/poll", get(poll))
|
||||
.route("/v1/stream", get(stream))
|
||||
.with_state(relay)
|
||||
}
|
||||
|
||||
@@ -377,9 +390,9 @@ struct PollParams {
|
||||
timeout_ms: Option<u64>,
|
||||
}
|
||||
|
||||
async fn poll(State(relay): State<Arc<Relay>>, Query(params): Query<PollParams>) -> Response {
|
||||
fn authenticate(relay: &Relay, params: &PollParams) -> Result<(), Response> {
|
||||
if !valid_key(¶ms.member) {
|
||||
return bad_request("valid member key required");
|
||||
return Err(bad_request("valid member key required"));
|
||||
}
|
||||
if verify_signature(
|
||||
¶ms.member,
|
||||
@@ -388,16 +401,72 @@ async fn poll(State(relay): State<Arc<Relay>>, Query(params): Query<PollParams>)
|
||||
)
|
||||
.is_err()
|
||||
{
|
||||
return unauthorized("invalid poll signature");
|
||||
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 => {}
|
||||
_ => return unauthorized("unknown, expired, or reused challenge"),
|
||||
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;
|
||||
|
||||
Reference in New Issue
Block a user