use super::*;
#[cfg(feature = "handlers")]
const STREAM_HEARTBEAT: Duration = Duration::from_secs(15);
#[cfg(feature = "handlers")]
const STREAM_IDLE_TIMEOUT: Duration = Duration::from_secs(600);
#[cfg(feature = "handlers")]
const DEFAULT_STREAM_CONNECTIONS: u32 = 256;
#[cfg(feature = "handlers")]
const MAX_STREAMS_PER_IP: u32 = 8;
#[cfg(feature = "handlers")]
pub(super) fn route_matches(route: &str, request_path: &str) -> bool {
let path = if request_path.starts_with('/') {
std::borrow::Cow::Borrowed(request_path)
} else {
std::borrow::Cow::Owned(format!("/{request_path}"))
};
Pattern::compile(route)
.map(|pattern| pattern.is_match(&path))
.unwrap_or(false)
}
#[cfg(feature = "handlers")]
struct IpStreamGuard {
counts: Arc<std::sync::Mutex<std::collections::HashMap<(String, IpAddr), u32>>>,
key: (String, IpAddr),
}
#[cfg(feature = "handlers")]
impl Drop for IpStreamGuard {
fn drop(&mut self) {
let mut counts = self.counts.lock().unwrap();
if let Some(n) = counts.get_mut(&self.key) {
*n -= 1;
if *n == 0 {
counts.remove(&self.key);
}
}
}
}
#[cfg(feature = "handlers")]
struct StreamConn {
events: futures::stream::BoxStream<'static, axum::response::sse::Event>,
_site_permit: tokio::sync::OwnedSemaphorePermit,
_ip_guard: IpStreamGuard,
}
#[cfg(feature = "handlers")]
fn acquire_stream_permit(
inner: &HandlerRuntimeInner,
scope: &str,
site_handlers: &boatramp_core::config::HandlersSiteConfig,
) -> Result<tokio::sync::OwnedSemaphorePermit, ()> {
let max = site_handlers
.max_stream_connections
.unwrap_or(DEFAULT_STREAM_CONNECTIONS)
.max(1) as usize;
let semaphore = {
let mut map = inner.stream_semaphores.lock().unwrap();
map.entry(scope.to_string())
.or_insert_with(|| Arc::new(tokio::sync::Semaphore::new(max)))
.clone()
};
semaphore.try_acquire_owned().map_err(|_| ())
}
#[cfg(feature = "handlers")]
fn acquire_stream_ip_slot(
inner: &HandlerRuntimeInner,
scope: &str,
ip: IpAddr,
) -> Result<IpStreamGuard, ()> {
let key = (scope.to_string(), ip);
{
let mut counts = inner.stream_ip_counts.lock().unwrap();
let n = counts.entry(key.clone()).or_insert(0);
if *n >= MAX_STREAMS_PER_IP {
return Err(());
}
*n += 1;
}
Ok(IpStreamGuard {
counts: inner.stream_ip_counts.clone(),
key,
})
}
#[cfg(feature = "handlers")]
#[allow(clippy::too_many_arguments)]
pub(super) async fn serve_stream(
inner: &Arc<HandlerRuntimeInner>,
site: &str,
site_handlers: &boatramp_core::config::HandlersSiteConfig,
stream: &boatramp_core::config::StreamConfig,
after: Option<String>,
client_ip: IpAddr,
preview: Option<&str>,
) -> Response {
use axum::response::sse::{Event, KeepAlive, Sse};
use futures::StreamExt;
let Some(messaging) = inner.messaging.clone() else {
return not_found();
};
let scope = match preview {
Some(id) => format!("{site}/_preview/{id}"),
None => site.to_string(),
};
let site_permit = match acquire_stream_permit(inner, &scope, site_handlers) {
Ok(permit) => permit,
Err(()) => {
return (
StatusCode::SERVICE_UNAVAILABLE,
"site stream connection limit reached\n",
)
.into_response()
}
};
let ip_guard = match acquire_stream_ip_slot(inner, &scope, client_ip) {
Ok(guard) => guard,
Err(()) => {
return (
StatusCode::TOO_MANY_REQUESTS,
"per-client stream connection limit reached\n",
)
.into_response()
}
};
let merged = futures::stream::select_all(stream.topics.iter().map(|topic| {
let namespaced = format!("{scope}/{topic}");
let label = topic.clone();
messaging
.subscribe(&namespaced, after.as_deref())
.map(move |event| stream_event(&label, &event.id, &event.payload))
.boxed()
}));
let conn = StreamConn {
events: merged.boxed(),
_site_permit: site_permit,
_ip_guard: ip_guard,
};
let body = futures::stream::unfold(conn, |mut conn| async move {
match tokio::time::timeout(STREAM_IDLE_TIMEOUT, conn.events.next()).await {
Ok(Some(event)) => Some((Ok::<Event, std::convert::Infallible>(event), conn)),
Ok(None) | Err(_) => None,
}
});
Sse::new(body)
.keep_alive(
KeepAlive::new()
.interval(STREAM_HEARTBEAT)
.text("keep-alive"),
)
.into_response()
}
#[cfg(feature = "handlers")]
#[allow(clippy::too_many_arguments)]
pub(super) async fn serve_graphql_subscription(
inner: &Arc<HandlerRuntimeInner>,
site: &str,
site_handlers: &boatramp_core::config::HandlersSiteConfig,
topic: &str,
after: Option<String>,
client_ip: IpAddr,
preview: Option<&str>,
) -> Response {
use axum::response::sse::{Event, KeepAlive, Sse};
use futures::StreamExt;
let Some(messaging) = inner.messaging.clone() else {
return not_found();
};
let scope = match preview {
Some(id) => format!("{site}/_preview/{id}"),
None => site.to_string(),
};
let site_permit = match acquire_stream_permit(inner, &scope, site_handlers) {
Ok(permit) => permit,
Err(()) => {
return (
StatusCode::SERVICE_UNAVAILABLE,
"site stream connection limit reached\n",
)
.into_response()
}
};
let ip_guard = match acquire_stream_ip_slot(inner, &scope, client_ip) {
Ok(guard) => guard,
Err(()) => {
return (
StatusCode::TOO_MANY_REQUESTS,
"per-client stream connection limit reached\n",
)
.into_response()
}
};
let namespaced = format!("{scope}/{topic}");
let events = messaging
.subscribe(&namespaced, after.as_deref())
.map(|event| graphql_sse_next(&event.id, &event.payload))
.boxed();
let conn = StreamConn {
events,
_site_permit: site_permit,
_ip_guard: ip_guard,
};
let body = futures::stream::unfold(conn, |mut conn| async move {
match tokio::time::timeout(STREAM_IDLE_TIMEOUT, conn.events.next()).await {
Ok(Some(event)) => Some((Ok::<Event, std::convert::Infallible>(event), conn)),
Ok(None) | Err(_) => None,
}
});
Sse::new(body)
.keep_alive(
KeepAlive::new()
.interval(STREAM_HEARTBEAT)
.text("keep-alive"),
)
.into_response()
}
#[cfg(feature = "handlers")]
fn graphql_sse_next(id: &str, payload: &[u8]) -> axum::response::sse::Event {
axum::response::sse::Event::default()
.id(id)
.event("next")
.data(graphql_sse_data(payload))
}
#[cfg(feature = "handlers")]
fn graphql_sse_data(payload: &[u8]) -> String {
match std::str::from_utf8(payload) {
Ok(text) if !text.contains('\r') => text.to_string(),
_ => r#"{"errors":[{"message":"subscription payload was not valid UTF-8 GraphQL JSON"}]}"#
.to_string(),
}
}
#[cfg(all(test, feature = "handlers"))]
mod tests {
use super::graphql_sse_data;
#[test]
fn graphql_sse_frames_a_json_result_verbatim() {
let payload = br#"{"data":{"messageAdded":{"id":"1","body":"hi"}}}"#;
assert_eq!(
graphql_sse_data(payload),
r#"{"data":{"messageAdded":{"id":"1","body":"hi"}}}"#
);
}
#[test]
fn graphql_sse_replaces_an_unframable_payload_with_an_error_result() {
for bad in [b"{\"data\":1}\r".to_vec(), vec![0xff, 0xfe]] {
let data = graphql_sse_data(&bad);
assert!(data.contains("\"errors\""), "unexpected: {data}");
let _: serde_json::Value = serde_json::from_str(&data).expect("valid JSON");
}
}
}
#[cfg(feature = "handlers")]
#[allow(clippy::too_many_arguments)]
pub(super) async fn serve_ws_stream(
inner: &Arc<HandlerRuntimeInner>,
site: &str,
site_handlers: &boatramp_core::config::HandlersSiteConfig,
stream: &boatramp_core::config::StreamConfig,
ws: axum::extract::ws::WebSocketUpgrade,
client_ip: IpAddr,
preview: Option<&str>,
) -> Response {
use futures::StreamExt;
let Some(messaging) = inner.messaging.clone() else {
return not_found();
};
let scope = match preview {
Some(id) => format!("{site}/_preview/{id}"),
None => site.to_string(),
};
let site_permit = match acquire_stream_permit(inner, &scope, site_handlers) {
Ok(permit) => permit,
Err(()) => {
return (
StatusCode::SERVICE_UNAVAILABLE,
"site stream connection limit reached\n",
)
.into_response()
}
};
let ip_guard = match acquire_stream_ip_slot(inner, &scope, client_ip) {
Ok(guard) => guard,
Err(()) => {
return (
StatusCode::TOO_MANY_REQUESTS,
"per-client stream connection limit reached\n",
)
.into_response()
}
};
let mut downstream = futures::stream::select_all(stream.topics.iter().map(|topic| {
let namespaced = format!("{scope}/{topic}");
messaging
.subscribe(&namespaced, None)
.map(|e| e.payload)
.boxed()
}));
let publish_topic = stream
.publish_topic
.as_ref()
.map(|topic| format!("{scope}/{topic}"));
ws.on_upgrade(move |socket| async move {
use axum::extract::ws::Message;
use futures::SinkExt;
let _permits = (site_permit, ip_guard);
let (mut sink, mut incoming) = socket.split();
loop {
tokio::select! {
event = downstream.next() => match event {
Some(payload) => {
if sink.send(Message::Binary(payload.into())).await.is_err() {
break; }
}
None => break, },
msg = incoming.next() => match msg {
Some(Ok(Message::Text(text))) => {
if let Some(topic) = &publish_topic {
let _ = messaging.publish(topic, text.as_bytes()).await;
}
}
Some(Ok(Message::Binary(bytes))) => {
if let Some(topic) = &publish_topic {
let _ = messaging.publish(topic, &bytes).await;
}
}
Some(Ok(_)) => {}
Some(Err(_)) | None => break, },
}
}
})
}
#[cfg(feature = "handlers")]
fn stream_event(topic: &str, id: &str, payload: &[u8]) -> axum::response::sse::Event {
use axum::response::sse::Event;
match std::str::from_utf8(payload) {
Ok(text) if !text.contains('\r') => Event::default().id(id).event(topic).data(text),
_ => {
use base64::Engine;
let encoded = base64::engine::general_purpose::STANDARD.encode(payload);
Event::default()
.id(id)
.event(format!("{topic}.b64"))
.data(encoded)
}
}
}