1use radixdb_core::time_compat::Instant;
40use std::borrow::Cow;
41use std::sync::{Arc, RwLock};
42
43use radixdb_core::SmartString;
44use rustc_hash::FxHashMap;
45
46use crate::context::ExecutionContext;
47use radixdb_core::{Error, Result};
48use radixdb_sql::ast::Statement;
49
50pub use crate::compiled_plan::{
51 CompiledCountDistinct, CompiledCountStar, CompiledExecution, CompiledInsert, CompiledPkDelete,
52 CompiledPkLookup, CompiledPkUpdate, CompiledUpdateColumn, PkValueSource, UpdateValueSource,
53};
54
55#[inline]
57fn to_lowercase_cow(s: &str) -> Cow<'_, str> {
58 if s.bytes().all(|b| !b.is_ascii_uppercase()) {
59 Cow::Borrowed(s)
60 } else {
61 Cow::Owned(s.to_lowercase())
62 }
63}
64
65#[derive(Debug, Clone, Default, PartialEq, Eq)]
67pub struct ParameterContract {
68 positional_count: usize,
69 named_params: Arc<Vec<SmartString>>,
70}
71
72impl ParameterContract {
73 #[doc(hidden)]
74 pub fn from_statement(statement: &Statement) -> Self {
75 let mut positional_count = 0;
76 let mut named_params = Vec::new();
77 radixdb_sql::ast::walk_statement_tree(statement, &mut |expression| {
78 if let radixdb_sql::ast::Expression::Parameter(parameter) = expression {
79 if let Some(name) = parameter.name.strip_prefix(':') {
80 named_params.push(SmartString::new(name));
81 } else {
82 positional_count = positional_count.max(parameter.index);
83 }
84 }
85 });
86 named_params.sort_unstable();
87 named_params.dedup();
88 Self {
89 positional_count,
90 named_params: Arc::new(named_params),
91 }
92 }
93
94 #[doc(hidden)]
95 pub fn from_statements(statements: &[Statement]) -> Self {
96 let mut positional_count = 0;
97 let mut named_params = Vec::new();
98 for statement in statements {
99 let contract = Self::from_statement(statement);
100 positional_count = positional_count.max(contract.positional_count);
101 named_params.extend(contract.named_params.iter().cloned());
102 }
103 named_params.sort_unstable();
104 named_params.dedup();
105 Self {
106 positional_count,
107 named_params: Arc::new(named_params),
108 }
109 }
110
111 pub fn positional_count(&self) -> usize {
113 self.positional_count
114 }
115
116 pub fn named_params(&self) -> &[SmartString] {
118 &self.named_params
119 }
120
121 pub fn has_params(&self) -> bool {
122 self.positional_count != 0 || !self.named_params.is_empty()
123 }
124
125 #[doc(hidden)]
126 pub fn validate(&self, context: &ExecutionContext) -> Result<()> {
127 let provided_positional = context.params().len();
128 if provided_positional != self.positional_count {
129 return Err(Error::invalid_argument(format!(
130 "statement requires exactly {} positional parameters, got {}",
131 self.positional_count, provided_positional
132 )));
133 }
134
135 let provided_named = context.named_params();
136 let required_user_count = self
137 .named_params
138 .iter()
139 .filter(|name| !crate::context::is_system_context_name(name.as_str()))
140 .count();
141 let provided_user_count = provided_named
142 .keys()
143 .filter(|name| !crate::context::is_system_context_name(name))
144 .count();
145 if provided_user_count != required_user_count
146 || self
147 .named_params
148 .iter()
149 .any(|name| !provided_named.contains_key(name.as_str()))
150 {
151 let required = self
152 .named_params
153 .iter()
154 .filter(|name| !crate::context::is_system_context_name(name.as_str()))
155 .map(SmartString::as_str)
156 .collect::<Vec<_>>()
157 .join(", ");
158 return Err(Error::invalid_argument(format!(
159 "statement requires exactly the named parameters [{required}]"
160 )));
161 }
162 Ok(())
163 }
164}
165
166pub const DEFAULT_CACHE_SIZE: usize = 1000;
168
169#[derive(Debug, Clone)]
172pub struct CachedPlanRef<B = ()> {
173 #[doc(hidden)]
175 pub statement: Arc<Statement>,
176 pub(crate) has_params: bool,
178 pub(crate) param_count: usize,
180 #[doc(hidden)]
182 pub parameter_contract: ParameterContract,
183 #[doc(hidden)]
185 pub compiled: Arc<RwLock<CompiledExecution>>,
186 #[doc(hidden)]
188 pub reference_expand: Arc<RwLock<B>>,
189 owner_token: Arc<()>,
190}
191
192impl<B> CachedPlanRef<B> {
193 pub fn statement(&self) -> &Statement {
195 &self.statement
196 }
197
198 pub fn has_params(&self) -> bool {
200 self.has_params
201 }
202
203 pub fn param_count(&self) -> usize {
205 self.param_count
206 }
207
208 pub fn parameter_contract(&self) -> &ParameterContract {
210 &self.parameter_contract
211 }
212
213 #[doc(hidden)]
215 pub fn compiled_state(&self) -> &Arc<RwLock<CompiledExecution>> {
216 &self.compiled
217 }
218
219 #[doc(hidden)]
221 pub fn binding_cache(&self) -> &Arc<RwLock<B>> {
222 &self.reference_expand
223 }
224}
225
226#[derive(Debug, Clone)]
228pub struct CachedQueryPlan<B = ()> {
229 pub statement: Arc<Statement>,
231 pub query_text: SmartString,
233 pub last_used: Instant,
235 pub usage_count: u64,
237 pub has_params: bool,
239 pub param_count: usize,
241 pub parameter_contract: ParameterContract,
243 pub normalized_query: SmartString,
245 pub compiled: Arc<RwLock<CompiledExecution>>,
247 #[doc(hidden)]
249 pub reference_expand: Arc<RwLock<B>>,
250}
251
252impl<B: Default> CachedQueryPlan<B> {
253 pub fn new(
255 statement: Arc<Statement>,
256 query_text: SmartString,
257 _has_params: bool,
258 _param_count: usize,
259 normalized_query: SmartString,
260 ) -> Self {
261 let parameter_contract = ParameterContract::from_statement(&statement);
262 let has_params = parameter_contract.has_params();
263 let param_count = parameter_contract.positional_count();
264 Self {
265 statement,
266 query_text,
267 last_used: Instant::now(),
268 usage_count: 1,
269 has_params,
270 param_count,
271 parameter_contract,
272 normalized_query,
273 compiled: Arc::new(RwLock::new(CompiledExecution::Unknown)),
274 reference_expand: Arc::new(RwLock::new(B::default())),
275 }
276 }
277}
278
279pub struct QueryCache<B = ()> {
284 plans: RwLock<FxHashMap<SmartString, CachedQueryPlan<B>>>,
286 max_size: usize,
288 prune_factor: f64,
290 owner_token: Arc<()>,
291}
292
293impl<B: Default> QueryCache<B> {
294 pub fn new(max_size: usize) -> Self {
296 Self {
297 plans: RwLock::new(FxHashMap::default()),
298 max_size,
299 prune_factor: 0.2, owner_token: Arc::new(()),
301 }
302 }
303
304 pub fn default_sized() -> Self {
306 Self::new(DEFAULT_CACHE_SIZE)
307 }
308
309 pub fn get(&self, query: &str) -> Option<CachedPlanRef<B>> {
316 let normalized = normalize_query(query);
317
318 let mut plans = self.plans.write().ok()?;
319 let plan = plans.get_mut(normalized.as_ref())?;
320 plan.last_used = Instant::now();
321 plan.usage_count = plan.usage_count.saturating_add(1);
322
323 Some(CachedPlanRef {
325 statement: plan.statement.clone(),
326 has_params: plan.has_params,
327 param_count: plan.param_count,
328 parameter_contract: plan.parameter_contract.clone(),
329 compiled: plan.compiled.clone(), reference_expand: plan.reference_expand.clone(),
331 owner_token: Arc::clone(&self.owner_token),
332 })
333 }
334
335 pub fn put(
341 &self,
342 query: &str,
343 statement: Arc<Statement>,
344 _has_params: bool,
345 _param_count: usize,
346 ) -> CachedPlanRef<B> {
347 let parameter_contract = ParameterContract::from_statement(&statement);
348 let has_params = parameter_contract.has_params();
349 let param_count = parameter_contract.positional_count();
350 let normalized = normalize_query(query);
351 let normalized_key: SmartString = match normalized {
353 Cow::Borrowed(s) => SmartString::new(s),
354 Cow::Owned(s) => SmartString::new(&s),
355 };
356
357 let compiled = Arc::new(RwLock::new(CompiledExecution::Unknown));
359 let reference_expand = Arc::new(RwLock::new(B::default()));
360
361 if self.max_size > 0 {
362 if let Ok(mut plans) = self.plans.write() {
363 if plans.len() >= self.max_size {
365 self.prune_cache(&mut plans);
366 }
367
368 let key_for_insert = normalized_key.clone();
371 plans.insert(
372 key_for_insert,
373 CachedQueryPlan {
374 statement: statement.clone(),
375 query_text: SmartString::new(query),
376 last_used: Instant::now(),
377 usage_count: 1,
378 has_params,
379 param_count,
380 parameter_contract: parameter_contract.clone(),
381 normalized_query: normalized_key, compiled: compiled.clone(), reference_expand: reference_expand.clone(),
384 },
385 );
386 }
387 }
388
389 CachedPlanRef {
391 statement,
392 has_params,
393 param_count,
394 parameter_contract,
395 compiled,
396 reference_expand,
397 owner_token: Arc::clone(&self.owner_token),
398 }
399 }
400
401 #[doc(hidden)]
402 pub fn owns(&self, plan: &CachedPlanRef<B>) -> bool {
403 Arc::ptr_eq(&self.owner_token, &plan.owner_token)
404 }
405
406 pub fn clear(&self) {
408 if let Ok(mut plans) = self.plans.write() {
409 plans.clear();
410 }
411 }
412
413 pub fn invalidate_table(&self, table_name: &str) {
416 let table_lower = to_lowercase_cow(table_name);
417 if let Ok(mut plans) = self.plans.write() {
418 plans.retain(|_key, plan| {
421 if let Ok(compiled) = plan.compiled.read() {
423 match &*compiled {
424 CompiledExecution::PkLookup(lookup)
425 if lookup.table_name == *table_lower =>
426 {
427 return false; }
429 CompiledExecution::CountDistinct(cd) if cd.table_name == *table_lower => {
430 return false; }
432 CompiledExecution::CountStar(cs) if cs.table_name == *table_lower => {
433 return false; }
435 _ => {}
436 }
437 }
438 let query_lower = to_lowercase_cow(&plan.query_text);
440 !query_lower.contains(&format!(" {} ", &*table_lower))
441 && !query_lower.contains(&format!(" {}\n", &*table_lower))
442 && !query_lower.contains(&format!(" {};", &*table_lower))
443 && !query_lower.contains(&format!("from {}", &*table_lower))
444 && !query_lower.contains(&format!("join {}", &*table_lower))
445 && !query_lower.contains(&format!("into {}", &*table_lower))
446 && !query_lower.contains(&format!("update {}", &*table_lower))
447 });
448 }
449 }
450
451 pub fn size(&self) -> usize {
453 self.plans.read().map(|p| p.len()).unwrap_or(0)
454 }
455
456 pub fn stats(&self) -> CacheStats {
458 let plans = match self.plans.read() {
459 Ok(p) => p,
460 Err(_) => {
461 return CacheStats {
462 size: 0,
463 max_size: self.max_size,
464 total_usage: 0,
465 avg_usage: 0.0,
466 }
467 }
468 };
469
470 let size = plans.len();
471 let total_usage: u64 = plans.values().map(|p| p.usage_count).sum();
472 let avg_usage = if size > 0 {
473 total_usage as f64 / size as f64
474 } else {
475 0.0
476 };
477
478 CacheStats {
479 size,
480 max_size: self.max_size,
481 total_usage,
482 avg_usage,
483 }
484 }
485
486 fn prune_cache(&self, plans: &mut FxHashMap<SmartString, CachedQueryPlan<B>>) {
488 let num_to_remove = ((self.max_size as f64) * self.prune_factor).ceil() as usize;
490 let num_to_remove = num_to_remove.max(1);
491
492 if plans.is_empty() {
493 return;
494 }
495
496 let mut entries: Vec<(&SmartString, Instant, u64)> = plans
499 .iter()
500 .map(|(k, p)| (k, p.last_used, p.usage_count))
501 .collect();
502
503 entries.sort_unstable_by(|a, b| a.1.cmp(&b.1).then_with(|| a.2.cmp(&b.2)));
505
506 let keys_to_remove: Vec<SmartString> = entries
508 .into_iter()
509 .take(num_to_remove.min(plans.len()))
510 .map(|(k, _, _)| k.clone())
511 .collect();
512
513 for key in keys_to_remove {
515 plans.remove(&key);
516 }
517 }
518}
519
520impl<B: Default> Default for QueryCache<B> {
521 fn default() -> Self {
522 Self::default_sized()
523 }
524}
525
526#[derive(Debug, Clone)]
528pub struct CacheStats {
529 pub size: usize,
531 pub max_size: usize,
533 pub total_usage: u64,
535 pub avg_usage: f64,
537}
538
539#[inline]
546fn normalize_query(query: &str) -> std::borrow::Cow<'_, str> {
547 std::borrow::Cow::Borrowed(query)
548}
549
550#[cfg(test)]
551mod tests {
552 use super::*;
553 use radixdb_sql::ast::{Expression, GroupByClause, SelectStatement, StarExpression};
554 use radixdb_sql::token::{Position, Token, TokenType};
555
556 fn dummy_token() -> Token {
557 Token::new(TokenType::Keyword, "SELECT", Position::new(0, 1, 1))
558 }
559
560 fn star_token() -> Token {
561 Token::new(TokenType::Operator, "*", Position::new(0, 1, 1))
562 }
563
564 fn create_test_statement() -> Arc<Statement> {
565 Arc::new(Statement::Select(SelectStatement {
566 token: dummy_token(),
567 with: None,
568 distinct: false,
569 distinct_on: vec![],
570 columns: vec![Expression::Star(StarExpression {
571 token: star_token(),
572 })],
573 table_expr: None,
574 where_clause: None,
575 group_by: GroupByClause::default(),
576 having: None,
577 window_defs: vec![],
578 order_by: vec![],
579 limit: None,
580 offset: None,
581 set_operations: vec![],
582 }))
583 }
584
585 #[test]
586 fn test_cache_put_get() {
587 let cache = QueryCache::<()>::new(100);
588 let stmt = create_test_statement();
589
590 cache.put("SELECT * FROM users", stmt.clone(), false, 0);
592 assert_eq!(cache.size(), 1);
593
594 let plan = cache.get("SELECT * FROM users");
596 assert!(plan.is_some());
597
598 let plan = plan.unwrap();
599 assert!(!plan.has_params);
600 assert_eq!(plan.param_count, 0);
601 }
602
603 #[test]
604 fn test_cache_miss() {
605 let cache = QueryCache::<()>::new(100);
606
607 let plan = cache.get("SELECT * FROM users");
608 assert!(plan.is_none());
609 }
610
611 #[test]
612 fn test_cache_usage_count() {
613 let cache = QueryCache::<()>::new(100);
614 let stmt = create_test_statement();
615
616 cache.put("SELECT * FROM users", stmt, false, 0);
617
618 for _ in 0..5 {
620 cache.get("SELECT * FROM users");
621 }
622
623 let stats = cache.stats();
624 assert_eq!(stats.total_usage, 6);
625 }
626
627 #[test]
628 fn r5_l03_cache_budgets_and_lru_follow_runtime_usage_query_plan() {
629 let cache = QueryCache::<()>::new(2);
630 let stmt = create_test_statement();
631 cache.put("SELECT 'a'", stmt.clone(), false, 0);
632 std::thread::sleep(std::time::Duration::from_millis(1));
633 cache.put("SELECT 'b'", stmt.clone(), false, 0);
634 assert!(cache.get("SELECT 'a'").is_some());
635 std::thread::sleep(std::time::Duration::from_millis(1));
636 cache.put("SELECT 'c'", stmt, false, 0);
637
638 assert!(cache.get("SELECT 'a'").is_some(), "hot plan was evicted");
639 assert!(cache.get("SELECT 'b'").is_none(), "cold plan survived");
640 assert!(cache.get("SELECT 'c'").is_some());
641 }
642
643 #[test]
644 fn test_cache_clear() {
645 let cache = QueryCache::<()>::new(100);
646 let stmt = create_test_statement();
647
648 cache.put("SELECT * FROM users", stmt, false, 0);
649 assert_eq!(cache.size(), 1);
650
651 cache.clear();
652 assert_eq!(cache.size(), 0);
653 }
654
655 #[test]
656 fn test_cache_pruning() {
657 let cache = QueryCache::<()>::new(5);
658 let stmt = create_test_statement();
659
660 for i in 0..10 {
662 let query = format!("SELECT * FROM table{}", i);
663 cache.put(&query, stmt.clone(), false, 0);
664 }
665
666 assert!(cache.size() <= 5);
668 }
669
670 #[test]
671 fn test_normalize_query() {
672 assert_eq!(
673 normalize_query(" SELECT * FROM users "),
674 " SELECT * FROM users "
675 );
676 assert_eq!(
677 normalize_query("SELECT\n*\nFROM\nusers"),
678 "SELECT\n*\nFROM\nusers"
679 );
680 assert_ne!(
681 normalize_query("SELECT 'a b'"),
682 normalize_query("SELECT 'a b'")
683 );
684 }
685
686 #[test]
687 fn test_normalize_query_utf8() {
688 assert_eq!(
690 normalize_query("SELECT * FROM t WHERE name = '日本語'"),
691 "SELECT * FROM t WHERE name = '日本語'"
692 );
693
694 assert_eq!(
696 normalize_query("SELECT * FROM t WHERE name = '日本語'"),
697 "SELECT * FROM t WHERE name = '日本語'"
698 );
699
700 assert_eq!(
702 normalize_query("SELECT\t*\tFROM t WHERE city = '東京' AND country = '中国'"),
703 "SELECT\t*\tFROM t WHERE city = '東京' AND country = '中国'"
704 );
705
706 assert_eq!(
708 normalize_query("SELECT * FROM t WHERE emoji = '🎉'"),
709 "SELECT * FROM t WHERE emoji = '🎉'"
710 );
711 }
712
713 #[test]
714 fn test_distinct_source_has_distinct_cache_key() {
715 let cache = QueryCache::<()>::new(100);
716 let stmt = create_test_statement();
717
718 cache.put("SELECT * FROM users", stmt, false, 0);
720
721 let plan = cache.get(" SELECT * FROM users ");
723 assert!(plan.is_none());
724 }
725
726 #[test]
727 fn test_parameterized_query() {
728 let cache = QueryCache::<()>::new(100);
729 let stmt = Arc::new(
730 radixdb_sql::parse_sql("SELECT * FROM users WHERE id = $1")
731 .expect("parse parameterized statement")
732 .into_iter()
733 .next()
734 .expect("one statement"),
735 );
736
737 cache.put("SELECT * FROM users WHERE id = $1", stmt, false, 0);
738
739 let plan = cache.get("SELECT * FROM users WHERE id = $1").unwrap();
740 assert!(plan.has_params);
741 assert_eq!(plan.param_count, 1);
742 }
743
744 #[test]
745 fn test_cache_stats() {
746 let cache = QueryCache::<()>::new(100);
747 let stmt = create_test_statement();
748
749 cache.put("SELECT 1", stmt.clone(), false, 0);
750 cache.put("SELECT 2", stmt.clone(), false, 0);
751
752 for _ in 0..5 {
754 cache.get("SELECT 1");
755 }
756
757 let stats = cache.stats();
758 assert_eq!(stats.size, 2);
759 assert_eq!(stats.max_size, 100);
760 assert_eq!(stats.total_usage, 7);
761 }
762
763 #[test]
764 fn v2_r5_zero_and_one_capacity_are_hard_bounds() {
765 let stmt = create_test_statement();
766 let disabled = QueryCache::<()>::new(0);
767 disabled.put("SELECT 1", stmt.clone(), false, 0);
768 assert_eq!(disabled.size(), 0);
769 assert!(disabled.get("SELECT 1").is_none());
770
771 let one = QueryCache::<()>::new(1);
772 one.put("SELECT 1", stmt.clone(), false, 0);
773 one.put("SELECT 2", stmt, false, 0);
774 assert_eq!(one.size(), 1);
775 assert!(one.get("SELECT 2").is_some());
776 }
777
778 #[test]
779 fn test_cache_thread_safety() {
780 use std::sync::Arc;
781 use std::thread;
782
783 let cache = Arc::new(QueryCache::<()>::new(1000));
784 let stmt = create_test_statement();
785
786 cache.put("SELECT * FROM users", stmt.clone(), false, 0);
788
789 let mut handles = vec![];
790
791 for _ in 0..10 {
793 let cache = Arc::clone(&cache);
794 handles.push(thread::spawn(move || {
795 for _ in 0..100 {
796 cache.get("SELECT * FROM users");
797 }
798 }));
799 }
800
801 for i in 0..5 {
803 let cache = Arc::clone(&cache);
804 let stmt = stmt.clone();
805 handles.push(thread::spawn(move || {
806 for j in 0..20 {
807 let query = format!("SELECT * FROM table{}_{}", i, j);
808 cache.put(&query, stmt.clone(), false, 0);
809 }
810 }));
811 }
812
813 for handle in handles {
814 handle.join().unwrap();
815 }
816
817 assert!(cache.get("SELECT * FROM users").is_some());
819 }
820}