use std::sync::Arc;
use bamboo_agent_core::AgentEvent;
use dashmap::DashMap;
use tokio::sync::broadcast;
use tokio_util::sync::CancellationToken;
#[derive(Default)]
pub struct SessionWatchers {
counts: DashMap<String, usize>,
relays: DashMap<String, RelayRegistration>,
}
struct RelayRegistration {
channel: broadcast::WeakSender<AgentEvent>,
generation: Arc<()>,
cancel: CancellationToken,
started_at: chrono::DateTime<chrono::Utc>,
}
pub(crate) struct NotificationRelaySubscription {
watchers: Arc<SessionWatchers>,
service: Arc<bamboo_notification::NotificationService>,
session_id: String,
generation: Arc<()>,
cancel: CancellationToken,
receiver: broadcast::Receiver<AgentEvent>,
}
impl NotificationRelaySubscription {
pub(crate) async fn recv(&mut self) -> Result<AgentEvent, broadcast::error::RecvError> {
tokio::select! {
biased;
_ = self.cancel.cancelled() => Err(broadcast::error::RecvError::Closed),
event = self.receiver.recv() => event,
}
}
}
impl Drop for NotificationRelaySubscription {
fn drop(&mut self) {
if let dashmap::mapref::entry::Entry::Occupied(entry) =
self.watchers.relays.entry(self.session_id.clone())
{
if Arc::ptr_eq(&entry.get().generation, &self.generation) {
self.service.end_relay(&self.session_id);
entry.remove();
}
}
}
}
impl SessionWatchers {
pub fn new() -> Arc<Self> {
Arc::new(Self::default())
}
fn watch(&self, session_id: &str) {
*self.counts.entry(session_id.to_string()).or_insert(0) += 1;
}
fn unwatch(&self, session_id: &str) {
let mut drop_entry = false;
if let Some(mut count) = self.counts.get_mut(session_id) {
*count = count.saturating_sub(1);
drop_entry = *count == 0;
}
if drop_entry {
self.counts.remove(session_id);
}
}
pub fn has_watcher(&self, session_id: &str) -> bool {
self.counts.get(session_id).is_some_and(|count| *count > 0)
}
pub(crate) fn begin_notification_relay(
self: &Arc<Self>,
service: Arc<bamboo_notification::NotificationService>,
session_id: &str,
sender: &broadcast::Sender<AgentEvent>,
) -> Option<NotificationRelaySubscription> {
let entry = self.relays.entry(session_id.to_string());
if let dashmap::mapref::entry::Entry::Occupied(occupied) = &entry {
if occupied
.get()
.channel
.upgrade()
.is_some_and(|channel| channel.same_channel(sender))
{
return None;
}
occupied.get().cancel.cancel();
service.end_relay(session_id);
}
let receiver = sender.subscribe();
let generation = Arc::new(());
let cancel = CancellationToken::new();
let registration = RelayRegistration {
channel: sender.downgrade(),
generation: generation.clone(),
cancel: cancel.clone(),
started_at: chrono::Utc::now(),
};
service.try_begin_relay(session_id);
entry.insert(registration);
Some(NotificationRelaySubscription {
watchers: self.clone(),
service,
session_id: session_id.to_string(),
generation,
cancel,
receiver,
})
}
pub(crate) fn has_external_receivers(
&self,
session_id: &str,
sender: &broadcast::Sender<AgentEvent>,
) -> bool {
let relay = self.relays.get(session_id);
let internal = usize::from(relay.as_ref().is_some_and(|relay| {
relay
.channel
.upgrade()
.is_some_and(|channel| channel.same_channel(sender))
}));
sender.receiver_count() > internal
}
pub(crate) fn relay_started_at(
&self,
session_id: &str,
sender: &broadcast::Sender<AgentEvent>,
) -> Option<chrono::DateTime<chrono::Utc>> {
let relay = self.relays.get(session_id)?;
relay
.channel
.upgrade()
.filter(|channel| channel.same_channel(sender))
.map(|_| relay.started_at)
}
}
pub struct WatcherGuard {
watchers: Arc<SessionWatchers>,
session_id: String,
}
impl WatcherGuard {
pub fn new(watchers: Arc<SessionWatchers>, session_id: &str) -> Self {
watchers.watch(session_id);
Self {
watchers,
session_id: session_id.to_string(),
}
}
}
impl Drop for WatcherGuard {
fn drop(&mut self) {
self.watchers.unwatch(&self.session_id);
}
}
#[cfg(test)]
mod tests {
use super::*;
fn notification_service() -> (
Arc<bamboo_notification::NotificationService>,
tempfile::TempDir,
) {
let dir = tempfile::tempdir().unwrap();
(
Arc::new(bamboo_notification::NotificationService::new(
dir.path().join("prefs.json"),
)),
dir,
)
}
#[tokio::test]
async fn relay_generation_does_not_hide_replacement_channel_subscribers() {
let watchers = SessionWatchers::new();
let (service, _dir) = notification_service();
let (old, _) = broadcast::channel(8);
let old_relay = watchers
.begin_notification_relay(service.clone(), "s", &old)
.unwrap();
assert!(!watchers.has_external_receivers("s", &old));
let (replacement, _ui_receiver) = broadcast::channel(8);
assert!(watchers.has_external_receivers("s", &replacement));
let new_relay = watchers
.begin_notification_relay(service.clone(), "s", &replacement)
.unwrap();
drop(old_relay);
assert!(
!service.try_begin_relay("s"),
"old Drop cannot clear the new registration"
);
assert!(watchers.has_external_receivers("s", &replacement));
drop(new_relay);
assert!(service.try_begin_relay("s"));
service.end_relay("s");
}
#[tokio::test]
async fn aborted_and_dropped_relay_tasks_release_receiver_and_registration() {
let watchers = SessionWatchers::new();
let (service, _dir) = notification_service();
let (sender, _) = broadcast::channel(8);
let subscription = watchers
.begin_notification_relay(service.clone(), "aborted", &sender)
.unwrap();
let task = tokio::spawn(async move {
let _subscription = subscription;
std::future::pending::<()>().await;
});
task.abort();
assert!(task.await.unwrap_err().is_cancelled());
assert_eq!(sender.receiver_count(), 0);
assert!(service.try_begin_relay("aborted"));
service.end_relay("aborted");
let subscription = watchers
.begin_notification_relay(service.clone(), "failed-setup", &sender)
.unwrap();
drop(subscription);
assert_eq!(sender.receiver_count(), 0);
assert!(service.try_begin_relay("failed-setup"));
service.end_relay("failed-setup");
}
#[test]
fn watch_and_unwatch_tracks_presence() {
let watchers = SessionWatchers::new();
assert!(!watchers.has_watcher("s1"));
let guard_one = WatcherGuard::new(watchers.clone(), "s1");
assert!(watchers.has_watcher("s1"));
let guard_two = WatcherGuard::new(watchers.clone(), "s1");
assert!(watchers.has_watcher("s1"));
drop(guard_one);
assert!(watchers.has_watcher("s1"), "one watcher remains");
drop(guard_two);
assert!(!watchers.has_watcher("s1"));
}
#[test]
fn distinct_sessions_are_tracked_independently() {
let watchers = SessionWatchers::new();
let _guard = WatcherGuard::new(watchers.clone(), "s1");
assert!(watchers.has_watcher("s1"));
assert!(!watchers.has_watcher("s2"));
}
#[test]
fn repeated_watch_unwatch_never_underflows_or_panics() {
let watchers = SessionWatchers::new();
for _ in 0..5 {
let guard = WatcherGuard::new(watchers.clone(), "s3");
drop(guard);
}
assert!(!watchers.has_watcher("s3"));
}
}