use std::collections::{HashMap, HashSet};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, RwLock};
use bytes::Bytes;
use tokio::sync::mpsc::OwnedPermit;
use tokio::sync::{Semaphore, mpsc};
use weida_core::{Error, Limits, TraceContext};
use weida_protocol::header::OrderingMode;
use weida_protocol::{DataHeader, filter};
use crate::conn::ConnHandle;
use crate::ordering::Sequencer;
use crate::transfer::write_data_preamble;
use crate::transport::SendHalf;
use weida_protocol::codes;
const WRITER_QUEUE: usize = 1024;
struct PubMsg {
topic: Arc<str>,
payload: Bytes,
trace: Option<TraceContext>,
sequence: Option<u64>,
}
enum PubItem {
Whole(PubMsg),
Begin { id: u64, head: PubMsg },
Chunk { id: u64, payload: Bytes },
Finish { id: u64 },
Abort { id: u64 },
}
struct SubEntry {
conn_id: usize,
filters: HashSet<String>,
tx: mpsc::Sender<PubItem>,
budget: Arc<Semaphore>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub(crate) enum DropCause {
SubscriberBudget,
SubscriberQueue,
NoParkedConnection,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct TopicDrops {
pub topic: Arc<str>,
pub subscriber_budget: u64,
pub subscriber_queue: u64,
pub no_parked_connection: u64,
}
impl TopicDrops {
pub fn total(&self) -> u64 {
self.subscriber_budget + self.subscriber_queue + self.no_parked_connection
}
}
#[derive(Default)]
struct Causes {
budget: AtomicU64,
queue: AtomicU64,
no_parked: AtomicU64,
}
impl Causes {
fn record(&self, cause: DropCause) {
let counter = match cause {
DropCause::SubscriberBudget => &self.budget,
DropCause::SubscriberQueue => &self.queue,
DropCause::NoParkedConnection => &self.no_parked,
};
counter.fetch_add(1, Ordering::Relaxed);
}
fn snapshot(&self, topic: &Arc<str>) -> TopicDrops {
TopicDrops {
topic: Arc::clone(topic),
subscriber_budget: self.budget.load(Ordering::Relaxed),
subscriber_queue: self.queue.load(Ordering::Relaxed),
no_parked_connection: self.no_parked.load(Ordering::Relaxed),
}
}
}
struct DropTable {
total: AtomicU64,
per_topic: RwLock<HashMap<Arc<str>, Causes>>,
max_topics: usize,
}
impl DropTable {
fn new(max_topics: usize) -> DropTable {
DropTable {
total: AtomicU64::new(0),
per_topic: RwLock::new(HashMap::new()),
max_topics,
}
}
fn record(&self, topic: &Arc<str>, cause: DropCause) {
self.total.fetch_add(1, Ordering::Relaxed);
{
let table = self.per_topic.read().expect("drop table poisoned");
if let Some(causes) = table.get(topic) {
causes.record(cause);
return;
}
}
let mut table = self.per_topic.write().expect("drop table poisoned");
if let Some(causes) = table.get(topic) {
causes.record(cause);
} else if table.len() < self.max_topics {
let causes = Causes::default();
causes.record(cause);
table.insert(Arc::clone(topic), causes);
}
}
fn on_topic(&self, topic: &str) -> Option<TopicDrops> {
let table = self.per_topic.read().expect("drop table poisoned");
table
.get_key_value(topic)
.map(|(topic, causes)| causes.snapshot(topic))
}
fn by_topic(&self) -> Vec<TopicDrops> {
let table = self.per_topic.read().expect("drop table poisoned");
table
.iter()
.map(|(topic, causes)| causes.snapshot(topic))
.collect()
}
}
struct PathState {
subs: Vec<SubEntry>,
drops: Arc<DropTable>,
}
pub(crate) struct SubRegistry {
paths: RwLock<HashMap<Arc<str>, PathState>>,
per_conn: RwLock<HashMap<usize, usize>>,
limits: Limits,
sequencer: Sequencer,
next_stream: AtomicU64,
}
impl SubRegistry {
pub(crate) fn new(limits: Limits, ordering: OrderingMode) -> SubRegistry {
SubRegistry {
paths: RwLock::new(HashMap::new()),
per_conn: RwLock::new(HashMap::new()),
limits,
sequencer: Sequencer::new(ordering),
next_stream: AtomicU64::new(0),
}
}
pub(crate) fn subscribe(
&self,
path: &str,
ctx: &ConnHandle,
filter: String,
) -> Result<(), Error> {
let conn_id = ctx.conn.stable_id();
let mut paths = self.paths.write().expect("subscription lock poisoned");
let mut per_conn = self.per_conn.write().expect("subscription lock poisoned");
let state = paths.entry(Arc::from(path)).or_insert_with(|| PathState {
subs: Vec::new(),
drops: Arc::new(DropTable::new(self.limits.max_sequence_scopes)),
});
if let Some(entry) = state.subs.iter_mut().find(|e| e.conn_id == conn_id) {
if entry.filters.contains(&filter) {
return Ok(());
}
let count = per_conn.entry(conn_id).or_insert(0);
if *count >= self.limits.max_subscriptions {
return Err(Error::LimitExceeded);
}
*count += 1;
entry.filters.insert(filter);
return Ok(());
}
let count = per_conn.entry(conn_id).or_insert(0);
if *count >= self.limits.max_subscriptions {
return Err(Error::LimitExceeded);
}
*count += 1;
let (tx, rx) = mpsc::channel(WRITER_QUEUE);
let budget = Arc::new(Semaphore::new(self.limits.subscriber_buffer_bytes));
let drops = Arc::clone(&state.drops);
state.subs.push(SubEntry {
conn_id,
filters: HashSet::from([filter]),
tx,
budget: Arc::clone(&budget),
});
ctx.exec
.spawn(writer(Arc::clone(ctx), Arc::from(path), rx, budget, drops));
Ok(())
}
pub(crate) fn reserve(&self, conn_id: usize) -> Result<(), Error> {
let mut per_conn = self.per_conn.write().expect("subscription lock poisoned");
let count = per_conn.entry(conn_id).or_insert(0);
if *count >= self.limits.max_subscriptions {
return Err(Error::LimitExceeded);
}
*count += 1;
Ok(())
}
pub(crate) fn release(&self, conn_id: usize) {
let mut per_conn = self.per_conn.write().expect("subscription lock poisoned");
decrement(&mut per_conn, conn_id, 1);
}
pub(crate) fn unsubscribe(&self, path: &str, conn_id: usize, filter: &str) {
let mut paths = self.paths.write().expect("subscription lock poisoned");
let mut per_conn = self.per_conn.write().expect("subscription lock poisoned");
let Some(state) = paths.get_mut(path) else {
return;
};
let Some(index) = state.subs.iter().position(|e| e.conn_id == conn_id) else {
return;
};
if !state.subs[index].filters.remove(filter) {
return;
}
decrement(&mut per_conn, conn_id, 1);
if state.subs[index].filters.is_empty() {
state.subs.swap_remove(index);
}
}
pub(crate) fn remove_connection(&self, conn_id: usize) {
let mut paths = self.paths.write().expect("subscription lock poisoned");
let mut per_conn = self.per_conn.write().expect("subscription lock poisoned");
for state in paths.values_mut() {
if let Some(index) = state.subs.iter().position(|e| e.conn_id == conn_id) {
state.subs.swap_remove(index);
}
}
per_conn.remove(&conn_id);
}
pub(crate) fn publish(
&self,
path: &str,
topic: &str,
payload: Bytes,
trace: Option<TraceContext>,
want: u32,
) -> usize {
let paths = self.paths.read().expect("subscription lock poisoned");
let Some(state) = paths.get(path) else {
return 0;
};
let topic: Arc<str> = Arc::from(topic);
let sequence = self.sequencer.next(&topic);
let mut sent = 0usize;
for entry in &state.subs {
if !entry.filters.iter().any(|f| filter::matches(&topic, f)) {
continue;
}
let Ok(permit) = entry.budget.try_acquire_many(want) else {
state.drops.record(&topic, DropCause::SubscriberBudget);
tracing::debug!(path, %topic, "subscriber budget exhausted; message dropped");
continue;
};
let msg = PubMsg {
topic: Arc::clone(&topic),
payload: payload.clone(),
trace,
sequence,
};
match entry.tx.try_send(PubItem::Whole(msg)) {
Ok(()) => {
permit.forget();
sent += 1;
}
Err(_) => {
state.drops.record(&topic, DropCause::SubscriberQueue);
tracing::debug!(path, %topic, "subscriber queue full; message dropped");
}
}
}
sent
}
pub(crate) fn open(&self, path: &str, topic: &str, trace: Option<TraceContext>) -> FanOut {
let paths = self.paths.read().expect("subscription lock poisoned");
let topic: Arc<str> = Arc::from(topic);
let id = self.next_stream.fetch_add(1, Ordering::Relaxed);
let Some(state) = paths.get(path) else {
return FanOut::empty(id, topic);
};
let sequence = self.sequencer.next(&topic);
let mut targets = Vec::new();
for entry in &state.subs {
if !entry.filters.iter().any(|f| filter::matches(&topic, f)) {
continue;
}
let Ok(ending) = entry.tx.clone().try_reserve_owned() else {
state.drops.record(&topic, DropCause::SubscriberQueue);
continue;
};
let head = PubMsg {
topic: Arc::clone(&topic),
payload: Bytes::new(),
trace,
sequence,
};
if entry.tx.try_send(PubItem::Begin { id, head }).is_err() {
state.drops.record(&topic, DropCause::SubscriberQueue);
continue;
}
targets.push(Target {
tx: entry.tx.clone(),
budget: Arc::clone(&entry.budget),
ending: Some(ending),
});
}
FanOut {
id,
topic,
targets,
drops: Some(Arc::clone(&state.drops)),
finished: false,
}
}
pub(crate) fn subscriber_count(&self, path: &str) -> usize {
self.paths
.read()
.expect("subscription lock poisoned")
.get(path)
.map_or(0, |s| s.subs.len())
}
pub(crate) fn filter_count(&self, path: &str) -> usize {
self.paths
.read()
.expect("subscription lock poisoned")
.get(path)
.map_or(0, |s| s.subs.iter().map(|e| e.filters.len()).sum())
}
pub(crate) fn dropped(&self, path: &str) -> u64 {
self.paths
.read()
.expect("subscription lock poisoned")
.get(path)
.map_or(0, |s| s.drops.total.load(Ordering::Relaxed))
}
pub(crate) fn dropped_on(&self, path: &str, topic: &str) -> Option<TopicDrops> {
self.paths
.read()
.expect("subscription lock poisoned")
.get(path)
.and_then(|s| s.drops.on_topic(topic))
}
pub(crate) fn drops(&self, path: &str) -> Vec<TopicDrops> {
self.paths
.read()
.expect("subscription lock poisoned")
.get(path)
.map_or_else(Vec::new, |s| s.drops.by_topic())
}
}
struct Target {
tx: mpsc::Sender<PubItem>,
budget: Arc<Semaphore>,
ending: Option<OwnedPermit<PubItem>>,
}
pub struct FanOut {
id: u64,
topic: Arc<str>,
targets: Vec<Target>,
drops: Option<Arc<DropTable>>,
finished: bool,
}
impl FanOut {
fn empty(id: u64, topic: Arc<str>) -> FanOut {
FanOut {
id,
topic,
targets: Vec::new(),
drops: None,
finished: false,
}
}
pub fn topic(&self) -> &str {
&self.topic
}
pub fn subscribers(&self) -> usize {
self.targets.len()
}
pub async fn write_within(
&mut self,
chunk: impl Into<Bytes>,
limit: std::time::Duration,
) -> Result<usize, Error> {
let chunk = chunk.into();
let want = u32::try_from(chunk.len()).map_err(|_| Error::LimitExceeded)?;
let mut kept = Vec::with_capacity(self.targets.len());
for mut target in std::mem::take(&mut self.targets) {
let budget = Arc::clone(&target.budget);
let acquired = tokio::select! {
permit = tokio::time::timeout(limit, budget.acquire_many_owned(want)) => {
match permit {
Ok(Ok(permit)) => Some(permit),
_ => None,
}
}
() = target.tx.closed() => None,
};
let Some(permit) = acquired else {
self.abort_one(&mut target, DropCause::SubscriberBudget);
continue;
};
match self.enqueue(&target, &chunk) {
Ok(()) => {
permit.forget();
kept.push(target);
}
Err(()) => {
drop(permit);
self.abort_one(&mut target, DropCause::SubscriberQueue);
}
}
}
self.targets = kept;
Ok(self.targets.len())
}
pub fn write_now(&mut self, chunk: impl Into<Bytes>) -> Result<usize, Error> {
let chunk = chunk.into();
let want = u32::try_from(chunk.len()).map_err(|_| Error::LimitExceeded)?;
let mut kept = Vec::with_capacity(self.targets.len());
for mut target in std::mem::take(&mut self.targets) {
let Ok(permit) = target.budget.try_acquire_many(want) else {
self.abort_one(&mut target, DropCause::SubscriberBudget);
continue;
};
match self.enqueue(&target, &chunk) {
Ok(()) => {
permit.forget();
kept.push(target);
}
Err(()) => {
drop(permit);
self.abort_one(&mut target, DropCause::SubscriberQueue);
}
}
}
self.targets = kept;
Ok(self.targets.len())
}
fn enqueue(&self, target: &Target, chunk: &Bytes) -> Result<(), ()> {
let item = PubItem::Chunk {
id: self.id,
payload: chunk.clone(),
};
target.tx.try_send(item).map_err(|_| ())
}
pub fn finish(mut self) -> usize {
self.finished = true;
let id = self.id;
let delivered = self.targets.len();
for target in &mut self.targets {
if let Some(ending) = target.ending.take() {
ending.send(PubItem::Finish { id });
}
}
delivered
}
fn abort_one(&self, target: &mut Target, cause: DropCause) {
if let Some(drops) = &self.drops {
drops.record(&self.topic, cause);
}
if let Some(ending) = target.ending.take() {
ending.send(PubItem::Abort { id: self.id });
}
tracing::debug!(topic = %self.topic, ?cause, "streamed fan-out copy aborted");
}
}
impl Drop for FanOut {
fn drop(&mut self) {
if self.finished {
return;
}
for target in &mut self.targets {
if let Some(ending) = target.ending.take() {
ending.send(PubItem::Abort { id: self.id });
}
}
}
}
fn decrement(counts: &mut HashMap<usize, usize>, conn_id: usize, by: usize) {
if let Some(count) = counts.get_mut(&conn_id) {
*count = count.saturating_sub(by);
if *count == 0 {
counts.remove(&conn_id);
}
}
}
async fn writer(
ctx: ConnHandle,
path: Arc<str>,
mut rx: mpsc::Receiver<PubItem>,
budget: Arc<Semaphore>,
drops: Arc<DropTable>,
) {
let mut streaming: HashMap<u64, SendHalf> = HashMap::new();
loop {
let item = tokio::select! {
item = rx.recv() => match item {
Some(item) => item,
None => break,
},
_ = ctx.conn.closed() => break,
};
match item {
PubItem::Whole(msg) => {
let len = msg.payload.len();
let outcome = write_one(&ctx, &path, &msg).await;
budget.add_permits(len);
match outcome {
Ok(()) => {}
Err(Error::NoParkedConnection) => {
drops.record(&msg.topic, DropCause::NoParkedConnection);
tracing::debug!(path = %path, "no parked connection; copy dropped");
}
Err(e) => {
tracing::debug!(path = %path, error = %e, "fan-out write failed; subscriber writer ending");
break;
}
}
}
PubItem::Begin { id, head } => match begin_one(&ctx, &path, &head).await {
Ok(stream) => {
streaming.insert(id, stream);
}
Err(Error::NoParkedConnection) => {
drops.record(&head.topic, DropCause::NoParkedConnection);
tracing::debug!(path = %path, "no parked connection; streamed copy dropped");
}
Err(e) => {
tracing::debug!(path = %path, error = %e, "fan-out open failed; subscriber writer ending");
break;
}
},
PubItem::Chunk { id, payload } => {
let len = payload.len();
if let Some(stream) = streaming.get_mut(&id)
&& let Err(e) = stream.write_all(&payload).await
{
tracing::debug!(path = %path, error = %e, "streamed fan-out write failed");
if let Some(mut stream) = streaming.remove(&id) {
stream.reset(codes::CANCELED);
}
}
budget.add_permits(len);
}
PubItem::Finish { id } => {
if let Some(mut stream) = streaming.remove(&id)
&& stream.finish().is_ok()
&& ctx.parked.park(stream.stopped())
{
ctx.shared.drain.evict();
}
}
PubItem::Abort { id } => {
if let Some(mut stream) = streaming.remove(&id) {
stream.reset(codes::CANCELED);
}
}
}
}
for (_, mut stream) in streaming {
stream.reset(codes::CANCELED);
}
tracing::debug!(path = %path, "subscriber writer ended");
}
async fn begin_one(ctx: &ConnHandle, path: &str, head: &PubMsg) -> Result<SendHalf, Error> {
let mut header = DataHeader::addressed(path);
header.topic = Some(head.topic.to_string());
header.traceparent = head.trace.map(|t| t.to_traceparent());
header.sequence = head.sequence;
let mut stream = ctx.open_uni().await?;
write_data_preamble(&mut stream, &header).await?;
Ok(stream)
}
async fn write_one(ctx: &ConnHandle, path: &str, msg: &PubMsg) -> Result<(), Error> {
let mut header = DataHeader::addressed(path);
header.topic = Some(msg.topic.to_string());
header.content_len = Some(msg.payload.len() as u64);
header.traceparent = msg.trace.map(|t| t.to_traceparent());
header.sequence = msg.sequence;
let mut stream = ctx.open_uni().await?;
write_data_preamble(&mut stream, &header).await?;
stream.write_all(&msg.payload).await?;
stream.finish()?;
if ctx.parked.park(stream.stopped()) {
ctx.shared.drain.evict();
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_literal_filter_matches_whole_segments_only() {
assert!(filter::matches("px.eur", "px.eur"));
assert!(!filter::matches("px.eur", "px.eur.spot"));
assert!(!filter::matches("fx.usd", "px.eur"));
assert!(!filter::matches("px.eur", "px."));
assert!(!filter::matches("sensors.temperature", "sensors.temp"));
}
#[test]
fn the_empty_filter_and_the_rest_wildcard_take_everything() {
assert!(filter::matches("", ""));
assert!(filter::matches("anything.at.all", ""));
assert!(filter::matches("anything.at.all", "#"));
assert!(filter::matches("", "#"));
}
#[test]
fn one_segment_wildcard_matches_exactly_one() {
assert!(filter::matches("px.eur", "px.*"));
assert!(filter::matches("sensors.a.temp", "sensors.*.temp"));
assert!(filter::matches("px.eur", "*.eur"));
assert!(!filter::matches("px", "px.*"));
assert!(!filter::matches("px.eur.spot", "px.*"));
}
#[test]
fn the_rest_wildcard_matches_zero_or_more_trailing_segments() {
assert!(filter::matches("px", "px.#"));
assert!(filter::matches("px.eur", "px.#"));
assert!(filter::matches("px.eur.spot", "px.#"));
assert!(!filter::matches("fx", "px.#"));
assert!(!filter::matches("pxx", "px.#"));
}
#[test]
fn a_topic_is_never_a_pattern() {
assert!(filter::matches("px.*", "px.*"));
assert!(!filter::matches("px.*", "px.eur"));
assert!(filter::matches("px.#", "px.#"));
assert!(filter::matches("px.*", "*.*"));
}
#[test]
fn empty_segments_match_only_empty_segments() {
assert!(filter::matches("px.", "px."));
assert!(!filter::matches("px.eur", "px."));
assert!(filter::matches("px.", "px.*"));
}
#[test]
fn matching_is_byte_exact_not_case_folded() {
assert!(!filter::matches("PX.EUR", "px.*"));
assert!(!filter::matches("px.eur", "PX.*"));
}
}