use std::{
collections::HashMap,
sync::{
Arc,
atomic::{AtomicU8, Ordering},
},
};
use arc_swap::ArcSwap;
use dashmap::DashMap;
use fraiseql_core::{runtime::Executor, security::ActorType};
use fraiseql_error::FraiseQLError;
use serde::{Deserialize, Serialize};
use tokio::sync::Semaphore;
#[cfg(feature = "auth")]
use crate::auth::rate_limiting::{AuthRateLimitConfig, Clock, KeyedRateLimiter};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum TenantStatus {
Active = 0,
Suspended = 1,
}
impl TenantStatus {
const fn from_u8(v: u8) -> Self {
match v {
1 => Self::Suspended,
_ => Self::Active,
}
}
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::Active => "active",
Self::Suspended => "suspended",
}
}
}
pub trait TenantStatusSource: Send + Sync {
fn is_suspended(&self, tenant_key: &str) -> bool;
}
impl TenantStatusSource for TenantExecutorRegistry {
fn is_suspended(&self, tenant_key: &str) -> bool {
matches!(self.tenant_status(tenant_key), Ok(TenantStatus::Suspended))
}
}
#[derive(Debug, Clone, Default, Deserialize)]
pub struct TenantQuota {
#[serde(default)]
pub max_requests_per_sec: Option<u32>,
#[serde(default)]
pub max_concurrent: Option<u32>,
#[serde(default)]
pub max_storage_bytes_advisory: Option<u64>,
#[serde(default)]
pub cost_budget: Option<usize>,
#[serde(default)]
pub cost_budget_per_minute: Option<usize>,
#[serde(default, skip_serializing_if = "HashMap::is_empty")]
pub cost_budget_per_actor: HashMap<ActorType, ActorCostBudget>,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ActorCostBudget {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub per_request: Option<usize>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub per_minute: Option<usize>,
}
impl ActorCostBudget {
#[must_use]
pub const fn is_empty(&self) -> bool {
self.per_request.is_none() && self.per_minute.is_none()
}
}
impl TenantQuota {
pub fn validate(&self) -> Result<(), String> {
for (actor, budget) in &self.cost_budget_per_actor {
if budget.is_empty() {
return Err(format!(
"cost_budget_per_actor['{}'] sets neither per_request nor per_minute, so it \
is a budget that can never refuse anything. Give it a number or remove it.",
actor.as_str()
));
}
}
Ok(())
}
}
struct CostWindow {
budget: u64,
state: std::sync::Mutex<(u64, u64)>,
}
const COST_WINDOW_SECS: u64 = 60;
impl CostWindow {
const fn new(budget: u64) -> Self {
Self {
budget,
state: std::sync::Mutex::new((0, 0)),
}
}
fn try_charge(&self, now_secs: u64, cost: u64) -> Result<(), u64> {
#[allow(clippy::unwrap_used)]
let mut state = self.state.lock().unwrap();
if now_secs >= state.0 + COST_WINDOW_SECS {
*state = (now_secs, 0);
}
if state.1.saturating_add(cost) > self.budget {
return Err((state.0 + COST_WINDOW_SECS).saturating_sub(now_secs));
}
state.1 = state.1.saturating_add(cost);
Ok(())
}
}
struct TenantEntry {
executor: Arc<ArcSwap<Executor>>,
status: AtomicU8,
concurrency: Option<Arc<Semaphore>>,
#[cfg(feature = "auth")]
rps: Option<Arc<KeyedRateLimiter>>,
cost_window: Option<CostWindow>,
cost_windows_by_actor: HashMap<ActorType, CostWindow>,
quota: TenantQuota,
}
impl TenantEntry {
fn new(executor: Arc<Executor>) -> Self {
Self {
executor: Arc::new(ArcSwap::from(executor)),
status: AtomicU8::new(TenantStatus::Active as u8),
concurrency: None,
#[cfg(feature = "auth")]
rps: None,
cost_window: None,
cost_windows_by_actor: HashMap::new(),
quota: TenantQuota::default(),
}
}
fn with_quota(mut self, quota: TenantQuota) -> Self {
self.concurrency = quota.max_concurrent.map(|n| Arc::new(Semaphore::new(n as usize)));
#[cfg(feature = "auth")]
{
self.rps = quota.max_requests_per_sec.map(|n| {
Arc::new(KeyedRateLimiter::new(AuthRateLimitConfig {
enabled: true,
max_requests: n,
window_secs: 1,
}))
});
}
#[cfg(not(feature = "auth"))]
if quota.max_requests_per_sec.is_some() {
tracing::warn!(
"Tenant quota sets `max_requests_per_sec`, but per-second rate limiting requires \
the `auth` feature; the limit will NOT be enforced in this build."
);
}
self.cost_window =
quota.cost_budget_per_minute.map(|budget| CostWindow::new(budget as u64));
self.cost_windows_by_actor = quota
.cost_budget_per_actor
.iter()
.filter_map(|(actor, budget)| {
budget.per_minute.map(|b| (*actor, CostWindow::new(b as u64)))
})
.collect();
self.quota = quota;
self
}
fn with_default_minute_budget(mut self, default: Option<u64>) -> Self {
if self.cost_window.is_none() {
self.cost_window = default.map(CostWindow::new);
}
self
}
fn status(&self) -> TenantStatus {
TenantStatus::from_u8(self.status.load(Ordering::Relaxed))
}
fn set_status(&self, status: TenantStatus) {
self.status.store(status as u8, Ordering::Relaxed);
}
}
const SUSPENDED_RETRY_AFTER_SECS: u64 = 60;
pub struct TenantExecutorRegistry {
default: Arc<ArcSwap<Executor>>,
tenants: DashMap<String, TenantEntry>,
default_minute_budget: Option<u64>,
}
impl TenantExecutorRegistry {
#[must_use]
pub fn new(default: Arc<ArcSwap<Executor>>) -> Self {
let default_minute_budget = default
.load()
.schema()
.security
.as_ref()
.and_then(|s| s.cost_budget.as_ref())
.and_then(|c| c.per_tenant_per_minute_default);
Self {
default,
tenants: DashMap::new(),
default_minute_budget,
}
}
pub fn executor_for(
&self,
tenant_key: Option<&str>,
) -> fraiseql_error::Result<arc_swap::Guard<Arc<Executor>>> {
match tenant_key {
None => Ok(self.default.load()),
Some(key) => {
let entry = self.tenants.get(key).ok_or_else(|| {
FraiseQLError::unauthorized(format!("Tenant '{key}' is not registered"))
})?;
self.require_active(key, entry.value())?;
Ok(entry.value().executor.load())
},
}
}
fn require_active(&self, key: &str, entry: &TenantEntry) -> fraiseql_error::Result<()> {
if entry.status() == TenantStatus::Suspended {
return Err(FraiseQLError::ServiceUnavailable {
message: format!("Tenant '{key}' is suspended"),
retry_after: Some(SUSPENDED_RETRY_AFTER_SECS),
});
}
Ok(())
}
pub fn executor_for_admin(
&self,
key: &str,
) -> fraiseql_error::Result<arc_swap::Guard<Arc<Executor>>> {
let entry = self.tenants.get(key).ok_or_else(|| {
FraiseQLError::unauthorized(format!("Tenant '{key}' is not registered"))
})?;
Ok(entry.value().executor.load())
}
pub fn upsert(&self, key: impl Into<String>, executor: Arc<Executor>) -> bool {
let key = key.into();
if let Some(existing) = self.tenants.get(&key) {
existing.value().executor.store(executor);
false
} else {
self.tenants.insert(
key,
TenantEntry::new(executor).with_default_minute_budget(self.default_minute_budget),
);
true
}
}
pub fn upsert_with_quota(
&self,
key: impl Into<String>,
executor: Arc<Executor>,
quota: TenantQuota,
) -> bool {
let key = key.into();
if let Some(existing) = self.tenants.get(&key) {
let prev_status = existing.value().status();
drop(existing);
self.tenants.remove(&key);
let entry = TenantEntry::new(executor)
.with_quota(quota)
.with_default_minute_budget(self.default_minute_budget);
entry.set_status(prev_status);
self.tenants.insert(key, entry);
false
} else {
self.tenants.insert(
key,
TenantEntry::new(executor)
.with_quota(quota)
.with_default_minute_budget(self.default_minute_budget),
);
true
}
}
pub fn try_acquire_concurrency(
&self,
key: &str,
) -> fraiseql_error::Result<Option<tokio::sync::OwnedSemaphorePermit>> {
let entry = self.tenants.get(key).ok_or_else(|| FraiseQLError::not_found("tenant", key))?;
if let Some(ref sem) = entry.value().concurrency {
match sem.clone().try_acquire_owned() {
Ok(permit) => Ok(Some(permit)),
Err(_) => Err(FraiseQLError::RateLimited {
message: format!(
"Tenant '{key}' concurrency limit reached (max {})",
entry.value().quota.max_concurrent.unwrap_or(0)
),
retry_after_secs: 1,
}),
}
} else {
Ok(None)
}
}
#[cfg(feature = "auth")]
pub fn try_acquire_rps(&self, key: &str) -> fraiseql_error::Result<()> {
let entry = self.tenants.get(key).ok_or_else(|| FraiseQLError::not_found("tenant", key))?;
match entry.value().rps {
Some(ref limiter) => rps_gate_check(
limiter.as_ref(),
key,
entry.value().quota.max_requests_per_sec.unwrap_or(0),
),
None => Ok(()),
}
}
#[must_use]
pub fn has_cost_budget(&self, key: &str) -> bool {
self.tenants.get(key).is_some_and(|e| {
e.value().quota.cost_budget.is_some()
|| e.value().cost_window.is_some()
|| !e.value().quota.cost_budget_per_actor.is_empty()
})
}
pub fn check_cost_budget(
&self,
key: &str,
actor: Option<ActorType>,
cost: usize,
) -> fraiseql_error::Result<()> {
let entry = self.tenants.get(key).ok_or_else(|| FraiseQLError::not_found("tenant", key))?;
let quota = &entry.value().quota;
let (budget, scope) = actor
.and_then(|a| {
quota.cost_budget_per_actor.get(&a).and_then(|b| b.per_request).map(|b| (b, a))
})
.map_or((quota.cost_budget, "per-request"), |(b, a)| (Some(b), a.as_str()));
match budget {
Some(budget) if cost > budget => Err(FraiseQLError::CostExceeded {
message: format!(
"Tenant '{key}' operation cost {cost} exceeds the {scope} cost budget of \
{budget}"
),
cost: cost as u64,
limit: budget as u64,
retry_after_secs: None,
}),
_ => Ok(()),
}
}
pub fn charge_cost_window(
&self,
key: &str,
actor: Option<ActorType>,
cost: usize,
) -> fraiseql_error::Result<()> {
let entry = self.tenants.get(key).ok_or_else(|| FraiseQLError::not_found("tenant", key))?;
let value = entry.value();
let window = actor
.and_then(|a| value.cost_windows_by_actor.get(&a))
.or(value.cost_window.as_ref());
let Some(window) = window else {
return Ok(());
};
let now_secs = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map_or(0, |d| d.as_secs());
window.try_charge(now_secs, cost as u64).map_err(|retry_after_secs| {
FraiseQLError::CostExceeded {
message: format!(
"Tenant '{key}' per-minute cost budget of {budget} is exhausted; retry in \
{retry_after_secs}s",
budget = window.budget
),
cost: cost as u64,
limit: window.budget,
retry_after_secs: Some(retry_after_secs),
}
})
}
pub fn tenant_quota(&self, key: &str) -> fraiseql_error::Result<TenantQuota> {
let entry = self.tenants.get(key).ok_or_else(|| FraiseQLError::not_found("tenant", key))?;
Ok(entry.value().quota.clone())
}
pub fn remove(&self, key: &str) -> fraiseql_error::Result<()> {
self.tenants
.remove(key)
.map(|_| ())
.ok_or_else(|| FraiseQLError::not_found("tenant", key))
}
pub fn suspend(&self, key: &str) -> fraiseql_error::Result<()> {
let entry = self.tenants.get(key).ok_or_else(|| FraiseQLError::not_found("tenant", key))?;
entry.value().set_status(TenantStatus::Suspended);
Ok(())
}
pub fn resume(&self, key: &str) -> fraiseql_error::Result<()> {
let entry = self.tenants.get(key).ok_or_else(|| FraiseQLError::not_found("tenant", key))?;
entry.value().set_status(TenantStatus::Active);
Ok(())
}
pub fn tenant_status(&self, key: &str) -> fraiseql_error::Result<TenantStatus> {
let entry = self.tenants.get(key).ok_or_else(|| FraiseQLError::not_found("tenant", key))?;
Ok(entry.value().status())
}
#[must_use]
pub fn tenant_keys(&self) -> Vec<String> {
self.tenants.iter().map(|e| e.key().clone()).collect()
}
#[must_use]
pub fn len(&self) -> usize {
self.tenants.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.tenants.is_empty()
}
#[must_use]
pub fn default_executor(&self) -> arc_swap::Guard<Arc<Executor>> {
self.default.load()
}
pub async fn health_check(&self, key: &str) -> fraiseql_error::Result<()> {
let entry = self.tenants.get(key).ok_or_else(|| FraiseQLError::not_found("tenant", key))?;
let executor = entry.value().executor.load();
executor.health_check().await
}
}
#[cfg(feature = "auth")]
fn rps_gate_check<C: Clock>(
limiter: &KeyedRateLimiter<C>,
key: &str,
max_per_sec: u32,
) -> fraiseql_error::Result<()> {
limiter.check(key).map_err(|_| FraiseQLError::RateLimited {
message: format!(
"Tenant '{key}' request-rate limit reached (max {max_per_sec} req/s)"
),
retry_after_secs: 1,
})
}
#[cfg(test)]
mod cost_window_tests {
#![allow(clippy::unwrap_used, clippy::panic)]
use super::{COST_WINDOW_SECS, CostWindow};
#[test]
fn admits_to_budget_then_rejects_without_charging() {
let w = CostWindow::new(1_000);
assert!(w.try_charge(1_000, 600).is_ok());
assert!(w.try_charge(1_010, 400).is_ok());
assert_eq!(w.try_charge(1_030, 1), Err(30));
assert!(w.try_charge(1_030, 0).is_ok());
}
#[test]
fn window_resets_after_sixty_seconds() {
let w = CostWindow::new(100);
assert!(w.try_charge(1_000, 100).is_ok());
assert_eq!(w.try_charge(1_059, 1), Err(1), "still inside the window");
assert!(
w.try_charge(1_000 + COST_WINDOW_SECS, 100).is_ok(),
"a fresh window admits a full budget"
);
}
}
#[cfg(all(test, feature = "auth"))]
mod rps_tests {
#![allow(clippy::unwrap_used, clippy::panic)]
use std::sync::{
Arc,
atomic::{AtomicU64, Ordering},
};
use fraiseql_error::FraiseQLError;
use super::rps_gate_check;
use crate::auth::rate_limiting::{AuthRateLimitConfig, KeyedRateLimiter};
const fn one_second_window(max: u32) -> AuthRateLimitConfig {
AuthRateLimitConfig {
enabled: true,
max_requests: max,
window_secs: 1,
}
}
#[test]
fn admits_up_to_limit_then_rejects_within_window() {
let now = Arc::new(AtomicU64::new(1_000));
let clock = move || now.load(Ordering::Relaxed);
let limiter = KeyedRateLimiter::with_clock(one_second_window(2), clock);
assert!(rps_gate_check(&limiter, "acme", 2).is_ok());
assert!(rps_gate_check(&limiter, "acme", 2).is_ok());
assert!(matches!(
rps_gate_check(&limiter, "acme", 2),
Err(FraiseQLError::RateLimited { .. })
));
}
#[test]
fn window_resets_on_next_second() {
let now = Arc::new(AtomicU64::new(1_000));
let clock = {
let n = now.clone();
move || n.load(Ordering::Relaxed)
};
let limiter = KeyedRateLimiter::with_clock(one_second_window(1), clock);
assert!(rps_gate_check(&limiter, "acme", 1).is_ok());
assert!(rps_gate_check(&limiter, "acme", 1).is_err());
now.store(1_001, Ordering::Relaxed); assert!(rps_gate_check(&limiter, "acme", 1).is_ok(), "window must reset");
}
}