use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use serde::{Deserialize, Serialize};
use serde_json::Value as JsonValue;
use crate::channels::Channels;
static NEXT_CONNECTION_ID: AtomicU64 = AtomicU64::new(1);
fn next_connection_id() -> u64 {
NEXT_CONNECTION_ID.fetch_add(1, Ordering::Relaxed)
}
const DEFAULT_TTL_SECS: u64 = 30;
#[derive(Clone)]
struct ConnectionPresence {
connection_id: u64,
meta: JsonValue,
last_heartbeat: Instant,
}
struct PresenceInner {
entries: HashMap<String, std::collections::BTreeMap<String, Vec<ConnectionPresence>>>,
ttl: Duration,
}
impl PresenceInner {
fn new(ttl: Duration) -> Self {
Self {
entries: HashMap::new(),
ttl,
}
}
fn add(&mut self, topic: &str, key: &str, connection_id: u64, meta: JsonValue) {
self.entries
.entry(topic.to_owned())
.or_default()
.entry(key.to_owned())
.or_default()
.push(ConnectionPresence {
connection_id,
meta,
last_heartbeat: Instant::now(),
});
}
fn remove(&mut self, topic: &str, key: &str, connection_id: u64) -> bool {
let mut key_fully_removed = false;
if let Some(by_key) = self.entries.get_mut(topic) {
if let Some(conns) = by_key.get_mut(key) {
conns.retain(|c| c.connection_id != connection_id);
if conns.is_empty() {
by_key.remove(key);
key_fully_removed = true;
}
}
if by_key.is_empty() {
self.entries.remove(topic);
}
}
key_fully_removed
}
fn list(&self, topic: &str) -> Vec<PresenceEntry> {
let Some(by_key) = self.entries.get(topic) else {
return Vec::new();
};
by_key
.iter()
.map(|(key, conns)| PresenceEntry {
key: key.clone(),
metas: conns.iter().map(|c| c.meta.clone()).collect(),
})
.collect()
}
fn refresh(&mut self, topic: &str, key: &str, connection_id: u64) {
if let Some(by_key) = self.entries.get_mut(topic)
&& let Some(conns) = by_key.get_mut(key)
{
for c in conns.iter_mut() {
if c.connection_id == connection_id {
c.last_heartbeat = Instant::now();
}
}
}
}
fn sweep_expired(&mut self) -> Vec<(String, String)> {
let ttl = self.ttl;
let now = Instant::now();
let mut removed = Vec::new();
self.entries.retain(|topic, by_key| {
by_key.retain(|key, conns| {
conns.retain(|c| now.duration_since(c.last_heartbeat) < ttl);
if conns.is_empty() {
removed.push((topic.clone(), key.clone()));
false
} else {
true
}
});
!by_key.is_empty()
});
removed
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PresenceEntry {
pub key: String,
pub metas: Vec<JsonValue>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "lowercase")]
pub enum PresenceEvent {
Join {
key: String,
meta: JsonValue,
},
Leave {
key: String,
},
}
#[derive(Clone)]
pub struct Presence {
inner: Arc<Mutex<PresenceInner>>,
channels: Channels,
}
impl Presence {
#[must_use]
pub fn new(channels: Channels) -> Self {
Self::with_ttl(channels, Duration::from_secs(DEFAULT_TTL_SECS))
}
#[must_use]
pub fn with_ttl(channels: Channels, ttl: Duration) -> Self {
Self {
inner: Arc::new(Mutex::new(PresenceInner::new(ttl))),
channels,
}
}
pub fn track(
&self,
topic: impl Into<String>,
key: impl Into<String>,
meta: impl Into<JsonValue>,
) -> PresenceHandle {
let topic = topic.into();
let key = key.into();
let meta = meta.into();
let connection_id = next_connection_id();
{
let mut inner = self.inner.lock().expect("presence lock poisoned");
inner.add(&topic, &key, connection_id, meta.clone());
}
let event = PresenceEvent::Join {
key: key.clone(),
meta,
};
self.publish_event(&topic, &event);
PresenceHandle {
topic,
key,
connection_id,
inner: Arc::clone(&self.inner),
channels: self.channels.clone(),
}
}
#[must_use]
pub fn list(&self, topic: &str) -> Vec<PresenceEntry> {
self.inner
.lock()
.expect("presence lock poisoned")
.list(topic)
}
pub fn sweep_expired(&self) {
let removed = {
let mut inner = self.inner.lock().expect("presence lock poisoned");
inner.sweep_expired()
};
for (topic, key) in removed {
self.publish_event(&topic, &PresenceEvent::Leave { key });
}
}
fn publish_event(&self, topic: &str, event: &PresenceEvent) {
let json = serde_json::to_string(event).unwrap_or_default();
if let Err(e) = self.channels.publish(&format!("presence:{topic}"), json) {
tracing::warn!(topic, error = ?e, "presence: failed to publish event");
}
}
}
pub struct PresenceHandle {
topic: String,
key: String,
connection_id: u64,
inner: Arc<Mutex<PresenceInner>>,
channels: Channels,
}
impl PresenceHandle {
pub fn refresh(&self) {
let mut inner = self.inner.lock().expect("presence lock poisoned");
inner.refresh(&self.topic, &self.key, self.connection_id);
}
#[must_use]
pub fn topic(&self) -> &str {
&self.topic
}
#[must_use]
pub fn key(&self) -> &str {
&self.key
}
}
impl Drop for PresenceHandle {
fn drop(&mut self) {
let key_fully_removed = {
let mut inner = self.inner.lock().expect("presence lock poisoned");
inner.remove(&self.topic, &self.key, self.connection_id)
};
if key_fully_removed {
let event = PresenceEvent::Leave {
key: self.key.clone(),
};
let json = serde_json::to_string(&event).unwrap_or_default();
let _ = self
.channels
.publish(&format!("presence:{}", self.topic), json);
}
}
}
#[cfg(all(feature = "ws", feature = "maud", feature = "htmx"))]
pub fn presence_stream(
state: &crate::state::AppState,
topic: &str,
) -> axum::response::sse::Sse<
impl tokio_stream::Stream<Item = Result<axum::response::sse::Event, std::convert::Infallible>>
+ use<>,
> {
use tokio_stream::StreamExt;
let topic = topic.to_owned();
let presence = state.presence().clone();
let subscriber = state.channels().subscribe(&format!("presence:{topic}"));
let stream = subscriber.into_stream().map(move |_msg| {
let count = presence.list(&topic).len();
let badge_html = presence_badge(count).into_string();
let data = crate::channels::inject_oob_attr(&badge_html, "outerHTML");
Ok::<_, std::convert::Infallible>(axum::response::sse::Event::default().data(data))
});
axum::response::sse::Sse::new(stream).keep_alive(crate::sse::keep_alive())
}
#[must_use]
#[cfg(feature = "maud")]
pub fn presence_badge(count: usize) -> maud::Markup {
maud::html! {
span id="presence-badge" { (count) " online" }
}
}
impl axum::extract::FromRequestParts<crate::state::AppState> for Presence {
type Rejection = std::convert::Infallible;
async fn from_request_parts(
_parts: &mut http::request::Parts,
state: &crate::state::AppState,
) -> Result<Self, Self::Rejection> {
Ok(state.presence().clone())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_presence() -> Presence {
let channels = Channels::new(16);
Presence::new(channels)
}
#[test]
fn track_adds_one_entry() {
let presence = make_presence();
let _handle = presence.track("room:1", "alice", serde_json::json!({"color": "blue"}));
let entries = presence.list("room:1");
assert_eq!(entries.len(), 1);
assert_eq!(entries[0].key, "alice");
assert_eq!(entries[0].metas.len(), 1);
assert_eq!(entries[0].metas[0]["color"], "blue");
}
#[test]
fn drop_handle_removes_entry() {
let presence = make_presence();
{
let _handle = presence.track("room:1", "alice", serde_json::json!({}));
assert_eq!(presence.list("room:1").len(), 1);
}
assert_eq!(presence.list("room:1").len(), 0);
}
#[test]
fn same_key_multiple_connections_collapsed() {
let presence = make_presence();
let _h1 = presence.track("room:1", "alice", serde_json::json!({"tab": 1}));
let _h2 = presence.track("room:1", "alice", serde_json::json!({"tab": 2}));
let entries = presence.list("room:1");
assert_eq!(entries.len(), 1, "same key should collapse into one entry");
assert_eq!(entries[0].key, "alice");
assert_eq!(entries[0].metas.len(), 2);
}
#[test]
fn dropping_one_connection_keeps_other() {
let presence = make_presence();
let _h1 = presence.track("room:1", "alice", serde_json::json!({"tab": 1}));
{
let _h2 = presence.track("room:1", "alice", serde_json::json!({"tab": 2}));
assert_eq!(presence.list("room:1")[0].metas.len(), 2);
}
let entries = presence.list("room:1");
assert_eq!(entries.len(), 1);
assert_eq!(entries[0].metas.len(), 1);
}
#[test]
fn different_keys_are_separate_entries() {
let presence = make_presence();
let _h1 = presence.track("room:1", "alice", serde_json::json!({}));
let _h2 = presence.track("room:1", "bob", serde_json::json!({}));
let mut entries = presence.list("room:1");
entries.sort_by_key(|e| e.key.clone());
assert_eq!(entries.len(), 2);
assert_eq!(entries[0].key, "alice");
assert_eq!(entries[1].key, "bob");
}
#[test]
fn list_unknown_topic_returns_empty() {
let presence = make_presence();
assert!(presence.list("nonexistent").is_empty());
}
#[test]
fn sweep_removes_stale_entries() {
let channels = Channels::new(16);
let presence = Presence::with_ttl(channels, Duration::from_nanos(1));
let _handle = presence.track("room:1", "alice", serde_json::json!({}));
std::thread::sleep(Duration::from_millis(1));
presence.sweep_expired();
assert!(presence.list("room:1").is_empty());
}
#[test]
fn sweep_respects_refreshed_entries() {
let channels = Channels::new(16);
let presence = Presence::with_ttl(channels, Duration::from_millis(500));
let handle = presence.track("room:1", "alice", serde_json::json!({}));
std::thread::sleep(Duration::from_millis(100));
handle.refresh(); std::thread::sleep(Duration::from_millis(100));
presence.sweep_expired();
assert_eq!(presence.list("room:1").len(), 1);
}
#[test]
fn presence_handle_exposes_topic_and_key() {
let presence = make_presence();
let handle = presence.track("chat:42", "user_7", serde_json::json!({}));
assert_eq!(handle.topic(), "chat:42");
assert_eq!(handle.key(), "user_7");
}
#[tokio::test]
async fn no_leave_event_while_other_connections_remain() {
let channels = Channels::new(16);
let presence = Presence::new(channels.clone());
let mut rx = channels.subscribe("presence:room:1");
let h1 = presence.track("room:1", "alice", serde_json::json!({"tab": 1}));
let _h2 = presence.track("room:1", "alice", serde_json::json!({"tab": 2}));
let _ = tokio::time::timeout(std::time::Duration::from_millis(50), rx.recv()).await;
let _ = tokio::time::timeout(std::time::Duration::from_millis(50), rx.recv()).await;
drop(h1);
let result = tokio::time::timeout(std::time::Duration::from_millis(50), rx.recv()).await;
assert!(
result.is_err(),
"no Leave event should be emitted while another connection is open"
);
assert_eq!(presence.list("room:1")[0].metas.len(), 1);
}
#[tokio::test]
async fn join_event_broadcast_on_track() {
let channels = Channels::new(16);
let mut rx = channels.subscribe("presence:room:1");
let presence = Presence::new(channels);
let _handle = presence.track("room:1", "alice", serde_json::json!({"name": "Alice"}));
let msg = tokio::time::timeout(std::time::Duration::from_millis(200), rx.recv())
.await
.expect("join event timed out")
.expect("channel closed");
let event: serde_json::Value = serde_json::from_str(msg.as_str()).unwrap();
assert_eq!(event["type"], "join");
assert_eq!(event["key"], "alice");
assert_eq!(event["meta"]["name"], "Alice");
}
#[tokio::test]
async fn leave_event_broadcast_on_drop() {
let channels = Channels::new(16);
let presence = Presence::new(channels.clone());
let mut rx = channels.subscribe("presence:room:1");
{
let _handle = presence.track("room:1", "alice", serde_json::json!({}));
let _ = tokio::time::timeout(std::time::Duration::from_millis(50), rx.recv()).await;
}
let msg = tokio::time::timeout(std::time::Duration::from_millis(200), rx.recv())
.await
.expect("leave event timed out")
.expect("channel closed");
let event: serde_json::Value = serde_json::from_str(msg.as_str()).unwrap();
assert_eq!(event["type"], "leave");
assert_eq!(event["key"], "alice");
}
#[tokio::test]
async fn sweep_broadcasts_leave_events() {
let channels = Channels::new(16);
let presence = Presence::with_ttl(channels.clone(), Duration::from_nanos(1));
let mut rx = channels.subscribe("presence:room:1");
presence.track("room:1", "alice", serde_json::json!({}));
let _ = tokio::time::timeout(std::time::Duration::from_millis(50), rx.recv()).await;
std::thread::sleep(Duration::from_millis(1));
presence.sweep_expired();
let msg = tokio::time::timeout(std::time::Duration::from_millis(200), rx.recv())
.await
.expect("sweep leave event timed out")
.expect("channel closed");
let event: serde_json::Value = serde_json::from_str(msg.as_str()).unwrap();
assert_eq!(event["type"], "leave");
}
#[test]
fn presence_event_join_serializes_correctly() {
let event = PresenceEvent::Join {
key: "alice".to_owned(),
meta: serde_json::json!({"role": "admin"}),
};
let json = serde_json::to_string(&event).unwrap();
let parsed: serde_json::Value = serde_json::from_str(&json).unwrap();
assert_eq!(parsed["type"], "join");
assert_eq!(parsed["key"], "alice");
assert_eq!(parsed["meta"]["role"], "admin");
}
#[test]
fn presence_event_leave_serializes_correctly() {
let event = PresenceEvent::Leave {
key: "alice".to_owned(),
};
let json = serde_json::to_string(&event).unwrap();
let parsed: serde_json::Value = serde_json::from_str(&json).unwrap();
assert_eq!(parsed["type"], "leave");
assert_eq!(parsed["key"], "alice");
}
#[cfg(feature = "maud")]
#[test]
fn presence_badge_renders_count() {
let badge = super::presence_badge(3);
let html = badge.into_string();
assert!(html.contains("presence-badge"), "missing id");
assert!(html.contains('3'), "missing count");
assert!(html.contains("online"), "missing label");
}
#[cfg(feature = "maud")]
#[test]
fn presence_badge_zero_renders() {
let html = super::presence_badge(0).into_string();
assert!(html.contains('0'));
assert!(html.contains("online"));
}
#[cfg(all(feature = "ws", feature = "maud", feature = "htmx"))]
#[tokio::test]
async fn presence_stream_emits_join_with_count() {
let state = crate::AppState::for_test();
let presence_svc = state.presence().clone();
let mut rx = state.channels().subscribe("presence:stream-test");
let _handle = presence_svc.track("stream-test", "bob", serde_json::json!({}));
let msg = tokio::time::timeout(std::time::Duration::from_millis(200), rx.recv())
.await
.expect("timed out waiting for join event")
.expect("channel closed");
let val: serde_json::Value = serde_json::from_str(msg.as_str()).unwrap();
assert_eq!(val["type"], "join");
assert_eq!(val["key"], "bob");
let count = presence_svc.list("stream-test").len();
assert_eq!(count, 1);
let badge = super::presence_badge(count);
let html = badge.into_string();
assert!(html.contains('1'));
assert!(html.contains("presence-badge"));
let _sse = presence_stream(&state, "stream-test");
}
#[cfg(all(feature = "ws", feature = "maud", feature = "htmx"))]
#[tokio::test]
async fn presence_stream_body_emits_oob_badge_on_join() {
use axum::response::IntoResponse;
use http_body_util::BodyExt;
let state = crate::AppState::for_test();
let presence_svc = state.presence().clone();
let sse = presence_stream(&state, "body-poll-test");
let response = sse.into_response();
let _handle = presence_svc.track("body-poll-test", "frank", serde_json::json!({}));
let mut body = response.into_body();
let frame_result =
tokio::time::timeout(std::time::Duration::from_millis(500), body.frame())
.await
.expect("timed out waiting for SSE frame from presence_stream");
match frame_result {
Some(Ok(frame)) => {
if let Ok(data) = frame.into_data() {
let text = String::from_utf8_lossy(&data);
assert!(
text.contains("presence-badge") || text.contains("hx-swap-oob"),
"expected OOB badge in SSE frame, got: {text}"
);
}
}
Some(Err(e)) => panic!("SSE body error: {e}"),
None => panic!("SSE body ended without yielding a frame"),
}
}
#[test]
fn presence_event_round_trips() {
let events = vec![
PresenceEvent::Join {
key: "bob".to_owned(),
meta: serde_json::json!({"tab": 2}),
},
PresenceEvent::Leave {
key: "bob".to_owned(),
},
];
for event in events {
let json = serde_json::to_string(&event).unwrap();
let _parsed: PresenceEvent = serde_json::from_str(&json).unwrap();
}
}
}