use radixdb_core::time_compat::Instant;
use std::borrow::Cow;
use std::sync::{Arc, RwLock};
use radixdb_core::SmartString;
use rustc_hash::FxHashMap;
use crate::context::ExecutionContext;
use radixdb_core::{Error, Result};
use radixdb_sql::ast::Statement;
pub use crate::compiled_plan::{
CompiledCountDistinct, CompiledCountStar, CompiledExecution, CompiledInsert, CompiledPkDelete,
CompiledPkLookup, CompiledPkUpdate, CompiledUpdateColumn, PkValueSource, UpdateValueSource,
};
#[inline]
fn to_lowercase_cow(s: &str) -> Cow<'_, str> {
if s.bytes().all(|b| !b.is_ascii_uppercase()) {
Cow::Borrowed(s)
} else {
Cow::Owned(s.to_lowercase())
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct ParameterContract {
positional_count: usize,
named_params: Arc<Vec<SmartString>>,
}
impl ParameterContract {
#[doc(hidden)]
pub fn from_statement(statement: &Statement) -> Self {
let mut positional_count = 0;
let mut named_params = Vec::new();
radixdb_sql::ast::walk_statement_tree(statement, &mut |expression| {
if let radixdb_sql::ast::Expression::Parameter(parameter) = expression {
if let Some(name) = parameter.name.strip_prefix(':') {
named_params.push(SmartString::new(name));
} else {
positional_count = positional_count.max(parameter.index);
}
}
});
named_params.sort_unstable();
named_params.dedup();
Self {
positional_count,
named_params: Arc::new(named_params),
}
}
#[doc(hidden)]
pub fn from_statements(statements: &[Statement]) -> Self {
let mut positional_count = 0;
let mut named_params = Vec::new();
for statement in statements {
let contract = Self::from_statement(statement);
positional_count = positional_count.max(contract.positional_count);
named_params.extend(contract.named_params.iter().cloned());
}
named_params.sort_unstable();
named_params.dedup();
Self {
positional_count,
named_params: Arc::new(named_params),
}
}
pub fn positional_count(&self) -> usize {
self.positional_count
}
pub fn named_params(&self) -> &[SmartString] {
&self.named_params
}
pub fn has_params(&self) -> bool {
self.positional_count != 0 || !self.named_params.is_empty()
}
#[doc(hidden)]
pub fn validate(&self, context: &ExecutionContext) -> Result<()> {
let provided_positional = context.params().len();
if provided_positional != self.positional_count {
return Err(Error::invalid_argument(format!(
"statement requires exactly {} positional parameters, got {}",
self.positional_count, provided_positional
)));
}
let provided_named = context.named_params();
let required_user_count = self
.named_params
.iter()
.filter(|name| !crate::context::is_system_context_name(name.as_str()))
.count();
let provided_user_count = provided_named
.keys()
.filter(|name| !crate::context::is_system_context_name(name))
.count();
if provided_user_count != required_user_count
|| self
.named_params
.iter()
.any(|name| !provided_named.contains_key(name.as_str()))
{
let required = self
.named_params
.iter()
.filter(|name| !crate::context::is_system_context_name(name.as_str()))
.map(SmartString::as_str)
.collect::<Vec<_>>()
.join(", ");
return Err(Error::invalid_argument(format!(
"statement requires exactly the named parameters [{required}]"
)));
}
Ok(())
}
}
pub const DEFAULT_CACHE_SIZE: usize = 1000;
#[derive(Debug, Clone)]
pub struct CachedPlanRef<B = ()> {
#[doc(hidden)]
pub statement: Arc<Statement>,
pub(crate) has_params: bool,
pub(crate) param_count: usize,
#[doc(hidden)]
pub parameter_contract: ParameterContract,
#[doc(hidden)]
pub compiled: Arc<RwLock<CompiledExecution>>,
#[doc(hidden)]
pub reference_expand: Arc<RwLock<B>>,
owner_token: Arc<()>,
}
impl<B> CachedPlanRef<B> {
pub fn statement(&self) -> &Statement {
&self.statement
}
pub fn has_params(&self) -> bool {
self.has_params
}
pub fn param_count(&self) -> usize {
self.param_count
}
pub fn parameter_contract(&self) -> &ParameterContract {
&self.parameter_contract
}
#[doc(hidden)]
pub fn compiled_state(&self) -> &Arc<RwLock<CompiledExecution>> {
&self.compiled
}
#[doc(hidden)]
pub fn binding_cache(&self) -> &Arc<RwLock<B>> {
&self.reference_expand
}
}
#[derive(Debug, Clone)]
pub struct CachedQueryPlan<B = ()> {
pub statement: Arc<Statement>,
pub query_text: SmartString,
pub last_used: Instant,
pub usage_count: u64,
pub has_params: bool,
pub param_count: usize,
pub parameter_contract: ParameterContract,
pub normalized_query: SmartString,
pub compiled: Arc<RwLock<CompiledExecution>>,
#[doc(hidden)]
pub reference_expand: Arc<RwLock<B>>,
}
impl<B: Default> CachedQueryPlan<B> {
pub fn new(
statement: Arc<Statement>,
query_text: SmartString,
_has_params: bool,
_param_count: usize,
normalized_query: SmartString,
) -> Self {
let parameter_contract = ParameterContract::from_statement(&statement);
let has_params = parameter_contract.has_params();
let param_count = parameter_contract.positional_count();
Self {
statement,
query_text,
last_used: Instant::now(),
usage_count: 1,
has_params,
param_count,
parameter_contract,
normalized_query,
compiled: Arc::new(RwLock::new(CompiledExecution::Unknown)),
reference_expand: Arc::new(RwLock::new(B::default())),
}
}
}
pub struct QueryCache<B = ()> {
plans: RwLock<FxHashMap<SmartString, CachedQueryPlan<B>>>,
max_size: usize,
prune_factor: f64,
owner_token: Arc<()>,
}
impl<B: Default> QueryCache<B> {
pub fn new(max_size: usize) -> Self {
Self {
plans: RwLock::new(FxHashMap::default()),
max_size,
prune_factor: 0.2, owner_token: Arc::new(()),
}
}
pub fn default_sized() -> Self {
Self::new(DEFAULT_CACHE_SIZE)
}
pub fn get(&self, query: &str) -> Option<CachedPlanRef<B>> {
let normalized = normalize_query(query);
let mut plans = self.plans.write().ok()?;
let plan = plans.get_mut(normalized.as_ref())?;
plan.last_used = Instant::now();
plan.usage_count = plan.usage_count.saturating_add(1);
Some(CachedPlanRef {
statement: plan.statement.clone(),
has_params: plan.has_params,
param_count: plan.param_count,
parameter_contract: plan.parameter_contract.clone(),
compiled: plan.compiled.clone(), reference_expand: plan.reference_expand.clone(),
owner_token: Arc::clone(&self.owner_token),
})
}
pub fn put(
&self,
query: &str,
statement: Arc<Statement>,
_has_params: bool,
_param_count: usize,
) -> CachedPlanRef<B> {
let parameter_contract = ParameterContract::from_statement(&statement);
let has_params = parameter_contract.has_params();
let param_count = parameter_contract.positional_count();
let normalized = normalize_query(query);
let normalized_key: SmartString = match normalized {
Cow::Borrowed(s) => SmartString::new(s),
Cow::Owned(s) => SmartString::new(&s),
};
let compiled = Arc::new(RwLock::new(CompiledExecution::Unknown));
let reference_expand = Arc::new(RwLock::new(B::default()));
if self.max_size > 0 {
if let Ok(mut plans) = self.plans.write() {
if plans.len() >= self.max_size {
self.prune_cache(&mut plans);
}
let key_for_insert = normalized_key.clone();
plans.insert(
key_for_insert,
CachedQueryPlan {
statement: statement.clone(),
query_text: SmartString::new(query),
last_used: Instant::now(),
usage_count: 1,
has_params,
param_count,
parameter_contract: parameter_contract.clone(),
normalized_query: normalized_key, compiled: compiled.clone(), reference_expand: reference_expand.clone(),
},
);
}
}
CachedPlanRef {
statement,
has_params,
param_count,
parameter_contract,
compiled,
reference_expand,
owner_token: Arc::clone(&self.owner_token),
}
}
#[doc(hidden)]
pub fn owns(&self, plan: &CachedPlanRef<B>) -> bool {
Arc::ptr_eq(&self.owner_token, &plan.owner_token)
}
pub fn clear(&self) {
if let Ok(mut plans) = self.plans.write() {
plans.clear();
}
}
pub fn invalidate_table(&self, table_name: &str) {
let table_lower = to_lowercase_cow(table_name);
if let Ok(mut plans) = self.plans.write() {
plans.retain(|_key, plan| {
if let Ok(compiled) = plan.compiled.read() {
match &*compiled {
CompiledExecution::PkLookup(lookup)
if lookup.table_name == *table_lower =>
{
return false; }
CompiledExecution::CountDistinct(cd) if cd.table_name == *table_lower => {
return false; }
CompiledExecution::CountStar(cs) if cs.table_name == *table_lower => {
return false; }
_ => {}
}
}
let query_lower = to_lowercase_cow(&plan.query_text);
!query_lower.contains(&format!(" {} ", &*table_lower))
&& !query_lower.contains(&format!(" {}\n", &*table_lower))
&& !query_lower.contains(&format!(" {};", &*table_lower))
&& !query_lower.contains(&format!("from {}", &*table_lower))
&& !query_lower.contains(&format!("join {}", &*table_lower))
&& !query_lower.contains(&format!("into {}", &*table_lower))
&& !query_lower.contains(&format!("update {}", &*table_lower))
});
}
}
pub fn size(&self) -> usize {
self.plans.read().map(|p| p.len()).unwrap_or(0)
}
pub fn stats(&self) -> CacheStats {
let plans = match self.plans.read() {
Ok(p) => p,
Err(_) => {
return CacheStats {
size: 0,
max_size: self.max_size,
total_usage: 0,
avg_usage: 0.0,
}
}
};
let size = plans.len();
let total_usage: u64 = plans.values().map(|p| p.usage_count).sum();
let avg_usage = if size > 0 {
total_usage as f64 / size as f64
} else {
0.0
};
CacheStats {
size,
max_size: self.max_size,
total_usage,
avg_usage,
}
}
fn prune_cache(&self, plans: &mut FxHashMap<SmartString, CachedQueryPlan<B>>) {
let num_to_remove = ((self.max_size as f64) * self.prune_factor).ceil() as usize;
let num_to_remove = num_to_remove.max(1);
if plans.is_empty() {
return;
}
let mut entries: Vec<(&SmartString, Instant, u64)> = plans
.iter()
.map(|(k, p)| (k, p.last_used, p.usage_count))
.collect();
entries.sort_unstable_by(|a, b| a.1.cmp(&b.1).then_with(|| a.2.cmp(&b.2)));
let keys_to_remove: Vec<SmartString> = entries
.into_iter()
.take(num_to_remove.min(plans.len()))
.map(|(k, _, _)| k.clone())
.collect();
for key in keys_to_remove {
plans.remove(&key);
}
}
}
impl<B: Default> Default for QueryCache<B> {
fn default() -> Self {
Self::default_sized()
}
}
#[derive(Debug, Clone)]
pub struct CacheStats {
pub size: usize,
pub max_size: usize,
pub total_usage: u64,
pub avg_usage: f64,
}
#[inline]
fn normalize_query(query: &str) -> std::borrow::Cow<'_, str> {
std::borrow::Cow::Borrowed(query)
}
#[cfg(test)]
mod tests {
use super::*;
use radixdb_sql::ast::{Expression, GroupByClause, SelectStatement, StarExpression};
use radixdb_sql::token::{Position, Token, TokenType};
fn dummy_token() -> Token {
Token::new(TokenType::Keyword, "SELECT", Position::new(0, 1, 1))
}
fn star_token() -> Token {
Token::new(TokenType::Operator, "*", Position::new(0, 1, 1))
}
fn create_test_statement() -> Arc<Statement> {
Arc::new(Statement::Select(SelectStatement {
token: dummy_token(),
with: None,
distinct: false,
distinct_on: vec![],
columns: vec![Expression::Star(StarExpression {
token: star_token(),
})],
table_expr: None,
where_clause: None,
group_by: GroupByClause::default(),
having: None,
window_defs: vec![],
order_by: vec![],
limit: None,
offset: None,
set_operations: vec![],
}))
}
#[test]
fn test_cache_put_get() {
let cache = QueryCache::<()>::new(100);
let stmt = create_test_statement();
cache.put("SELECT * FROM users", stmt.clone(), false, 0);
assert_eq!(cache.size(), 1);
let plan = cache.get("SELECT * FROM users");
assert!(plan.is_some());
let plan = plan.unwrap();
assert!(!plan.has_params);
assert_eq!(plan.param_count, 0);
}
#[test]
fn test_cache_miss() {
let cache = QueryCache::<()>::new(100);
let plan = cache.get("SELECT * FROM users");
assert!(plan.is_none());
}
#[test]
fn test_cache_usage_count() {
let cache = QueryCache::<()>::new(100);
let stmt = create_test_statement();
cache.put("SELECT * FROM users", stmt, false, 0);
for _ in 0..5 {
cache.get("SELECT * FROM users");
}
let stats = cache.stats();
assert_eq!(stats.total_usage, 6);
}
#[test]
fn r5_l03_cache_budgets_and_lru_follow_runtime_usage_query_plan() {
let cache = QueryCache::<()>::new(2);
let stmt = create_test_statement();
cache.put("SELECT 'a'", stmt.clone(), false, 0);
std::thread::sleep(std::time::Duration::from_millis(1));
cache.put("SELECT 'b'", stmt.clone(), false, 0);
assert!(cache.get("SELECT 'a'").is_some());
std::thread::sleep(std::time::Duration::from_millis(1));
cache.put("SELECT 'c'", stmt, false, 0);
assert!(cache.get("SELECT 'a'").is_some(), "hot plan was evicted");
assert!(cache.get("SELECT 'b'").is_none(), "cold plan survived");
assert!(cache.get("SELECT 'c'").is_some());
}
#[test]
fn test_cache_clear() {
let cache = QueryCache::<()>::new(100);
let stmt = create_test_statement();
cache.put("SELECT * FROM users", stmt, false, 0);
assert_eq!(cache.size(), 1);
cache.clear();
assert_eq!(cache.size(), 0);
}
#[test]
fn test_cache_pruning() {
let cache = QueryCache::<()>::new(5);
let stmt = create_test_statement();
for i in 0..10 {
let query = format!("SELECT * FROM table{}", i);
cache.put(&query, stmt.clone(), false, 0);
}
assert!(cache.size() <= 5);
}
#[test]
fn test_normalize_query() {
assert_eq!(
normalize_query(" SELECT * FROM users "),
" SELECT * FROM users "
);
assert_eq!(
normalize_query("SELECT\n*\nFROM\nusers"),
"SELECT\n*\nFROM\nusers"
);
assert_ne!(
normalize_query("SELECT 'a b'"),
normalize_query("SELECT 'a b'")
);
}
#[test]
fn test_normalize_query_utf8() {
assert_eq!(
normalize_query("SELECT * FROM t WHERE name = '日本語'"),
"SELECT * FROM t WHERE name = '日本語'"
);
assert_eq!(
normalize_query("SELECT * FROM t WHERE name = '日本語'"),
"SELECT * FROM t WHERE name = '日本語'"
);
assert_eq!(
normalize_query("SELECT\t*\tFROM t WHERE city = '東京' AND country = '中国'"),
"SELECT\t*\tFROM t WHERE city = '東京' AND country = '中国'"
);
assert_eq!(
normalize_query("SELECT * FROM t WHERE emoji = '🎉'"),
"SELECT * FROM t WHERE emoji = '🎉'"
);
}
#[test]
fn test_distinct_source_has_distinct_cache_key() {
let cache = QueryCache::<()>::new(100);
let stmt = create_test_statement();
cache.put("SELECT * FROM users", stmt, false, 0);
let plan = cache.get(" SELECT * FROM users ");
assert!(plan.is_none());
}
#[test]
fn test_parameterized_query() {
let cache = QueryCache::<()>::new(100);
let stmt = Arc::new(
radixdb_sql::parse_sql("SELECT * FROM users WHERE id = $1")
.expect("parse parameterized statement")
.into_iter()
.next()
.expect("one statement"),
);
cache.put("SELECT * FROM users WHERE id = $1", stmt, false, 0);
let plan = cache.get("SELECT * FROM users WHERE id = $1").unwrap();
assert!(plan.has_params);
assert_eq!(plan.param_count, 1);
}
#[test]
fn test_cache_stats() {
let cache = QueryCache::<()>::new(100);
let stmt = create_test_statement();
cache.put("SELECT 1", stmt.clone(), false, 0);
cache.put("SELECT 2", stmt.clone(), false, 0);
for _ in 0..5 {
cache.get("SELECT 1");
}
let stats = cache.stats();
assert_eq!(stats.size, 2);
assert_eq!(stats.max_size, 100);
assert_eq!(stats.total_usage, 7);
}
#[test]
fn v2_r5_zero_and_one_capacity_are_hard_bounds() {
let stmt = create_test_statement();
let disabled = QueryCache::<()>::new(0);
disabled.put("SELECT 1", stmt.clone(), false, 0);
assert_eq!(disabled.size(), 0);
assert!(disabled.get("SELECT 1").is_none());
let one = QueryCache::<()>::new(1);
one.put("SELECT 1", stmt.clone(), false, 0);
one.put("SELECT 2", stmt, false, 0);
assert_eq!(one.size(), 1);
assert!(one.get("SELECT 2").is_some());
}
#[test]
fn test_cache_thread_safety() {
use std::sync::Arc;
use std::thread;
let cache = Arc::new(QueryCache::<()>::new(1000));
let stmt = create_test_statement();
cache.put("SELECT * FROM users", stmt.clone(), false, 0);
let mut handles = vec![];
for _ in 0..10 {
let cache = Arc::clone(&cache);
handles.push(thread::spawn(move || {
for _ in 0..100 {
cache.get("SELECT * FROM users");
}
}));
}
for i in 0..5 {
let cache = Arc::clone(&cache);
let stmt = stmt.clone();
handles.push(thread::spawn(move || {
for j in 0..20 {
let query = format!("SELECT * FROM table{}_{}", i, j);
cache.put(&query, stmt.clone(), false, 0);
}
}));
}
for handle in handles {
handle.join().unwrap();
}
assert!(cache.get("SELECT * FROM users").is_some());
}
}