use std::cell::RefCell;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use crate::error::{DbError, PrimaryCode};
use crate::DbResult;
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum Exceeded {
Rows,
Bytes,
Time,
Cancelled,
}
impl Exceeded {
pub fn name(self) -> &'static str {
match self {
Exceeded::Rows => "rows",
Exceeded::Bytes => "bytes",
Exceeded::Time => "time",
Exceeded::Cancelled => "cancelled",
}
}
}
#[derive(Clone, Debug)]
pub struct Limits {
pub rows: Option<u64>,
pub bytes: Option<u64>,
pub time: Option<Duration>,
}
impl Limits {
pub fn unbounded() -> Limits {
Limits {
rows: None,
bytes: None,
time: None,
}
}
pub fn served() -> Limits {
Limits {
rows: Some(10_000_000),
bytes: Some(256 * 1024 * 1024),
time: Some(Duration::from_secs(60)),
}
}
pub fn with_time(mut self, time: Option<Duration>) -> Limits {
self.time = time;
self
}
pub fn with_rows(mut self, rows: Option<u64>) -> Limits {
self.rows = rows;
self
}
pub fn with_bytes(mut self, bytes: Option<u64>) -> Limits {
self.bytes = bytes;
self
}
}
impl Default for Limits {
fn default() -> Limits {
Limits::unbounded()
}
}
#[derive(Debug)]
struct Spending {
limits: Limits,
deadline: Option<Instant>,
cancel: Arc<AtomicBool>,
rows: u64,
bytes: u64,
}
thread_local! {
static ACTIVE: RefCell<Option<Spending>> = const { RefCell::new(None) };
}
#[derive(Debug)]
pub struct Guard {
previous: Option<Spending>,
}
impl Drop for Guard {
fn drop(&mut self) {
let previous = self.previous.take();
ACTIVE.with(|held| {
*held.borrow_mut() = previous;
});
}
}
pub fn arm(limits: Limits, cancel: Arc<AtomicBool>) -> Guard {
cancel.store(false, Ordering::Relaxed);
arm_as_it_stands(limits, cancel)
}
pub fn arm_as_it_stands(limits: Limits, cancel: Arc<AtomicBool>) -> Guard {
let spending = Spending {
deadline: limits
.time
.and_then(|window| Instant::now().checked_add(window)),
limits,
cancel,
rows: 0,
bytes: 0,
};
let previous = ACTIVE.with(|held| held.borrow_mut().replace(spending));
Guard { previous }
}
pub fn armed() -> bool {
ACTIVE.with(|held| held.borrow().is_some())
}
pub fn check() -> DbResult<()> {
ACTIVE.with(|held| {
let Some(spending) = held
.borrow()
.as_ref()
.map(|spending| (spending.cancel.load(Ordering::Relaxed), spending.deadline))
else {
return Ok(());
};
let (cancelled, deadline) = spending;
if cancelled {
return Err(exceeded(Exceeded::Cancelled, 0, 0));
}
if let Some(deadline) = deadline {
if Instant::now() >= deadline {
return Err(exceeded(Exceeded::Time, 0, 0));
}
}
Ok(())
})
}
pub fn spend(rows: u64, bytes: u64) -> DbResult<()> {
ACTIVE.with(|held| {
let mut borrowed = held.borrow_mut();
let Some(spending) = borrowed.as_mut() else {
return Ok(());
};
spending.rows = spending.rows.saturating_add(rows);
spending.bytes = spending.bytes.saturating_add(bytes);
if let Some(most) = spending.limits.rows {
if spending.rows > most {
return Err(exceeded(Exceeded::Rows, spending.rows, most));
}
}
if let Some(most) = spending.limits.bytes {
if spending.bytes > most {
return Err(exceeded(Exceeded::Bytes, spending.bytes, most));
}
}
Ok(())
})
}
pub fn materialise(bytes: u64) -> DbResult<()> {
ACTIVE.with(|held| {
let mut borrowed = held.borrow_mut();
let Some(spending) = borrowed.as_mut() else {
return Ok(());
};
if spending.cancel.load(Ordering::Relaxed) {
return Err(exceeded(Exceeded::Cancelled, 0, 0));
}
if let Some(deadline) = spending.deadline {
if Instant::now() >= deadline {
return Err(exceeded(Exceeded::Time, 0, 0));
}
}
spending.bytes = spending.bytes.saturating_add(bytes);
if let Some(most) = spending.limits.bytes {
if spending.bytes > most {
return Err(exceeded(Exceeded::Bytes, spending.bytes, most));
}
}
Ok(())
})
}
pub fn spent() -> (u64, u64) {
ACTIVE.with(|held| {
held.borrow()
.as_ref()
.map(|spending| (spending.rows, spending.bytes))
.unwrap_or((0, 0))
})
}
fn exceeded(what: Exceeded, used: u64, most: u64) -> DbError {
let said = match what {
Exceeded::Cancelled => "this request was cancelled.".to_string(),
Exceeded::Time => "this request ran past the time it was allowed.".to_string(),
Exceeded::Rows => format!(
"this request produced {used} rows, past the {most} it was allowed. Ask for fewer \
with a WHERE clause or a LIMIT."
),
Exceeded::Bytes => format!(
"this request produced {used} bytes of row data, past the {most} it was allowed. Ask \
for fewer columns or fewer rows."
),
};
DbError::primary(PrimaryCode::Interrupt)
.with_message(said.clone())
.with_detail(format!("budget={} {said}", what.name()))
}
pub fn exceeded_kind(error: &DbError) -> Option<Exceeded> {
let detail = error.detail()?;
let rest = detail.strip_prefix("budget=")?;
let name = rest.split_whitespace().next()?;
match name {
"rows" => Some(Exceeded::Rows),
"bytes" => Some(Exceeded::Bytes),
"time" => Some(Exceeded::Time),
"cancelled" => Some(Exceeded::Cancelled),
_ => None,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn an_unarmed_thread_spends_nothing() {
assert!(!armed());
assert!(check().is_ok());
assert!(spend(u64::MAX, u64::MAX).is_ok());
assert_eq!(spent(), (0, 0));
}
#[test]
fn a_row_budget_refuses_and_says_which_one() {
let _guard = arm(
Limits::unbounded().with_rows(Some(10)),
Arc::new(AtomicBool::new(false)),
);
assert!(spend(9, 0).is_ok());
let refused = spend(2, 0).expect_err("the eleventh row is past ten");
assert_eq!(exceeded_kind(&refused), Some(Exceeded::Rows));
assert!(
refused.message().contains("11 rows"),
"{}",
refused.message()
);
}
#[test]
fn a_byte_budget_refuses_on_its_own() {
let _guard = arm(
Limits::unbounded().with_bytes(Some(100)),
Arc::new(AtomicBool::new(false)),
);
assert!(spend(1_000_000, 99).is_ok());
let refused = spend(0, 2).expect_err("101 bytes is past 100");
assert_eq!(exceeded_kind(&refused), Some(Exceeded::Bytes));
}
#[test]
fn a_passed_deadline_refuses() {
let _guard = arm(
Limits::unbounded().with_time(Some(Duration::from_millis(0))),
Arc::new(AtomicBool::new(false)),
);
std::thread::sleep(Duration::from_millis(2));
let refused = check().expect_err("the deadline has passed");
assert_eq!(exceeded_kind(&refused), Some(Exceeded::Time));
}
#[test]
fn a_cancel_from_another_thread_stops_it() {
let flag = Arc::new(AtomicBool::new(false));
let _guard = arm(Limits::unbounded(), Arc::clone(&flag));
assert!(check().is_ok());
let other = Arc::clone(&flag);
std::thread::spawn(move || other.store(true, Ordering::Relaxed))
.join()
.expect("the setter runs");
let refused = check().expect_err("a cancelled request stops");
assert_eq!(exceeded_kind(&refused), Some(Exceeded::Cancelled));
}
#[test]
fn a_stale_cancel_does_not_stop_the_next_request() {
let flag = Arc::new(AtomicBool::new(true));
let _guard = arm(Limits::unbounded(), Arc::clone(&flag));
assert!(check().is_ok(), "arming clears a flag from the last call");
}
#[test]
fn a_nested_budget_is_restored_when_it_ends() {
let outer = arm(
Limits::unbounded().with_rows(Some(5)),
Arc::new(AtomicBool::new(false)),
);
assert!(spend(4, 0).is_ok());
{
let _inner = arm(Limits::unbounded(), Arc::new(AtomicBool::new(false)));
assert!(spend(1_000, 0).is_ok(), "the inner budget is its own");
}
let refused = spend(2, 0).expect_err("the outer budget still counts its own four rows");
assert_eq!(exceeded_kind(&refused), Some(Exceeded::Rows));
drop(outer);
assert!(!armed());
}
#[test]
fn an_ordinary_failure_names_no_budget() {
assert_eq!(
exceeded_kind(&crate::error::refusal("something else")),
None
);
}
}