use std::collections::BTreeMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use bytes::Bytes;
use dashmap::DashMap;
use parking_lot::RwLock;
use crate::error::{Result, RiftError, TopicReject};
use crate::now_ms;
use crate::topic::profile::TopicProfile;
use crate::topic::retention::RetentionPolicy;
pub fn validate_name(name: &str) -> Result<()> {
if name.is_empty() {
return Err(RiftError::Topic(TopicReject::InvalidName(
"empty topic".into(),
)));
}
if name.len() > 256 {
return Err(RiftError::Topic(TopicReject::InvalidName(format!(
"name too long: {} > 256",
name.len()
))));
}
if name.starts_with('$') {
return Err(RiftError::Topic(TopicReject::InvalidName(format!(
"name starts with reserved '$' prefix: {}",
name
))));
}
if name.chars().any(|c| c.is_control()) {
return Err(RiftError::Topic(TopicReject::InvalidName(
"name contains control characters".into(),
)));
}
Ok(())
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct SubscriberId(pub u64);
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct LogEntry {
pub offset: i64,
pub publisher_session: Option<String>,
pub message_id: String,
pub class: String,
pub event: Option<String>,
pub payload: Bytes,
pub timestamp: i64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub appended_at: Option<i64>,
}
#[derive(Debug)]
pub struct TopicEntry {
pub name: String,
pub profile: RwLock<TopicProfile>,
pub closed: parking_lot::Mutex<bool>,
pub log: RwLock<Vec<LogEntry>>,
pub subscriber_count: AtomicU64,
pub publisher_count: AtomicU64,
pub latest_snapshot: RwLock<Option<LogEntry>>,
}
impl TopicEntry {
fn new(name: String, profile: TopicProfile) -> Self {
Self {
name,
profile: RwLock::new(profile),
closed: parking_lot::Mutex::new(false),
log: RwLock::new(Vec::new()),
subscriber_count: AtomicU64::new(0),
publisher_count: AtomicU64::new(0),
latest_snapshot: RwLock::new(None),
}
}
pub fn head_offset(&self) -> i64 {
self.log.read().last().map(|e| e.offset).unwrap_or(0)
}
pub fn can_subscribe(&self) -> bool {
let limit = self.profile.read().max_subscribers;
self.subscriber_count.load(Ordering::Relaxed) < limit as u64
}
pub fn can_publish(&self) -> bool {
let limit = self.profile.read().max_publishers;
self.publisher_count.load(Ordering::Relaxed) < limit as u64
}
pub fn inc_subscriber(&self) {
self.subscriber_count.fetch_add(1, Ordering::Relaxed);
}
pub fn dec_subscriber(&self) {
self.subscriber_count.fetch_sub(1, Ordering::Relaxed);
}
pub fn inc_publisher(&self) {
self.publisher_count.fetch_add(1, Ordering::Relaxed);
}
pub fn dec_publisher(&self) {
self.publisher_count.fetch_sub(1, Ordering::Relaxed);
}
pub fn try_inc_subscriber(&self) -> bool {
let limit = self.profile.read().max_subscribers as u64;
let mut cur = self.subscriber_count.load(Ordering::Relaxed);
loop {
if cur >= limit {
return false;
}
match self.subscriber_count.compare_exchange_weak(
cur,
cur + 1,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => return true,
Err(observed) => cur = observed,
}
}
}
pub fn try_inc_publisher(&self) -> bool {
let limit = self.profile.read().max_publishers as u64;
let mut cur = self.publisher_count.load(Ordering::Relaxed);
loop {
if cur >= limit {
return false;
}
match self.publisher_count.compare_exchange_weak(
cur,
cur + 1,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => return true,
Err(observed) => cur = observed,
}
}
}
pub fn append(&self, mut entry: LogEntry) {
let now = now_ms();
if entry.timestamp == 0 {
entry.timestamp = now;
}
entry.appended_at = Some(now);
let profile = self.profile.read().clone();
let mut log = self.log.write();
log.push(entry.clone());
match profile.retention {
RetentionPolicy::None => log.clear(),
RetentionPolicy::Count(n) => {
if log.len() > n {
let drop = log.len() - n;
log.drain(0..drop);
}
}
RetentionPolicy::Size(max_bytes) => {
let mut total: usize = log.iter().map(|e| e.payload.len()).sum();
let mut idx = 0;
while total > max_bytes && idx < log.len() {
total -= log[idx].payload.len();
idx += 1;
}
if idx > 0 {
log.drain(0..idx);
}
}
RetentionPolicy::Ttl(ttl) => {
let ttl_ms = ttl.as_millis() as i64;
log.retain(|e| {
let ts = e.appended_at.unwrap_or(e.timestamp);
now - ts <= ttl_ms
});
}
RetentionPolicy::Latest => {
log.retain(|e| e.offset == entry.offset);
}
RetentionPolicy::Durable => {
const DURABLE_MEMORY_CAP: usize = 10_000;
if log.len() > DURABLE_MEMORY_CAP {
let drop = log.len() - DURABLE_MEMORY_CAP;
log.drain(0..drop);
}
}
}
if profile.snapshot_enabled {
*self.latest_snapshot.write() = Some(entry);
}
}
pub fn range(&self, from: i64, to: i64) -> Vec<LogEntry> {
self.log
.read()
.iter()
.filter(|e| e.offset >= from && e.offset < to)
.cloned()
.collect()
}
pub fn snapshot(&self) -> Option<LogEntry> {
self.latest_snapshot.read().clone()
}
}
#[derive(Clone, Debug, Default)]
pub struct TopicStore {
inner: Arc<DashMap<String, Arc<TopicEntry>>>,
}
impl TopicStore {
pub fn new() -> Self {
Self::default()
}
pub fn get_or_create(
&self,
name: &str,
default_profile: TopicProfile,
) -> Result<Arc<TopicEntry>> {
validate_name(name)?;
Ok(self
.inner
.entry(name.to_string())
.or_insert_with(|| Arc::new(TopicEntry::new(name.to_string(), default_profile)))
.value()
.clone())
}
pub fn get(&self, name: &str) -> Option<Arc<TopicEntry>> {
self.inner.get(name).map(|e| e.clone())
}
pub fn exists(&self, name: &str) -> bool {
self.inner.contains_key(name)
}
pub fn remove(&self, name: &str) -> Option<Arc<TopicEntry>> {
self.inner.remove(name).map(|(_, e)| e)
}
pub fn names(&self) -> Vec<String> {
self.inner.iter().map(|kv| kv.key().clone()).collect()
}
pub fn stats(&self) -> BTreeMap<String, (u64, u64, i64)> {
self.inner
.iter()
.map(|kv| {
let e = kv.value();
(
kv.key().clone(),
(
e.subscriber_count.load(Ordering::Relaxed),
e.publisher_count.load(Ordering::Relaxed),
e.head_offset(),
),
)
})
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn sample_entry(offset: i64, payload: &[u8]) -> LogEntry {
LogEntry {
offset,
publisher_session: None,
message_id: format!("m-{offset}"),
class: "event".into(),
event: Some("e".into()),
payload: Bytes::copy_from_slice(payload),
timestamp: 0,
appended_at: None,
}
}
#[test]
fn name_validation() {
assert!(super::validate_name("room/1").is_ok());
assert!(super::validate_name("user/abc").is_ok());
assert!(super::validate_name("").is_err());
assert!(super::validate_name("$system").is_err());
assert!(super::validate_name(&"x".repeat(257)).is_err());
}
#[test]
fn get_or_create_idempotent() {
let store = TopicStore::new();
let a = store
.get_or_create("room/1", TopicProfile::default())
.unwrap();
let b = store
.get_or_create("room/1", TopicProfile::default())
.unwrap();
assert!(Arc::ptr_eq(&a, &b));
}
#[test]
fn append_count_retention() {
let store = TopicStore::new();
let entry = store
.get_or_create(
"t1",
TopicProfile {
retention: RetentionPolicy::Count(2),
..TopicProfile::default()
},
)
.unwrap();
for i in 1..=5 {
entry.append(sample_entry(i, b"x"));
}
let log = entry.log.read();
assert_eq!(log.len(), 2);
assert_eq!(log[0].offset, 4);
assert_eq!(log[1].offset, 5);
}
#[test]
fn append_latest_retention() {
let store = TopicStore::new();
let entry = store
.get_or_create(
"t1",
TopicProfile {
retention: RetentionPolicy::Latest,
..TopicProfile::default()
},
)
.unwrap();
entry.append(sample_entry(1, b"a"));
entry.append(sample_entry(2, b"b"));
let log = entry.log.read();
assert_eq!(log.len(), 1);
assert_eq!(log[0].offset, 2);
}
#[test]
fn range_query() {
let store = TopicStore::new();
let entry = store
.get_or_create(
"t1",
TopicProfile {
retention: RetentionPolicy::Count(100),
..TopicProfile::default()
},
)
.unwrap();
for i in 1..=5 {
entry.append(sample_entry(i, b"x"));
}
let got = entry.range(2, 4);
assert_eq!(got.iter().map(|e| e.offset).collect::<Vec<_>>(), vec![2, 3]);
let got_inclusive_end = entry.range(2, 5);
assert_eq!(
got_inclusive_end
.iter()
.map(|e| e.offset)
.collect::<Vec<_>>(),
vec![2, 3, 4]
);
}
#[test]
fn snapshot_keeps_latest() {
let store = TopicStore::new();
let entry = store
.get_or_create(
"t1",
TopicProfile {
retention: RetentionPolicy::Count(10),
snapshot_enabled: true,
..TopicProfile::default()
},
)
.unwrap();
entry.append(sample_entry(1, b"a"));
entry.append(sample_entry(2, b"b"));
let s = entry.snapshot().unwrap();
assert_eq!(s.offset, 2);
}
#[test]
fn head_offset_reflects_log() {
let store = TopicStore::new();
let entry = store
.get_or_create(
"t1",
TopicProfile {
retention: RetentionPolicy::Count(100),
..TopicProfile::default()
},
)
.unwrap();
assert_eq!(entry.head_offset(), 0);
entry.append(sample_entry(1, b"a"));
assert_eq!(entry.head_offset(), 1);
entry.append(sample_entry(5, b"b"));
assert_eq!(entry.head_offset(), 5);
}
#[test]
fn subscriber_limit_check() {
let store = TopicStore::new();
let entry = store
.get_or_create(
"t1",
TopicProfile {
max_subscribers: 2,
..TopicProfile::default()
},
)
.unwrap();
assert!(entry.can_subscribe());
entry.inc_subscriber();
entry.inc_subscriber();
assert!(!entry.can_subscribe());
}
#[test]
fn publisher_limit_check() {
let store = TopicStore::new();
let entry = store
.get_or_create(
"t1",
TopicProfile {
max_publishers: 1,
..TopicProfile::default()
},
)
.unwrap();
assert!(entry.can_publish());
entry.inc_publisher();
assert!(!entry.can_publish());
}
}