1use std::collections::HashMap;
15use std::sync::atomic::{AtomicU64, Ordering};
16use std::sync::Arc;
17use std::time::{Duration, Instant};
18
19use parking_lot::RwLock;
20use sqlparser::ast::Statement;
21use sqlparser::dialect::GenericDialect;
22use sqlparser::parser::Parser;
23use std::hash::Hasher;
24use twox_hash::XxHash64;
25
26pub struct SqlNormalizer;
33
34impl SqlNormalizer {
35 pub fn normalize(sql: &str) -> String {
41 let dialect = GenericDialect {};
42 match Parser::parse_sql(&dialect, sql) {
43 Ok(statements) => {
44 if statements.is_empty() {
45 return sql.trim().to_lowercase();
46 }
47 let normalized: Vec<String> =
48 statements.iter().map(|stmt| stmt.to_string()).collect();
49 normalized.join("; ")
50 }
51 Err(_) => {
52 let trimmed: String = sql.split_whitespace().collect::<Vec<_>>().join(" ");
53 trimmed.to_lowercase()
54 }
55 }
56 }
57
58 pub fn extract_tables(sql: &str) -> Vec<String> {
62 let dialect = GenericDialect {};
63 let mut tables = Vec::new();
64
65 if let Ok(statements) = Parser::parse_sql(&dialect, sql) {
66 for stmt in &statements {
67 Self::extract_tables_from_stmt(stmt, &mut tables);
68 }
69 }
70
71 if tables.is_empty() {
72 Self::extract_tables_from_str(sql, &mut tables);
73 }
74
75 tables.sort();
76 tables.dedup();
77 tables
78 }
79
80 fn extract_tables_from_stmt(stmt: &Statement, tables: &mut Vec<String>) {
81 use sqlparser::ast::SetExpr;
82
83 match stmt {
84 Statement::Query(query) => {
85 if let SetExpr::Select(select) = &*query.body {
86 for from in &select.from {
87 Self::extract_table_factor(&from.relation, tables);
88 for join in &from.joins {
89 Self::extract_table_factor(&join.relation, tables);
90 }
91 }
92 }
93 }
94 Statement::Insert(_) => {}
95 Statement::Update { table, .. } => {
96 Self::extract_table_factor(&table.relation, tables);
97 }
98 Statement::Delete(delete) => {
99 for table_name in &delete.tables {
100 let full_name = table_name
101 .0
102 .iter()
103 .map(|i| i.value.clone())
104 .collect::<Vec<_>>()
105 .join(".");
106 tables.push(full_name);
107 }
108 }
109 _ => {}
110 }
111 }
112
113 fn extract_table_factor(factor: &sqlparser::ast::TableFactor, tables: &mut Vec<String>) {
114 use sqlparser::ast::TableFactor;
115 if let TableFactor::Table { name, .. } = factor {
116 let full_name = name
117 .0
118 .iter()
119 .map(|i| i.value.clone())
120 .collect::<Vec<_>>()
121 .join(".");
122 tables.push(full_name);
123 }
124 }
125
126 fn extract_tables_from_str(sql: &str, tables: &mut Vec<String>) {
127 let lower = sql.to_lowercase();
128 for keyword in ["into ", "from ", "update ", "join "] {
129 let mut search_pos = 0;
130 while let Some(pos) = lower[search_pos..].find(keyword) {
131 let abs_pos = search_pos + pos;
132 let rest = &sql[abs_pos + keyword.len()..];
133 let table: String = rest
134 .chars()
135 .take_while(|c| c.is_alphanumeric() || *c == '_' || *c == '.')
136 .collect();
137 if !table.is_empty() && !table.chars().all(|c| c.is_numeric()) {
138 tables.push(table);
139 }
140 search_pos = abs_pos + keyword.len();
141 }
142 }
143 }
144}
145
146#[derive(Debug, Clone)]
153pub struct PlanCacheKey {
154 pub hash: u64,
156 pub sql_normalized: String,
158}
159
160impl PlanCacheKey {
161 pub fn from_sql(sql: &str) -> Self {
166 let sql_normalized = SqlNormalizer::normalize(sql);
167 let mut h = XxHash64::with_seed(0);
168 h.write(sql_normalized.as_bytes());
169 let hash = h.finish();
170 Self {
171 hash,
172 sql_normalized,
173 }
174 }
175}
176
177impl PartialEq for PlanCacheKey {
178 fn eq(&self, other: &Self) -> bool {
179 self.hash == other.hash && self.sql_normalized == other.sql_normalized
180 }
181}
182
183impl Eq for PlanCacheKey {}
184
185impl std::hash::Hash for PlanCacheKey {
186 fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
187 self.hash.hash(state);
188 }
189}
190
191pub struct PlanCacheEntry {
197 pub ast: Option<Arc<Statement>>,
199 pub analysis: Option<Arc<String>>,
201 pub created_at: Instant,
203 pub tables: Vec<String>,
205 pub ttl: Option<Duration>,
207}
208
209impl PlanCacheEntry {
210 pub fn is_expired(&self) -> bool {
212 if let Some(ttl) = self.ttl {
213 self.created_at.elapsed() >= ttl
214 } else {
215 false
216 }
217 }
218}
219
220pub struct PlanCacheStats {
226 parse_hits: AtomicU64,
227 parse_misses: AtomicU64,
228 optimize_hits: AtomicU64,
229 optimize_misses: AtomicU64,
230 evictions: AtomicU64,
231}
232
233impl PlanCacheStats {
234 fn new() -> Self {
235 Self {
236 parse_hits: AtomicU64::new(0),
237 parse_misses: AtomicU64::new(0),
238 optimize_hits: AtomicU64::new(0),
239 optimize_misses: AtomicU64::new(0),
240 evictions: AtomicU64::new(0),
241 }
242 }
243
244 pub fn parse_hits(&self) -> u64 {
246 self.parse_hits.load(Ordering::Relaxed)
247 }
248
249 pub fn parse_misses(&self) -> u64 {
251 self.parse_misses.load(Ordering::Relaxed)
252 }
253
254 pub fn optimize_hits(&self) -> u64 {
256 self.optimize_hits.load(Ordering::Relaxed)
257 }
258
259 pub fn optimize_misses(&self) -> u64 {
261 self.optimize_misses.load(Ordering::Relaxed)
262 }
263
264 pub fn evictions(&self) -> u64 {
266 self.evictions.load(Ordering::Relaxed)
267 }
268
269 pub fn parse_hit_rate(&self) -> f64 {
271 let hits = self.parse_hits();
272 let misses = self.parse_misses();
273 let total = hits + misses;
274 if total == 0 {
275 0.0
276 } else {
277 hits as f64 / total as f64
278 }
279 }
280
281 pub fn optimize_hit_rate(&self) -> f64 {
283 let hits = self.optimize_hits();
284 let misses = self.optimize_misses();
285 let total = hits + misses;
286 if total == 0 {
287 0.0
288 } else {
289 hits as f64 / total as f64
290 }
291 }
292}
293
294impl Default for PlanCacheStats {
295 fn default() -> Self {
296 Self::new()
297 }
298}
299
300#[derive(Debug, Clone)]
304pub struct PlanCacheStatsSnapshot {
305 pub parse_hits: u64,
307 pub parse_misses: u64,
309 pub optimize_hits: u64,
311 pub optimize_misses: u64,
313 pub evictions: u64,
315 pub size: usize,
317 pub parse_hit_rate: f64,
319 pub optimize_hit_rate: f64,
321}
322
323struct LruOrder64 {
327 nodes: Vec<LruNode64>,
328 free_list: Vec<usize>,
329 index: HashMap<u64, usize>,
330 head: Option<usize>,
331 tail: Option<usize>,
332}
333
334struct LruNode64 {
335 key: u64,
336 prev: Option<usize>,
337 next: Option<usize>,
338}
339
340impl LruOrder64 {
341 fn new() -> Self {
342 Self {
343 nodes: Vec::new(),
344 free_list: Vec::new(),
345 index: HashMap::new(),
346 head: None,
347 tail: None,
348 }
349 }
350
351 fn touch(&mut self, key: u64) {
352 if let Some(&idx) = self.index.get(&key) {
353 self.unlink(idx);
354 self.link_tail(idx);
355 } else {
356 let idx = self.alloc_node(key);
357 self.link_tail(idx);
358 self.index.insert(key, idx);
359 }
360 }
361
362 fn remove(&mut self, key: u64) {
363 if let Some(idx) = self.index.remove(&key) {
364 self.unlink(idx);
365 self.free_node(idx);
366 }
367 }
368
369 fn lru_key(&self) -> Option<u64> {
370 self.head.map(|idx| self.nodes[idx].key)
371 }
372
373 fn clear(&mut self) {
374 self.nodes.clear();
375 self.free_list.clear();
376 self.index.clear();
377 self.head = None;
378 self.tail = None;
379 }
380
381 fn len(&self) -> usize {
382 self.index.len()
383 }
384
385 fn alloc_node(&mut self, key: u64) -> usize {
386 if let Some(idx) = self.free_list.pop() {
387 self.nodes[idx] = LruNode64 {
388 key,
389 prev: None,
390 next: None,
391 };
392 idx
393 } else {
394 self.nodes.push(LruNode64 {
395 key,
396 prev: None,
397 next: None,
398 });
399 self.nodes.len() - 1
400 }
401 }
402
403 fn free_node(&mut self, idx: usize) {
404 self.free_list.push(idx);
405 }
406
407 fn unlink(&mut self, idx: usize) {
408 let prev = self.nodes[idx].prev;
409 let next = self.nodes[idx].next;
410 match prev {
411 Some(p) => self.nodes[p].next = next,
412 None => self.head = next,
413 }
414 match next {
415 Some(n) => self.nodes[n].prev = prev,
416 None => self.tail = prev,
417 }
418 self.nodes[idx].prev = None;
419 self.nodes[idx].next = None;
420 }
421
422 fn link_tail(&mut self, idx: usize) {
423 match self.tail {
424 Some(t) => {
425 self.nodes[idx].prev = Some(t);
426 self.nodes[t].next = Some(idx);
427 }
428 None => {
429 self.head = Some(idx);
430 }
431 }
432 self.tail = Some(idx);
433 }
434}
435
436pub struct PlanCache {
447 parse_cache: RwLock<HashMap<u64, PlanCacheEntry>>,
449 optimize_cache: RwLock<HashMap<u64, PlanCacheEntry>>,
451 access_order: RwLock<LruOrder64>,
453 table_index: RwLock<HashMap<String, Vec<u64>>>,
455 stats: PlanCacheStats,
457 max_size: usize,
459 default_ttl: Option<Duration>,
461 access_counts: RwLock<HashMap<u64, u32>>,
463 lru_k: u32,
465 min_capacity: usize,
467 max_capacity: usize,
469}
470
471impl PlanCache {
472 pub fn new(max_size: usize, default_ttl: Option<Duration>) -> Self {
477 Self {
478 parse_cache: RwLock::new(HashMap::new()),
479 optimize_cache: RwLock::new(HashMap::new()),
480 access_order: RwLock::new(LruOrder64::new()),
481 table_index: RwLock::new(HashMap::new()),
482 stats: PlanCacheStats::new(),
483 max_size,
484 default_ttl,
485 access_counts: RwLock::new(HashMap::new()),
486 lru_k: 2,
487 min_capacity: max_size / 4,
488 max_capacity: max_size * 4,
489 }
490 }
491
492 pub fn get_or_parse(&self, sql: &str) -> Result<Arc<Statement>, String> {
497 let key = PlanCacheKey::from_sql(sql);
498
499 let hit = {
502 let cache = self.parse_cache.read();
503 match cache.get(&key.hash) {
504 Some(entry) if !entry.is_expired() => entry.ast.clone(),
505 _ => None,
506 }
507 };
508 if let Some(ast) = hit {
509 self.stats.parse_hits.fetch_add(1, Ordering::Relaxed);
510 self.access_order.write().touch(key.hash);
511 self.access_counts.write().entry(key.hash).and_modify(|c| *c += 1).or_insert(1);
512 return Ok(ast);
513 }
514
515 self.stats.parse_misses.fetch_add(1, Ordering::Relaxed);
516
517 let dialect = GenericDialect {};
518 let statements = Parser::parse_sql(&dialect, sql).map_err(|e| e.to_string())?;
519 if statements.is_empty() {
520 return Err("empty SQL".to_string());
521 }
522 let ast = Arc::new(statements.into_iter().next().unwrap());
523 let tables = SqlNormalizer::extract_tables(sql);
524
525 {
526 let mut access_order = self.access_order.write();
527 let mut optimize_cache = self.optimize_cache.write();
529 let mut cache = self.parse_cache.write();
530 let mut table_index = self.table_index.write();
531
532 if cache.len() >= self.max_size {
533 if let Some(lru_hash) = access_order.lru_key() {
534 access_order.remove(lru_hash);
535 cache.remove(&lru_hash);
536 optimize_cache.remove(&lru_hash);
537 Self::remove_from_table_index(&mut table_index, lru_hash);
538 self.stats.evictions.fetch_add(1, Ordering::Relaxed);
539 }
540 }
541
542 let entry = PlanCacheEntry {
543 ast: Some(ast.clone()),
544 analysis: None,
545 created_at: Instant::now(),
546 tables: tables.clone(),
547 ttl: self.default_ttl,
548 };
549 cache.insert(key.hash, entry);
550 access_order.touch(key.hash);
551
552 for table in &tables {
553 table_index.entry(table.clone()).or_default().push(key.hash);
554 }
555 }
556
557 Ok(ast)
558 }
559
560 pub fn get_or_optimize(&self, sql: &str) -> Option<Arc<String>> {
565 let key = PlanCacheKey::from_sql(sql);
566
567 let hit = {
568 let cache = self.optimize_cache.read();
569 match cache.get(&key.hash) {
570 Some(entry) if !entry.is_expired() => entry.analysis.clone(),
571 _ => None,
572 }
573 };
574 if let Some(analysis) = hit {
575 self.stats.optimize_hits.fetch_add(1, Ordering::Relaxed);
576 self.access_order.write().touch(key.hash);
577 return Some(analysis);
578 }
579
580 self.stats.optimize_misses.fetch_add(1, Ordering::Relaxed);
581 None
582 }
583
584 pub fn store_optimize(&self, sql: &str, analysis: Arc<String>) {
588 let key = PlanCacheKey::from_sql(sql);
589 let tables = SqlNormalizer::extract_tables(sql);
590
591 let mut access_order = self.access_order.write();
592 let mut cache = self.optimize_cache.write();
593 let mut parse_cache = self.parse_cache.write();
594 let mut table_index = self.table_index.write();
595
596 if cache.len() >= self.max_size {
597 if let Some(lru_hash) = access_order.lru_key() {
598 access_order.remove(lru_hash);
599 cache.remove(&lru_hash);
600 parse_cache.remove(&lru_hash);
601 Self::remove_from_table_index(&mut table_index, lru_hash);
602 self.stats.evictions.fetch_add(1, Ordering::Relaxed);
603 }
604 }
605
606 let entry = PlanCacheEntry {
607 ast: None,
608 analysis: Some(analysis),
609 created_at: Instant::now(),
610 tables: tables.clone(),
611 ttl: self.default_ttl,
612 };
613 cache.insert(key.hash, entry);
614 access_order.touch(key.hash);
615
616 for table in &tables {
617 table_index.entry(table.clone()).or_default().push(key.hash);
618 }
619 }
620
621 pub fn invalidate_table(&self, table: &str) -> usize {
625 let mut access_order = self.access_order.write();
627 let mut optimize_cache = self.optimize_cache.write();
628 let mut parse_cache = self.parse_cache.write();
629 let mut table_index = self.table_index.write();
630 let keys = table_index.remove(table).unwrap_or_default();
631
632 if keys.is_empty() {
633 return 0;
634 }
635
636 let count = keys.len();
637
638 for &hash in &keys {
639 access_order.remove(hash);
640 parse_cache.remove(&hash);
641 optimize_cache.remove(&hash);
642 }
643
644 for remaining_keys in table_index.values_mut() {
645 remaining_keys.retain(|k| !keys.contains(k));
646 }
647
648 self.stats
649 .evictions
650 .fetch_add(count as u64, Ordering::Relaxed);
651 count
652 }
653
654 pub fn invalidate_all(&self) {
656 self.parse_cache.write().clear();
657 self.optimize_cache.write().clear();
658 self.access_order.write().clear();
659 self.table_index.write().clear();
660 }
661
662 pub fn stats(&self) -> PlanCacheStatsSnapshot {
664 let size = self.access_order.read().len();
665 PlanCacheStatsSnapshot {
666 parse_hits: self.stats.parse_hits(),
667 parse_misses: self.stats.parse_misses(),
668 optimize_hits: self.stats.optimize_hits(),
669 optimize_misses: self.stats.optimize_misses(),
670 evictions: self.stats.evictions(),
671 size,
672 parse_hit_rate: self.stats.parse_hit_rate(),
673 optimize_hit_rate: self.stats.optimize_hit_rate(),
674 }
675 }
676
677 pub fn size(&self) -> usize {
679 self.access_order.read().len()
680 }
681
682 fn remove_from_table_index(table_index: &mut HashMap<String, Vec<u64>>, hash: u64) {
683 for keys in table_index.values_mut() {
684 keys.retain(|k| *k != hash);
685 }
686 }
687}
688
689impl Default for PlanCache {
690 fn default() -> Self {
691 Self::new(1024, None)
692 }
693}
694
695#[derive(Debug, Clone)]
701pub struct PlanCacheConfig {
702 pub capacity: usize,
704 pub ttl: Duration,
706 pub include_param_types: bool,
708}
709
710impl Default for PlanCacheConfig {
711 fn default() -> Self {
712 Self {
713 capacity: 256,
714 ttl: Duration::from_secs(300),
715 include_param_types: true,
716 }
717 }
718}
719
720impl PlanCacheConfig {
721 pub fn new() -> Self {
723 Self::default()
724 }
725
726 pub fn with_capacity(mut self, capacity: usize) -> Self {
728 self.capacity = capacity;
729 self
730 }
731
732 pub fn with_ttl(mut self, ttl: Duration) -> Self {
734 self.ttl = ttl;
735 self
736 }
737
738 pub fn with_param_types(mut self, include: bool) -> Self {
740 self.include_param_types = include;
741 self
742 }
743}
744
745pub fn fingerprint_with_types(sql: &str, param_types: &[crate::DbType]) -> u64 {
754 let mut hasher = XxHash64::with_seed(0);
755 hasher.write(sql.as_bytes());
756 for pt in param_types {
757 hasher.write(&[*pt as u8]);
758 }
759 hasher.finish()
760}
761
762impl PlanCache {
763 pub fn eviction_count(&self) -> u64 {
765 self.stats.evictions.load(Ordering::Relaxed)
766 }
767
768 pub fn with_config(config: &PlanCacheConfig) -> Self {
770 Self::new(config.capacity, Some(config.ttl))
771 }
772
773 pub fn get_or_parse_with_types(
777 &self,
778 sql: &str,
779 param_types: &[crate::DbType],
780 ) -> Result<Arc<Statement>, String> {
781 let type_fingerprint = fingerprint_with_types(sql, param_types);
782
783 let hit = {
784 let cache = self.parse_cache.read();
785 match cache.get(&type_fingerprint) {
786 Some(entry) if !entry.is_expired() => entry.ast.clone(),
787 _ => None,
788 }
789 };
790 if let Some(ast) = hit {
791 self.stats.parse_hits.fetch_add(1, Ordering::Relaxed);
792 self.access_order.write().touch(type_fingerprint);
793 return Ok(ast);
794 }
795
796 self.stats.parse_misses.fetch_add(1, Ordering::Relaxed);
797
798 let dialect = GenericDialect {};
799 let statements = Parser::parse_sql(&dialect, sql).map_err(|e| e.to_string())?;
800 if statements.is_empty() {
801 return Err("empty SQL".to_string());
802 }
803 let ast = Arc::new(statements.into_iter().next().unwrap());
804 let tables = SqlNormalizer::extract_tables(sql);
805
806 {
807 let mut access_order = self.access_order.write();
808 let mut optimize_cache = self.optimize_cache.write();
809 let mut cache = self.parse_cache.write();
810 let mut table_index = self.table_index.write();
811
812 if cache.len() >= self.max_size {
813 if let Some(lru_hash) = access_order.lru_key() {
814 access_order.remove(lru_hash);
815 cache.remove(&lru_hash);
816 optimize_cache.remove(&lru_hash);
817 Self::remove_from_table_index(&mut table_index, lru_hash);
818 self.stats.evictions.fetch_add(1, Ordering::Relaxed);
819 }
820 }
821
822 let entry = PlanCacheEntry {
823 ast: Some(ast.clone()),
824 analysis: None,
825 created_at: Instant::now(),
826 tables: tables.clone(),
827 ttl: self.default_ttl,
828 };
829 cache.insert(type_fingerprint, entry);
830 access_order.touch(type_fingerprint);
831
832 for table in &tables {
833 table_index
834 .entry(table.clone())
835 .or_default()
836 .push(type_fingerprint);
837 }
838 }
839
840 Ok(ast)
841 }
842
843 pub fn access_count(&self, hash: u64) -> u32 {
845 self.access_counts.read().get(&hash).copied().unwrap_or(0)
846 }
847
848 pub fn lru_k(&self) -> u32 {
850 self.lru_k
851 }
852
853 pub fn parse_hit_rate(&self) -> f64 {
855 let hits = self.stats.parse_hits.load(Ordering::Relaxed);
856 let misses = self.stats.parse_misses.load(Ordering::Relaxed);
857 let total = hits + misses;
858 if total == 0 {
859 0.0
860 } else {
861 hits as f64 / total as f64
862 }
863 }
864
865 pub fn adjust_capacity(&mut self) -> bool {
872 let hit_rate = self.parse_hit_rate();
873 let old_size = self.max_size;
874 if hit_rate > 0.8 && self.max_size < self.max_capacity {
875 self.max_size = (self.max_size * 3 / 2).min(self.max_capacity);
876 } else if hit_rate < 0.3 && self.max_size > self.min_capacity {
877 self.max_size = (self.max_size * 2 / 3).max(self.min_capacity);
878 }
879 self.max_size != old_size
880 }
881
882 pub fn capacity(&self) -> usize {
884 self.max_size
885 }
886}
887
888#[cfg(test)]
891mod tests {
892 use super::*;
893
894 #[test]
895 fn test_sql_normalizer_basic() {
896 let sql1 = "SELECT * FROM users WHERE id = ?";
897 let sql2 = "select * from users where id = ?";
898 let n1 = SqlNormalizer::normalize(sql1);
899 let n2 = SqlNormalizer::normalize(sql2);
900 assert_eq!(n1, n2, "大小写差异应归一化");
901 }
902
903 #[test]
904 fn test_sql_normalizer_whitespace() {
905 let sql1 = "SELECT * FROM users";
906 let sql2 = "SELECT * FROM users";
907 let n1 = SqlNormalizer::normalize(sql1);
908 let n2 = SqlNormalizer::normalize(sql2);
909 assert_eq!(n1, n2, "空白差异应归一化");
910 }
911
912 #[test]
913 fn test_sql_normalizer_different_semantics() {
914 let sql1 = "SELECT * FROM users WHERE id = ?";
915 let sql2 = "SELECT * FROM orders WHERE id = ?";
916 let n1 = SqlNormalizer::normalize(sql1);
917 let n2 = SqlNormalizer::normalize(sql2);
918 assert_ne!(n1, n2, "不同表名应产生不同归一化");
919 }
920
921 #[test]
922 fn test_sql_normalizer_parse_error_fallback() {
923 let sql = "this is not valid sql !!!";
924 let normalized = SqlNormalizer::normalize(sql);
925 assert!(!normalized.is_empty());
926 }
927
928 #[test]
929 fn test_extract_tables_select() {
930 let sql = "SELECT * FROM users JOIN orders ON users.id = orders.user_id";
931 let tables = SqlNormalizer::extract_tables(sql);
932 assert!(tables.contains(&"users".to_string()));
933 assert!(tables.contains(&"orders".to_string()));
934 }
935
936 #[test]
937 fn test_extract_tables_insert() {
938 let sql = "INSERT INTO products (name) VALUES (?)";
939 let tables = SqlNormalizer::extract_tables(sql);
940 assert!(tables.contains(&"products".to_string()));
941 }
942
943 #[test]
944 fn test_extract_tables_update() {
945 let sql = "UPDATE products SET name = ? WHERE id = ?";
946 let tables = SqlNormalizer::extract_tables(sql);
947 assert!(tables.contains(&"products".to_string()));
948 }
949
950 #[test]
951 fn test_extract_tables_delete() {
952 let sql = "DELETE FROM products WHERE id = ?";
953 let tables = SqlNormalizer::extract_tables(sql);
954 assert!(
955 tables.iter().any(|t| t.contains("products")),
956 "应包含 products 表,实际: {:?}",
957 tables
958 );
959 }
960
961 #[test]
962 fn test_plan_cache_key_same_sql() {
963 let k1 = PlanCacheKey::from_sql("SELECT * FROM users WHERE id = ?");
964 let k2 = PlanCacheKey::from_sql("select * from users where id = ?");
965 assert_eq!(k1.hash, k2.hash, "相同 SQL 模板应产生相同 hash");
966 }
967
968 #[test]
969 fn test_plan_cache_key_different_sql() {
970 let k1 = PlanCacheKey::from_sql("SELECT * FROM users");
971 let k2 = PlanCacheKey::from_sql("SELECT * FROM orders");
972 assert_ne!(k1.hash, k2.hash, "不同 SQL 应产生不同 hash");
973 }
974
975 #[test]
976 fn test_plan_cache_key_no_sensitive_data() {
977 let key = PlanCacheKey::from_sql("SELECT * FROM users WHERE password = ?");
978 assert!(
979 !key.sql_normalized.contains("secret123"),
980 "参数化查询缓存键不应包含参数值"
981 );
982 }
983
984 #[test]
985 fn test_plan_cache_new() {
986 let cache = PlanCache::new(1024, None);
987 assert_eq!(cache.size(), 0);
988 let stats = cache.stats();
989 assert_eq!(stats.size, 0);
990 assert_eq!(stats.parse_hits, 0);
991 assert_eq!(stats.parse_misses, 0);
992 }
993
994 #[test]
995 fn test_plan_cache_default() {
996 let cache = PlanCache::default();
997 assert_eq!(cache.size(), 0);
998 }
999
1000 #[test]
1001 fn test_get_or_parse_hit() {
1002 let cache = PlanCache::new(100, None);
1003 let sql = "SELECT * FROM users WHERE id = ?";
1004
1005 let ast1 = cache.get_or_parse(sql).expect("parse");
1006 assert_eq!(cache.stats().parse_misses, 1, "首次应 miss");
1007
1008 let ast2 = cache.get_or_parse(sql).expect("parse");
1009 assert_eq!(cache.stats().parse_hits, 1, "第二次应 hit");
1010 assert_eq!(cache.stats().parse_misses, 1, "misses 不变");
1011
1012 assert!(Arc::ptr_eq(&ast1, &ast2), "命中应返回相同 Arc");
1013 }
1014
1015 #[test]
1016 fn test_get_or_parse_different_params_same_template() {
1017 let cache = PlanCache::new(100, None);
1018 let sql1 = "SELECT * FROM users WHERE id = ?";
1019 let sql2 = "select * from users where id = ?";
1020
1021 cache.get_or_parse(sql1).expect("parse");
1022 cache.get_or_parse(sql2).expect("parse");
1023
1024 assert_eq!(cache.stats().parse_hits, 1, "相同模板不同写法应命中");
1025 assert_eq!(cache.stats().parse_misses, 1);
1026 }
1027
1028 #[test]
1029 fn test_get_or_parse_different_sql() {
1030 let cache = PlanCache::new(100, None);
1031 cache.get_or_parse("SELECT * FROM users").expect("parse");
1032 cache.get_or_parse("SELECT * FROM orders").expect("parse");
1033
1034 assert_eq!(cache.stats().parse_misses, 2, "不同 SQL 应各 miss 一次");
1035 assert_eq!(cache.stats().parse_hits, 0);
1036 }
1037
1038 #[test]
1039 fn test_get_or_optimize_miss_then_store_then_hit() {
1040 let cache = PlanCache::new(100, None);
1041 let sql = "SELECT * FROM users WHERE id = ?";
1042
1043 assert!(cache.get_or_optimize(sql).is_none(), "首次应 miss");
1044 assert_eq!(cache.stats().optimize_misses, 1);
1045
1046 cache.store_optimize(sql, Arc::new("optimized plan".to_string()));
1047
1048 let result = cache.get_or_optimize(sql);
1049 assert!(result.is_some(), "存储后应 hit");
1050 assert_eq!(cache.stats().optimize_hits, 1);
1051 assert_eq!(*result.unwrap().as_ref(), "optimized plan");
1052 }
1053
1054 #[test]
1055 fn test_invalidate_table_precise() {
1056 let cache = PlanCache::new(100, None);
1057 cache.get_or_parse("SELECT * FROM users").expect("parse");
1058 cache.get_or_parse("SELECT * FROM orders").expect("parse");
1059 assert_eq!(cache.size(), 2);
1060
1061 let evicted = cache.invalidate_table("users");
1062 assert_eq!(evicted, 1, "应失效 1 条");
1063 assert_eq!(cache.size(), 1, "应剩余 1 条");
1064
1065 let stats = cache.stats();
1066 assert!(stats.parse_hits == 0, "orders 缓存应不受影响");
1067 cache.get_or_parse("SELECT * FROM orders").expect("parse");
1068 assert_eq!(cache.stats().parse_hits, 1, "orders 应命中缓存");
1069 }
1070
1071 #[test]
1072 fn test_invalidate_table_nonexistent() {
1073 let cache = PlanCache::new(100, None);
1074 cache.get_or_parse("SELECT * FROM users").expect("parse");
1075 let evicted = cache.invalidate_table("nonexistent");
1076 assert_eq!(evicted, 0, "不存在的表应返回 0");
1077 assert_eq!(cache.size(), 1, "缓存不应受影响");
1078 }
1079
1080 #[test]
1081 fn test_invalidate_all() {
1082 let cache = PlanCache::new(100, None);
1083 cache.get_or_parse("SELECT * FROM users").expect("parse");
1084 cache.get_or_parse("SELECT * FROM orders").expect("parse");
1085 assert_eq!(cache.size(), 2);
1086
1087 cache.invalidate_all();
1088 assert_eq!(cache.size(), 0, "全量清空后 size 应为 0");
1089 }
1090
1091 #[test]
1092 fn test_lru_eviction() {
1093 let cache = PlanCache::new(3, None);
1094 cache.get_or_parse("SELECT * FROM t1").expect("parse");
1095 cache.get_or_parse("SELECT * FROM t2").expect("parse");
1096 cache.get_or_parse("SELECT * FROM t3").expect("parse");
1097 assert_eq!(cache.size(), 3);
1098
1099 cache.get_or_parse("SELECT * FROM t4").expect("parse");
1100 assert_eq!(cache.size(), 3, "max_size=3 应保持 3 条");
1101 assert!(cache.stats().evictions >= 1, "应有淘汰");
1102
1103 cache.get_or_parse("SELECT * FROM t1").expect("parse");
1104 assert!(cache.stats().parse_misses >= 4, "t1 被淘汰后应重新 miss");
1105 }
1106
1107 #[test]
1108 fn test_lru_eviction_max_size_1() {
1109 let cache = PlanCache::new(1, None);
1110 cache.get_or_parse("SELECT * FROM t1").expect("parse");
1111 assert_eq!(cache.size(), 1);
1112
1113 cache.get_or_parse("SELECT * FROM t2").expect("parse");
1114 assert_eq!(cache.size(), 1, "max_size=1 应保持 1 条");
1115
1116 cache.get_or_parse("SELECT * FROM t1").expect("parse");
1117 assert!(cache.stats().parse_misses >= 3, "t1 应被淘汰后重新 miss");
1118 }
1119
1120 #[test]
1121 fn test_stats_hit_rate() {
1122 let cache = PlanCache::new(100, None);
1123 let sql = "SELECT * FROM users";
1124
1125 cache.get_or_parse(sql).expect("parse");
1126 cache.get_or_parse(sql).expect("parse");
1127 cache.get_or_parse(sql).expect("parse");
1128
1129 let stats = cache.stats();
1130 assert_eq!(stats.parse_hits, 2);
1131 assert_eq!(stats.parse_misses, 1);
1132 assert!((stats.parse_hit_rate - (2.0 / 3.0)).abs() < 0.001);
1133 }
1134
1135 #[test]
1136 fn test_stats_hit_rate_empty() {
1137 let cache = PlanCache::new(100, None);
1138 let stats = cache.stats();
1139 assert_eq!(stats.parse_hit_rate, 0.0, "空缓存命中率应为 0.0");
1140 assert_eq!(stats.optimize_hit_rate, 0.0);
1141 }
1142
1143 #[test]
1144 fn test_ttl_expiration() {
1145 let cache = PlanCache::new(100, Some(Duration::from_nanos(1)));
1146 let sql = "SELECT * FROM users";
1147
1148 cache.get_or_parse(sql).expect("parse");
1149 std::thread::sleep(Duration::from_millis(10));
1150
1151 cache.get_or_parse(sql).expect("parse");
1152 assert!(cache.stats().parse_misses >= 2, "TTL 过期后应重新 miss");
1153 }
1154
1155 #[test]
1156 fn test_plan_cache_concurrent_same_sql() {
1157 use std::sync::Arc;
1158 use std::thread;
1159
1160 let cache = Arc::new(PlanCache::new(100, None));
1161 let sql = "SELECT * FROM users WHERE id = ?";
1162 let mut handles = Vec::new();
1163
1164 for _ in 0..10 {
1165 let cache = cache.clone();
1166 handles.push(thread::spawn(move || {
1167 cache.get_or_parse(sql).expect("parse");
1168 }));
1169 }
1170
1171 for h in handles {
1172 h.join().expect("thread");
1173 }
1174
1175 assert!(cache.size() >= 1, "并发后应至少有 1 条缓存");
1176 assert!(
1177 cache.stats().parse_misses + cache.stats().parse_hits >= 10,
1178 "应有 10 次访问记录"
1179 );
1180 }
1181}
1182#[derive(Debug, Clone)]
1183pub struct LruKReplacer {
1184 k: usize,
1185 capacity: usize,
1186 access_history: std::collections::HashMap<u64, Vec<std::time::Instant>>,
1187}
1188
1189impl LruKReplacer {
1190 pub fn new(k: usize, capacity: usize) -> Self {
1191 Self {
1192 k,
1193 capacity,
1194 access_history: std::collections::HashMap::new(),
1195 }
1196 }
1197
1198 pub fn access(&mut self, key: u64) {
1199 let history = self.access_history.entry(key).or_default();
1200 history.push(std::time::Instant::now());
1201 if history.len() > self.k {
1202 history.remove(0);
1203 }
1204 }
1205
1206 pub fn evict(&mut self) -> Option<u64> {
1207 if self.access_history.len() < self.capacity {
1208 return None;
1209 }
1210 let mut oldest_key: Option<u64> = None;
1211 let mut oldest_time: Option<std::time::Instant> = None;
1212
1213 for (&key, history) in &self.access_history {
1214 let ref_time = if history.len() >= self.k {
1215 history[0]
1216 } else {
1217 std::time::Instant::now()
1218 };
1219 if oldest_time.is_none() || ref_time < oldest_time.unwrap() {
1220 oldest_time = Some(ref_time);
1221 oldest_key = Some(key);
1222 }
1223 }
1224
1225 if let Some(key) = oldest_key {
1226 self.access_history.remove(&key);
1227 }
1228 oldest_key
1229 }
1230
1231 pub fn hit_rate(&self) -> f64 {
1232 let total: usize = self.access_history.values().map(|h| h.len()).sum();
1233 if total == 0 {
1234 return 0.0;
1235 }
1236 let hits: usize = self
1237 .access_history
1238 .values()
1239 .map(|h| if h.len() >= self.k { 1 } else { 0 })
1240 .sum();
1241 hits as f64 / self.access_history.len().max(1) as f64
1242 }
1243
1244 pub fn len(&self) -> usize {
1245 self.access_history.len()
1246 }
1247
1248 pub fn is_empty(&self) -> bool {
1249 self.access_history.is_empty()
1250 }
1251}
1252
1253#[derive(Debug, Clone)]
1254pub struct AdaptiveCacheCapacity {
1255 min: usize,
1256 max: usize,
1257 current: usize,
1258 hit_rate_window: std::collections::VecDeque<f64>,
1259 window_size: usize,
1260}
1261
1262impl AdaptiveCacheCapacity {
1263 pub fn new(min: usize, max: usize) -> Self {
1264 Self {
1265 min,
1266 max,
1267 current: min,
1268 hit_rate_window: std::collections::VecDeque::new(),
1269 window_size: 100,
1270 }
1271 }
1272
1273 pub fn record_hit_rate(&mut self, rate: f64) {
1274 if self.hit_rate_window.len() >= self.window_size {
1275 self.hit_rate_window.pop_front();
1276 }
1277 self.hit_rate_window.push_back(rate);
1278 }
1279
1280 pub fn adapt(&mut self) -> bool {
1281 if self.hit_rate_window.is_empty() {
1282 return false;
1283 }
1284 let avg: f64 =
1285 self.hit_rate_window.iter().sum::<f64>() / self.hit_rate_window.len() as f64;
1286 let old = self.current;
1287 if avg < 0.85 {
1288 self.current = (self.current as f64 * 1.5) as usize;
1289 } else if avg > 0.95 {
1290 self.current = (self.current as f64 * 0.8) as usize;
1291 }
1292 self.current = self.current.clamp(self.min, self.max);
1293 self.current != old
1294 }
1295
1296 pub fn current(&self) -> usize {
1297 self.current
1298 }
1299
1300 pub fn min(&self) -> usize {
1301 self.min
1302 }
1303
1304 pub fn max(&self) -> usize {
1305 self.max
1306 }
1307}