use super::event::{AuditEvent, QueryLogEntry};
use parking_lot::RwLock;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use tokio::sync::broadcast;
use tracing::{debug, trace, warn};
const DEFAULT_CAPACITY: usize = 1024;
#[derive(Debug, Clone)]
pub enum BusEvent {
QueryLog(QueryLogEntry),
Security(AuditEvent),
}
#[derive(Debug, Default)]
pub struct EventBusStats {
pub events_published: AtomicU64,
pub events_dropped: AtomicU64,
pub active_subscribers: AtomicUsize,
pub peak_subscribers: AtomicUsize,
}
impl EventBusStats {
pub fn snapshot(&self) -> EventBusStatsSnapshot {
EventBusStatsSnapshot {
events_published: self.events_published.load(Ordering::Relaxed),
events_dropped: self.events_dropped.load(Ordering::Relaxed),
active_subscribers: self.active_subscribers.load(Ordering::Relaxed),
peak_subscribers: self.peak_subscribers.load(Ordering::Relaxed),
}
}
}
#[derive(Debug, Clone, Default, serde::Serialize)]
pub struct EventBusStatsSnapshot {
pub events_published: u64,
pub events_dropped: u64,
pub active_subscribers: usize,
pub peak_subscribers: usize,
}
#[derive(Debug)]
pub struct AuditEventBus {
query_tx: broadcast::Sender<QueryLogEntry>,
security_tx: broadcast::Sender<AuditEvent>,
stats: Arc<EventBusStats>,
capacity: usize,
}
impl AuditEventBus {
pub fn new() -> Self {
Self::with_capacity(DEFAULT_CAPACITY)
}
pub fn with_capacity(capacity: usize) -> Self {
let (query_tx, _) = broadcast::channel(capacity);
let (security_tx, _) = broadcast::channel(capacity);
debug!(capacity, "Created audit event bus");
Self {
query_tx,
security_tx,
stats: Arc::new(EventBusStats::default()),
capacity,
}
}
pub fn publish_query(&self, entry: QueryLogEntry) -> usize {
match self.query_tx.send(entry) {
Ok(count) => {
self.stats.events_published.fetch_add(1, Ordering::Relaxed);
trace!(subscribers = count, "Published query log entry");
count
}
Err(_) => {
trace!("No subscribers for query log event");
0
}
}
}
pub fn publish_security(&self, event: AuditEvent) -> usize {
match self.security_tx.send(event) {
Ok(count) => {
self.stats.events_published.fetch_add(1, Ordering::Relaxed);
trace!(subscribers = count, "Published security event");
count
}
Err(_) => {
trace!("No subscribers for security event");
0
}
}
}
pub fn subscribe_queries(&self) -> QueryLogSubscriber {
let rx = self.query_tx.subscribe();
let stats = Arc::clone(&self.stats);
let current = stats.active_subscribers.fetch_add(1, Ordering::Relaxed) + 1;
let peak = stats.peak_subscribers.load(Ordering::Relaxed);
if current > peak {
stats.peak_subscribers.store(current, Ordering::Relaxed);
}
debug!(active = current, "New query log subscriber");
QueryLogSubscriber { rx, stats }
}
pub fn subscribe_security(&self) -> SecurityEventSubscriber {
let rx = self.security_tx.subscribe();
let stats = Arc::clone(&self.stats);
let current = stats.active_subscribers.fetch_add(1, Ordering::Relaxed) + 1;
let peak = stats.peak_subscribers.load(Ordering::Relaxed);
if current > peak {
stats.peak_subscribers.store(current, Ordering::Relaxed);
}
debug!(active = current, "New security event subscriber");
SecurityEventSubscriber { rx, stats }
}
pub fn stats(&self) -> EventBusStatsSnapshot {
self.stats.snapshot()
}
pub fn query_subscriber_count(&self) -> usize {
self.query_tx.receiver_count()
}
pub fn security_subscriber_count(&self) -> usize {
self.security_tx.receiver_count()
}
pub fn capacity(&self) -> usize {
self.capacity
}
}
impl Default for AuditEventBus {
fn default() -> Self {
Self::new()
}
}
impl Clone for AuditEventBus {
fn clone(&self) -> Self {
Self {
query_tx: self.query_tx.clone(),
security_tx: self.security_tx.clone(),
stats: Arc::clone(&self.stats),
capacity: self.capacity,
}
}
}
pub struct QueryLogSubscriber {
rx: broadcast::Receiver<QueryLogEntry>,
stats: Arc<EventBusStats>,
}
impl QueryLogSubscriber {
pub async fn recv(&mut self) -> Option<QueryLogEntry> {
loop {
match self.rx.recv().await {
Ok(entry) => return Some(entry),
Err(broadcast::error::RecvError::Lagged(count)) => {
warn!(
lagged = count,
"Query log subscriber lagged, dropped events"
);
self.stats
.events_dropped
.fetch_add(count, Ordering::Relaxed);
}
Err(broadcast::error::RecvError::Closed) => {
debug!("Query log channel closed");
return None;
}
}
}
}
pub fn try_recv(&mut self) -> Option<QueryLogEntry> {
loop {
match self.rx.try_recv() {
Ok(entry) => return Some(entry),
Err(broadcast::error::TryRecvError::Lagged(count)) => {
self.stats
.events_dropped
.fetch_add(count, Ordering::Relaxed);
}
Err(_) => return None,
}
}
}
}
impl Drop for QueryLogSubscriber {
fn drop(&mut self) {
self.stats
.active_subscribers
.fetch_sub(1, Ordering::Relaxed);
debug!("Query log subscriber dropped");
}
}
pub struct SecurityEventSubscriber {
rx: broadcast::Receiver<AuditEvent>,
stats: Arc<EventBusStats>,
}
impl SecurityEventSubscriber {
pub async fn recv(&mut self) -> Option<AuditEvent> {
loop {
match self.rx.recv().await {
Ok(event) => return Some(event),
Err(broadcast::error::RecvError::Lagged(count)) => {
warn!(
lagged = count,
"Security event subscriber lagged, dropped events"
);
self.stats
.events_dropped
.fetch_add(count, Ordering::Relaxed);
}
Err(broadcast::error::RecvError::Closed) => {
debug!("Security event channel closed");
return None;
}
}
}
}
pub fn try_recv(&mut self) -> Option<AuditEvent> {
loop {
match self.rx.try_recv() {
Ok(event) => return Some(event),
Err(broadcast::error::TryRecvError::Lagged(count)) => {
self.stats
.events_dropped
.fetch_add(count, Ordering::Relaxed);
}
Err(_) => return None,
}
}
}
}
impl Drop for SecurityEventSubscriber {
fn drop(&mut self) {
self.stats
.active_subscribers
.fetch_sub(1, Ordering::Relaxed);
debug!("Security event subscriber dropped");
}
}
static EVENT_BUS: once_cell::sync::Lazy<RwLock<Option<AuditEventBus>>> =
once_cell::sync::Lazy::new(|| RwLock::new(None));
pub fn init_event_bus(capacity: usize) {
let mut bus = EVENT_BUS.write();
*bus = Some(AuditEventBus::with_capacity(capacity));
debug!(capacity, "Initialized global event bus");
}
pub fn event_bus() -> Option<AuditEventBus> {
EVENT_BUS.read().clone()
}
pub fn publish_query(entry: QueryLogEntry) -> usize {
if let Some(bus) = EVENT_BUS.read().as_ref() {
bus.publish_query(entry)
} else {
0
}
}
pub fn publish_security(event: AuditEvent) -> usize {
if let Some(bus) = EVENT_BUS.read().as_ref() {
bus.publish_security(event)
} else {
0
}
}
#[cfg(test)]
mod tests {
use super::*;
fn sample_query_entry() -> QueryLogEntry {
QueryLogEntry::new(
1234,
"udp",
"example.com".to_string(),
"A".to_string(),
"IN".to_string(),
)
}
#[tokio::test]
async fn test_publish_without_subscribers() {
let bus = AuditEventBus::new();
let count = bus.publish_query(sample_query_entry());
assert_eq!(count, 0);
}
#[tokio::test]
async fn test_single_subscriber() {
let bus = AuditEventBus::new();
let mut sub = bus.subscribe_queries();
let entry = sample_query_entry();
let count = bus.publish_query(entry.clone());
assert_eq!(count, 1);
let received = sub.recv().await.unwrap();
assert_eq!(received.qname, "example.com");
}
#[tokio::test]
async fn test_multiple_subscribers() {
let bus = AuditEventBus::new();
let mut sub1 = bus.subscribe_queries();
let mut sub2 = bus.subscribe_queries();
let mut sub3 = bus.subscribe_queries();
assert_eq!(bus.query_subscriber_count(), 3);
let entry = sample_query_entry();
let count = bus.publish_query(entry);
assert_eq!(count, 3);
assert!(sub1.recv().await.is_some());
assert!(sub2.recv().await.is_some());
assert!(sub3.recv().await.is_some());
}
#[tokio::test]
async fn test_backpressure_handling() {
let bus = AuditEventBus::with_capacity(2);
let mut sub = bus.subscribe_queries();
for i in 0..5 {
let mut entry = sample_query_entry();
entry.query_id = i;
bus.publish_query(entry);
}
let received = sub.recv().await;
assert!(received.is_some());
let stats = bus.stats();
assert!(stats.events_dropped > 0 || stats.events_published == 5);
}
#[tokio::test]
async fn test_subscriber_drop_updates_count() {
let bus = AuditEventBus::new();
{
let _sub1 = bus.subscribe_queries();
let _sub2 = bus.subscribe_queries();
assert_eq!(bus.query_subscriber_count(), 2);
}
assert_eq!(bus.query_subscriber_count(), 0);
}
#[tokio::test]
async fn test_stats_tracking() {
let bus = AuditEventBus::new();
let _sub = bus.subscribe_queries();
for _ in 0..10 {
bus.publish_query(sample_query_entry());
}
let stats = bus.stats();
assert_eq!(stats.events_published, 10);
assert_eq!(stats.active_subscribers, 1);
assert_eq!(stats.peak_subscribers, 1);
}
}