use std::{
collections::HashMap,
sync::{
Arc, OnceLock,
atomic::{AtomicU64, Ordering},
},
};
use dashmap::DashMap;
use serde::Serialize;
use super::events::MutationAuditEvent;
static GLOBAL_USAGE_AGGREGATOR: OnceLock<Arc<UsageAggregator>> = OnceLock::new();
#[must_use]
pub fn global_aggregator() -> &'static Arc<UsageAggregator> {
GLOBAL_USAGE_AGGREGATOR.get_or_init(|| Arc::new(UsageAggregator::new()))
}
#[must_use]
pub fn validate_period(period: &str) -> bool {
let bytes = period.as_bytes();
if bytes.len() != 7 || bytes[4] != b'-' {
return false;
}
let year_str = &period[..4];
let month_str = &period[5..];
if !year_str.bytes().all(|b| b.is_ascii_digit()) {
return false;
}
if !month_str.bytes().all(|b| b.is_ascii_digit()) {
return false;
}
let month: u8 = month_str.parse().unwrap_or(0);
(1..=12).contains(&month)
}
#[non_exhaustive]
#[derive(Debug, Clone, Serialize)]
pub struct UsageSummary {
pub mutations: HashMap<String, u64>,
}
pub struct UsageAggregator {
counters: DashMap<(String, String, String), AtomicU64>,
backend: std::sync::RwLock<std::sync::Arc<dyn UsageBackend>>,
}
impl std::fmt::Debug for UsageAggregator {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("UsageAggregator")
.field("entry_count", &self.counters.len())
.finish_non_exhaustive()
}
}
impl UsageAggregator {
#[must_use]
pub fn new() -> Self {
Self {
counters: DashMap::new(),
backend: std::sync::RwLock::new(std::sync::Arc::new(NoopBackend)),
}
}
#[must_use]
pub fn new_with_backend(backend: std::sync::Arc<dyn UsageBackend>) -> Self {
Self {
counters: DashMap::new(),
backend: std::sync::RwLock::new(backend),
}
}
pub fn set_backend(&self, backend: std::sync::Arc<dyn UsageBackend>) {
*self.backend.write().expect("backend lock poisoned") = backend;
}
pub fn record(&self, event: &MutationAuditEvent) {
let key = (event.tenant_id.clone(), event.period.clone(), event.entity_type.clone());
self.counters
.entry(key)
.or_insert_with(|| AtomicU64::new(0))
.fetch_add(1, Ordering::Relaxed);
}
pub fn query(&self, tenant_id: &str, period: &str) -> UsageSummary {
let mut mutations: HashMap<String, u64> = HashMap::new();
for entry in &self.counters {
let (t, p, e) = entry.key();
if t == tenant_id && p == period {
mutations.insert(e.clone(), entry.value().load(Ordering::Relaxed));
}
}
UsageSummary { mutations }
}
#[must_use]
pub fn entry_count(&self) -> usize {
self.counters.len()
}
pub async fn flush_to_backend(&self) -> Result<(), String> {
let snapshot: HashMap<(String, String, String), u64> = self
.counters
.iter()
.map(|entry| (entry.key().clone(), entry.value().load(Ordering::Relaxed)))
.collect();
let backend = self.backend.read().expect("backend lock poisoned").clone();
backend.flush(&snapshot).await
}
pub async fn load_from_backend(&self) -> Result<(), String> {
let backend = self.backend.read().expect("backend lock poisoned").clone();
let persisted = backend.load().await?;
for (key, count) in persisted {
self.counters
.entry(key)
.or_insert_with(|| AtomicU64::new(0))
.fetch_add(count, Ordering::Relaxed);
}
Ok(())
}
}
impl Default for UsageAggregator {
fn default() -> Self {
Self::new()
}
}
#[async_trait::async_trait]
pub trait UsageBackend: Send + Sync {
async fn flush(
&self,
counters: &std::collections::HashMap<(String, String, String), u64>,
) -> Result<(), String>;
async fn load(
&self,
) -> Result<std::collections::HashMap<(String, String, String), u64>, String>;
}
#[derive(Debug, Default)]
pub struct NoopBackend;
#[cfg(feature = "redis-usage")]
#[derive(Debug, Clone)]
pub struct RedisBackend {
client: ::redis::aio::ConnectionManager,
}
#[cfg(feature = "redis-usage")]
impl RedisBackend {
#[must_use]
pub const fn new(client: ::redis::aio::ConnectionManager) -> Self {
Self { client }
}
fn redis_key(tenant_id: &str, period: &str) -> String {
format!("fraiseql:usage:{tenant_id}:{period}")
}
}
#[cfg(feature = "redis-usage")]
#[async_trait::async_trait]
impl UsageBackend for RedisBackend {
async fn flush(
&self,
counters: &std::collections::HashMap<(String, String, String), u64>,
) -> Result<(), String> {
use ::redis::AsyncCommands as _;
let mut grouped: std::collections::HashMap<String, Vec<(&str, u64)>> =
std::collections::HashMap::new();
for ((tenant, period, entity), &count) in counters {
let key = Self::redis_key(tenant, period);
grouped.entry(key).or_default().push((entity.as_str(), count));
}
let mut conn = self.client.clone();
for (key, fields) in &grouped {
if !fields.is_empty() {
conn.hset_multiple::<_, _, _, ()>(key, fields.as_slice())
.await
.map_err(|e| format!("Redis flush error: {e}"))?;
}
}
Ok(())
}
async fn load(
&self,
) -> Result<std::collections::HashMap<(String, String, String), u64>, String> {
use ::redis::AsyncCommands as _;
let mut conn = self.client.clone();
let mut result = std::collections::HashMap::new();
let keys: Vec<String> = conn
.keys("fraiseql:usage:*")
.await
.map_err(|e| format!("Redis load scan error: {e}"))?;
for key in &keys {
let parts: Vec<&str> = key.splitn(4, ':').collect();
if parts.len() != 4 {
continue;
}
let tenant = parts[2].to_owned();
let period = parts[3].to_owned();
let hash: std::collections::HashMap<String, u64> = conn
.hgetall(key)
.await
.map_err(|e| format!("Redis load hgetall error for {key}: {e}"))?;
for (entity, count) in hash {
result.insert((tenant.clone(), period.clone(), entity), count);
}
}
Ok(result)
}
}
#[async_trait::async_trait]
impl UsageBackend for NoopBackend {
async fn flush(
&self,
_counters: &std::collections::HashMap<(String, String, String), u64>,
) -> Result<(), String> {
Ok(())
}
async fn load(
&self,
) -> Result<std::collections::HashMap<(String, String, String), u64>, String> {
Ok(std::collections::HashMap::new())
}
}
#[derive(Debug, Clone)]
pub struct PostgresBackend {
pool: sqlx::PgPool,
}
impl PostgresBackend {
pub async fn new(pool: sqlx::PgPool) -> Result<Self, String> {
sqlx::query(
"CREATE TABLE IF NOT EXISTS fraiseql_usage_counters (
tenant_id TEXT NOT NULL,
period TEXT NOT NULL,
entity_type TEXT NOT NULL,
count BIGINT NOT NULL DEFAULT 0,
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
PRIMARY KEY (tenant_id, period, entity_type)
)",
)
.execute(&pool)
.await
.map_err(|e| format!("PostgresBackend schema migration failed: {e}"))?;
Ok(Self { pool })
}
}
#[async_trait::async_trait]
impl UsageBackend for PostgresBackend {
async fn flush(
&self,
counters: &std::collections::HashMap<(String, String, String), u64>,
) -> Result<(), String> {
if counters.is_empty() {
return Ok(());
}
for ((tenant_id, period, entity_type), &count) in counters {
sqlx::query(
"INSERT INTO fraiseql_usage_counters
(tenant_id, period, entity_type, count, updated_at)
VALUES ($1, $2, $3, $4, NOW())
ON CONFLICT (tenant_id, period, entity_type)
DO UPDATE SET count = EXCLUDED.count, updated_at = NOW()",
)
.bind(tenant_id)
.bind(period)
.bind(entity_type)
.bind(count.cast_signed())
.execute(&self.pool)
.await
.map_err(|e| format!("PostgresBackend flush error: {e}"))?;
}
Ok(())
}
async fn load(
&self,
) -> Result<std::collections::HashMap<(String, String, String), u64>, String> {
let rows: Vec<(String, String, String, i64)> = sqlx::query_as(
"SELECT tenant_id, period, entity_type, count
FROM fraiseql_usage_counters",
)
.fetch_all(&self.pool)
.await
.map_err(|e| format!("PostgresBackend load error: {e}"))?;
let result = rows
.into_iter()
.map(|(tenant_id, period, entity_type, count)| {
((tenant_id, period, entity_type), count.max(0).cast_unsigned())
})
.collect();
Ok(result)
}
}