use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex, MutexGuard, Weak};
use tokio::sync::{mpsc, oneshot};
use crate::station_link::{self, Event, Link, LinkError, Publication, SignedPublication};
use super::{Pool, PoolError, PoolInner};
const SUBSCRIPTION_BUFFER: usize = 256;
static NEXT_SUBSCRIPTION: AtomicU64 = AtomicU64::new(1);
pub struct Subscription {
inner: Arc<SubInner>,
events: mpsc::Receiver<Event>,
unsubscribed: bool,
}
pub(super) struct SubInner {
id: u64,
pool: Weak<PoolInner>,
realm: [u8; 32],
topic: String,
held: Mutex<Held>,
dropped: AtomicU64,
}
struct Held {
events: Option<mpsc::Sender<Event>>,
on_links: HashMap<u64, Forwarder>,
}
struct Forwarder {
stop: oneshot::Sender<()>,
unsubscribed: oneshot::Receiver<Result<(), LinkError>>,
}
impl Pool {
pub async fn subscribe(
&self,
realm: &[u8; 32],
topic: &str,
) -> Result<Subscription, PoolError> {
let (events_tx, events) = mpsc::channel(SUBSCRIPTION_BUFFER);
let sub = Arc::new(SubInner {
id: NEXT_SUBSCRIPTION.fetch_add(1, Ordering::Relaxed),
pool: Arc::downgrade(&self.inner),
realm: *realm,
topic: topic.to_string(),
held: Mutex::new(Held {
events: Some(events_tx),
on_links: HashMap::new(),
}),
dropped: AtomicU64::new(0),
});
{
let mut state = self.inner.lock();
if state.closed {
return Err(PoolError::Closed);
}
state.subs.insert(sub.id, sub.clone());
}
for link in self.inner.links() {
sub.attach(&link).await;
}
Ok(Subscription {
inner: sub,
events,
unsubscribed: false,
})
}
pub async fn publish(&self, p: Publication) -> Result<(), PoolError> {
let links = self.inner.links();
if links.is_empty() {
return Err(PoolError::NoLink(Vec::new()));
}
let signed =
SignedPublication::sign(&self.inner.opts.identity, &self.inner.publication_seq, p)?;
let mut errors = Vec::new();
let mut sent = 0;
for link in links.iter().take(self.inner.opts.replication_factor) {
match link.publish_signed(&signed).await {
Ok(()) => sent += 1,
Err(e) => errors.push(e),
}
}
if sent == 0 {
return Err(PoolError::NoLink(errors));
}
Ok(())
}
}
impl Subscription {
pub async fn recv(&mut self) -> Option<Event> {
self.events.recv().await
}
pub fn dropped(&self) -> u64 {
self.inner.dropped.load(Ordering::Relaxed)
}
pub async fn unsubscribe(&mut self) -> Result<(), LinkError> {
self.unsubscribed = true;
if let Some(pool) = self.inner.pool.upgrade() {
pool.lock().subs.remove(&self.inner.id);
}
self.inner.end().await
}
}
impl Drop for Subscription {
fn drop(&mut self) {
if self.unsubscribed {
return;
}
if let Some(pool) = self.inner.pool.upgrade() {
pool.lock().subs.remove(&self.inner.id);
}
let inner = self.inner.clone();
if let Ok(runtime) = tokio::runtime::Handle::try_current() {
runtime.spawn(async move {
let _ = inner.end().await;
});
}
}
}
impl SubInner {
fn lock(&self) -> MutexGuard<'_, Held> {
self.held.lock().unwrap_or_else(|p| p.into_inner())
}
pub(super) async fn end(&self) -> Result<(), LinkError> {
let forwarders = {
let mut held = self.lock();
if held.events.take().is_none() {
return Ok(());
}
std::mem::take(&mut held.on_links)
};
let mut result = Ok(());
for (_, f) in forwarders {
let _ = f.stop.send(());
if let Ok(Err(e)) = f.unsubscribed.await {
result = Err(e);
}
}
result
}
pub(super) async fn attach(self: &Arc<Self>, link: &Link) {
let events = {
let held = self.lock();
match &held.events {
Some(events) if !held.on_links.contains_key(&link.serial()) => events.clone(),
_ => return,
}
};
let Ok(on_link) = link.subscribe(&self.realm, &self.topic).await else {
return;
};
let (stop, stopped) = oneshot::channel();
let (unsubscribed_tx, unsubscribed) = oneshot::channel();
let kept = {
let mut held = self.lock();
let wanted = held.events.is_some() && !held.on_links.contains_key(&link.serial());
if wanted {
held.on_links
.insert(link.serial(), Forwarder { stop, unsubscribed });
}
wanted
};
if !kept {
let _ = on_link.unsubscribe().await;
return;
}
tokio::spawn(forward(
self.clone(),
link.serial(),
on_link,
events,
stopped,
unsubscribed_tx,
));
}
}
async fn forward(
sub: Arc<SubInner>,
serial: u64,
mut on_link: station_link::Subscription,
events: mpsc::Sender<Event>,
mut stop: oneshot::Receiver<()>,
unsubscribed: oneshot::Sender<Result<(), LinkError>>,
) {
loop {
tokio::select! {
_ = &mut stop => {
let _ = unsubscribed.send(on_link.unsubscribe().await);
return;
}
event = on_link.recv() => match event {
Some(event) => {
if events.try_send(event).is_err() {
sub.dropped.fetch_add(1, Ordering::Relaxed);
}
}
None => break,
},
}
}
sub.lock().on_links.remove(&serial);
let _ = unsubscribed.send(Ok(()));
}