#![allow(clippy::cast_possible_truncation)]
use std::sync::atomic::{AtomicU64, AtomicU8, AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use super::limits::{GuardRailViolation, QueryLimits};
#[derive(Debug)]
pub struct QueryContext {
pub limits: QueryLimits,
start_time: Instant,
current_depth: AtomicU64,
current_cardinality: AtomicUsize,
memory_used: AtomicUsize,
traversal_nodes_visited: AtomicU64,
traversal_edges_traversed: AtomicU64,
executed_filter_strategy: Arc<ExecutedStrategyCell>,
}
#[derive(Debug, Default)]
pub(crate) struct ExecutedStrategyCell(AtomicU8);
impl ExecutedStrategyCell {
pub(crate) fn record(&self, strategy: crate::velesql::FilterStrategy) {
use crate::velesql::FilterStrategy as F;
let raw = match strategy {
F::None => 0,
F::PreFilter => 1,
F::PreFilterExact => 2,
F::PostFilter => 3,
};
self.0.store(raw, Ordering::Relaxed);
}
pub(crate) fn get(&self) -> Option<crate::velesql::FilterStrategy> {
use crate::velesql::FilterStrategy as F;
match self.0.load(Ordering::Relaxed) {
1 => Some(F::PreFilter),
2 => Some(F::PreFilterExact),
3 => Some(F::PostFilter),
_ => None,
}
}
}
impl QueryContext {
#[must_use]
pub fn new(limits: QueryLimits) -> Self {
Self {
limits,
start_time: Instant::now(),
current_depth: AtomicU64::new(0),
current_cardinality: AtomicUsize::new(0),
memory_used: AtomicUsize::new(0),
traversal_nodes_visited: AtomicU64::new(0),
traversal_edges_traversed: AtomicU64::new(0),
executed_filter_strategy: Arc::new(ExecutedStrategyCell::default()),
}
}
#[must_use]
pub(crate) fn executed_strategy_slot(&self) -> Arc<ExecutedStrategyCell> {
Arc::clone(&self.executed_filter_strategy)
}
#[must_use]
pub(crate) fn executed_filter_strategy(&self) -> Option<crate::velesql::FilterStrategy> {
self.executed_filter_strategy.get()
}
pub fn check_timeout(&self) -> Result<(), GuardRailViolation> {
if self.limits.timeout_ms == 0 {
return Ok(());
}
let elapsed_ms = self.start_time.elapsed().as_millis() as u64;
if elapsed_ms >= self.limits.timeout_ms {
return Err(GuardRailViolation::Timeout {
max_ms: self.limits.timeout_ms,
elapsed_ms,
});
}
Ok(())
}
pub fn check_depth(&self, depth: u32) -> Result<(), GuardRailViolation> {
self.current_depth
.store(u64::from(depth), Ordering::Relaxed);
if depth > self.limits.max_depth {
return Err(GuardRailViolation::DepthExceeded {
max: self.limits.max_depth,
actual: depth,
});
}
Ok(())
}
pub fn check_cardinality(&self, count: usize) -> Result<(), GuardRailViolation> {
let current = self.current_cardinality.fetch_add(count, Ordering::Relaxed) + count;
if current > self.limits.max_cardinality {
return Err(GuardRailViolation::CardinalityExceeded {
max: self.limits.max_cardinality,
actual: current,
});
}
Ok(())
}
pub fn check_memory(&self, bytes: usize) -> Result<(), GuardRailViolation> {
let current = self.memory_used.fetch_add(bytes, Ordering::Relaxed) + bytes;
if current > self.limits.memory_limit_bytes {
return Err(GuardRailViolation::MemoryExceeded {
max_bytes: self.limits.memory_limit_bytes,
used_bytes: current,
});
}
Ok(())
}
#[must_use]
pub fn elapsed(&self) -> Duration {
self.start_time.elapsed()
}
#[must_use]
pub fn memory_used(&self) -> usize {
self.memory_used.load(Ordering::Relaxed)
}
pub fn add_traversal(&self, nodes_visited: u64, edges_traversed: u64) {
self.traversal_nodes_visited
.fetch_add(nodes_visited, Ordering::Relaxed);
self.traversal_edges_traversed
.fetch_add(edges_traversed, Ordering::Relaxed);
}
#[must_use]
pub fn traversal_nodes_visited(&self) -> u64 {
self.traversal_nodes_visited.load(Ordering::Relaxed)
}
#[must_use]
pub fn traversal_edges_traversed(&self) -> u64 {
self.traversal_edges_traversed.load(Ordering::Relaxed)
}
}