use std::convert::Infallible;
use std::sync::Arc;
use axum::extract::State;
use axum::http::{HeaderMap, StatusCode};
use axum::response::sse::{Event, KeepAlive, Sse};
use axum::response::{IntoResponse, Response};
use axum::routing::get;
use axum::Router;
use tokio_stream::wrappers::ReceiverStream;
use uuid::Uuid;
use crate::application::service::chatter_acl::{MessagingIdentity, ThreadAccessResolver};
use crate::domain::event::constants::{discuss_channel, guest_channel, partner_channel};
use crate::infrastructure::persistence::channel_repository::ChannelRepository;
use crate::presentation::http::thread_routes::ApiState;
use crate::presentation::middleware::{unauthorized, WireIdentity};
use crate::realtime::session::{mint_session_token, SESSION_TTL_SECS};
use crate::realtime::SessionSecret;
use crate::realtime::tailer::parse_record_key;
use super::registry::{identity_key, ConnectionGuard, RealtimeRegistry, TailedEvent};
const MAX_STREAMS_PER_IDENTITY: usize = 4;
const MAX_MEMBERSHIP_CHANNELS: usize = 64;
struct Allowlist {
own: String,
membership: std::collections::HashSet<String>,
}
impl Allowlist {
async fn build(pool: &sqlx::PgPool, identity: &MessagingIdentity) -> Self {
let mut membership = std::collections::HashSet::new();
match ChannelRepository::membership_channel_ids(pool, identity.partner_id(), identity.guest_id()).await {
Ok(ids) => {
for id in ids {
if membership.len() >= MAX_MEMBERSHIP_CHANNELS {
tracing::warn!(
target: "mail::realtime_stream",
identity = %identity_key(identity),
"membership allowlist truncated at MAX_MEMBERSHIP_CHANNELS"
);
break;
}
membership.insert(discuss_channel(id));
}
}
Err(e) => {
tracing::error!(target: "mail::realtime_stream", error = %e, "membership lookup failed");
}
}
Self { own: identity.channel(), membership }
}
async fn permits(
&self,
pool: &sqlx::PgPool,
identity: &MessagingIdentity,
acl: &crate::application::service::chatter_acl::ThreadAclSlot,
event: &TailedEvent,
) -> bool {
if event.channel == self.own {
return true;
}
if self.membership.contains(&event.channel) {
return true;
}
match parse_record_key(&event.channel) {
Some((model, res_id)) => acl.can_read(pool, identity, model, res_id).await,
None => false,
}
}
}
fn sse_event(event: &TailedEvent) -> Event {
Event::default()
.id(event.id.to_string())
.event(event.message_type.as_str())
.data(event.payload.to_string())
}
async fn stream(
State(app): State<ApiState>,
axum::Extension(identity): axum::Extension<WireIdentity>,
axum::Extension(secret): axum::Extension<SessionSecret>,
headers: HeaderMap,
) -> Response {
let Ok(id) = crate::presentation::http::thread_routes::require_identity(&identity) else {
return unauthorized();
};
let registry = Arc::clone(&app.realtime_registry);
if registry.reserve_connection(&identity_key(&id), MAX_STREAMS_PER_IDENTITY).is_err() {
return (
StatusCode::TOO_MANY_REQUESTS,
axum::Json(serde_json::json!({
"error": format!("too many concurrent streams for this identity (max {MAX_STREAMS_PER_IDENTITY})")
})),
)
.into_response();
}
let _guard = ConnectionGuard::new(Arc::clone(®istry), &id);
let last_id = headers
.get("last-event-id")
.and_then(|v| v.to_str().ok())
.and_then(|v| Uuid::parse_str(v).ok());
let replay = registry.replay_after(last_id);
let mut live = registry.subscribe();
let pool = app.db_pool();
let acl = app.thread_acl.clone();
let allow = Allowlist::build(&pool, &id).await;
let proof = mint_session_token(&secret.0, &id, SESSION_TTL_SECS);
let (tx, rx) = tokio::sync::mpsc::channel::<Result<Event, Infallible>>(64);
tokio::spawn(async move {
let _guard = _guard; let send = |ev: Event| tx.send(Ok(ev));
if send(Event::default().event("session").data(proof)).await.is_err() {
return; }
for event in &replay {
if allow.permits(&pool, &id, &acl, event).await {
if send(sse_event(event)).await.is_err() {
return;
}
} else {
tracing::warn!(
target: "mail::realtime_stream",
channel = %event.channel,
outbox_id = %event.id,
"BUS-B2: dropping out-of-allowlist event for this identity"
);
}
}
loop {
match live.recv().await {
Ok(event) => {
if allow.permits(&pool, &id, &acl, &event).await {
if send(sse_event(&event)).await.is_err() {
return; }
} else {
tracing::warn!(
target: "mail::realtime_stream",
channel = %event.channel,
outbox_id = %event.id,
"BUS-B2: dropping out-of-allowlist event for this identity"
);
}
}
Err(tokio::sync::broadcast::error::RecvError::Lagged(n)) => {
tracing::warn!(
target: "mail::realtime_stream",
missed = n,
"stream lagged; client should reconnect with Last-Event-ID"
);
}
Err(tokio::sync::broadcast::error::RecvError::Closed) => return,
}
}
});
let stream: ReceiverStream<Result<Event, Infallible>> = ReceiverStream::new(rx);
let mut response = Sse::new(stream)
.keep_alive(KeepAlive::new().interval(std::time::Duration::from_secs(15)))
.into_response();
response.headers_mut().insert(
axum::http::HeaderName::from_static("x-accel-buffering"),
axum::http::HeaderValue::from_static("no"),
);
response
}
pub fn composer() -> Router<ApiState> {
Router::new().route("/mail/realtime/stream", get(stream))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn own_and_membership_keys_are_direct_allows_by_construction() {
let p = Uuid::new_v4();
assert_eq!(partner_channel(p), format!("res.partner_{p}"));
let g = Uuid::new_v4();
assert_eq!(guest_channel(g), format!("mail.guest_{g}"));
let c = Uuid::new_v4();
assert_eq!(discuss_channel(c), format!("discuss.channel_{c}"));
}
}