Skip to main content

sz_orm_core/
plan_cache.rs

1//! v3.2.0 查询计划缓存
2//!
3//! 缓存 SQL 解析结果(AST)与查询优化结果,相同 SQL 模板第二次起跳过解析/优化(≤1μs)。
4//! 与 L2Cache(数据缓存)职责分离:L2Cache 缓存查询结果数据,PlanCache 缓存查询计划。
5//!
6//! 特性:
7//! - SQL 归一化(忽略空白/注释/参数顺序,参数值替换为占位符)
8//! - xxHash 64bit 键生成(无碰撞差分测试验证)
9//! - 双缓存(parse_cache + optimize_cache)
10//! - LRU 淘汰(arena 双向链表 O(1))
11//! - 表级精确失效(table_index 索引)
12//! - 命中率统计(原子计数器无锁)
13
14use 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 xxhash_rust::xxh64::xxh64;
24
25// ─── SqlNormalizer ───────────────────────────────────────────────
26
27/// SQL 归一化器
28///
29/// 将 SQL 文本归一化为标准形式:忽略大小写差异、空白差异、参数值差异。
30/// 相同语义不同写法的 SQL 归一化后产生相同文本,用于缓存键生成。
31pub struct SqlNormalizer;
32
33impl SqlNormalizer {
34    /// 归一化 SQL 文本
35    ///
36    /// 返回归一化后的 SQL 文本。参数值(如 WHERE id = 42 中的 42)
37    /// 不会被特殊处理——调用方应使用参数化查询(WHERE id = ?),
38    /// 归一化仅处理大小写和空白差异。
39    pub fn normalize(sql: &str) -> String {
40        let dialect = GenericDialect {};
41        match Parser::parse_sql(&dialect, sql) {
42            Ok(statements) => {
43                if statements.is_empty() {
44                    return sql.trim().to_lowercase();
45                }
46                let normalized: Vec<String> =
47                    statements.iter().map(|stmt| stmt.to_string()).collect();
48                normalized.join("; ")
49            }
50            Err(_) => {
51                let trimmed: String = sql.split_whitespace().collect::<Vec<_>>().join(" ");
52                trimmed.to_lowercase()
53            }
54        }
55    }
56
57    /// 从 SQL 中提取依赖的表名列表
58    ///
59    /// 遍历 AST 提取所有表引用(FROM / JOIN / INTO / UPDATE 等)。
60    pub fn extract_tables(sql: &str) -> Vec<String> {
61        let dialect = GenericDialect {};
62        let mut tables = Vec::new();
63
64        if let Ok(statements) = Parser::parse_sql(&dialect, sql) {
65            for stmt in &statements {
66                Self::extract_tables_from_stmt(stmt, &mut tables);
67            }
68        }
69
70        if tables.is_empty() {
71            Self::extract_tables_from_str(sql, &mut tables);
72        }
73
74        tables.sort();
75        tables.dedup();
76        tables
77    }
78
79    fn extract_tables_from_stmt(stmt: &Statement, tables: &mut Vec<String>) {
80        use sqlparser::ast::SetExpr;
81
82        match stmt {
83            Statement::Query(query) => {
84                if let SetExpr::Select(select) = &*query.body {
85                    for from in &select.from {
86                        Self::extract_table_factor(&from.relation, tables);
87                        for join in &from.joins {
88                            Self::extract_table_factor(&join.relation, tables);
89                        }
90                    }
91                }
92            }
93            Statement::Insert(_) => {}
94            Statement::Update { table, .. } => {
95                Self::extract_table_factor(&table.relation, tables);
96            }
97            Statement::Delete(delete) => {
98                for table_name in &delete.tables {
99                    let full_name = table_name
100                        .0
101                        .iter()
102                        .map(|i| i.value.clone())
103                        .collect::<Vec<_>>()
104                        .join(".");
105                    tables.push(full_name);
106                }
107            }
108            _ => {}
109        }
110    }
111
112    fn extract_table_factor(factor: &sqlparser::ast::TableFactor, tables: &mut Vec<String>) {
113        use sqlparser::ast::TableFactor;
114        if let TableFactor::Table { name, .. } = factor {
115            let full_name = name
116                .0
117                .iter()
118                .map(|i| i.value.clone())
119                .collect::<Vec<_>>()
120                .join(".");
121            tables.push(full_name);
122        }
123    }
124
125    fn extract_tables_from_str(sql: &str, tables: &mut Vec<String>) {
126        let lower = sql.to_lowercase();
127        for keyword in ["into ", "from ", "update ", "join "] {
128            let mut search_pos = 0;
129            while let Some(pos) = lower[search_pos..].find(keyword) {
130                let abs_pos = search_pos + pos;
131                let rest = &sql[abs_pos + keyword.len()..];
132                let table: String = rest
133                    .chars()
134                    .take_while(|c| c.is_alphanumeric() || *c == '_' || *c == '.')
135                    .collect();
136                if !table.is_empty() && !table.chars().all(|c| c.is_numeric()) {
137                    tables.push(table);
138                }
139                search_pos = abs_pos + keyword.len();
140            }
141        }
142    }
143}
144
145// ─── PlanCacheKey ────────────────────────────────────────────────
146
147/// 查询计划缓存键
148///
149/// 由归一化 SQL 的 xxHash 64bit 哈希值 + 归一化 SQL 文本组成。
150/// 哈希值用于快速查找,SQL 文本用于二次校验(防哈希碰撞)。
151#[derive(Debug, Clone)]
152pub struct PlanCacheKey {
153    /// xxHash 64bit 哈希值
154    pub hash: u64,
155    /// 归一化 SQL 文本(用于碰撞校验)
156    pub sql_normalized: String,
157}
158
159impl PlanCacheKey {
160    /// 从原始 SQL 生成缓存键
161    ///
162    /// 1. 归一化 SQL(忽略大小写/空白差异)
163    /// 2. 计算归一化 SQL 的 xxHash 64bit 哈希
164    pub fn from_sql(sql: &str) -> Self {
165        let sql_normalized = SqlNormalizer::normalize(sql);
166        let hash = xxh64(sql_normalized.as_bytes(), 0);
167        Self {
168            hash,
169            sql_normalized,
170        }
171    }
172}
173
174impl PartialEq for PlanCacheKey {
175    fn eq(&self, other: &Self) -> bool {
176        self.hash == other.hash && self.sql_normalized == other.sql_normalized
177    }
178}
179
180impl Eq for PlanCacheKey {}
181
182impl std::hash::Hash for PlanCacheKey {
183    fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
184        self.hash.hash(state);
185    }
186}
187
188// ─── PlanCacheEntry ──────────────────────────────────────────────
189
190/// 查询计划缓存条目
191///
192/// 存储解析后的 AST 和/或优化后的查询分析结果。
193pub struct PlanCacheEntry {
194    /// 解析后的 AST(parse_cache 条目)
195    pub ast: Option<Arc<Statement>>,
196    /// 优化后的查询分析(optimize_cache 条目)
197    pub analysis: Option<Arc<String>>,
198    /// 创建时间
199    pub created_at: Instant,
200    /// 依赖的表名列表(用于表级失效)
201    pub tables: Vec<String>,
202    /// TTL(可选,过期自动失效)
203    pub ttl: Option<Duration>,
204}
205
206impl PlanCacheEntry {
207    /// 检查条目是否已过期
208    pub fn is_expired(&self) -> bool {
209        if let Some(ttl) = self.ttl {
210            self.created_at.elapsed() >= ttl
211        } else {
212            false
213        }
214    }
215}
216
217// ─── PlanCacheStats ──────────────────────────────────────────────
218
219/// 查询计划缓存统计
220///
221/// 原子计数器无锁统计命中/未命中/淘汰次数。
222pub struct PlanCacheStats {
223    parse_hits: AtomicU64,
224    parse_misses: AtomicU64,
225    optimize_hits: AtomicU64,
226    optimize_misses: AtomicU64,
227    evictions: AtomicU64,
228}
229
230impl PlanCacheStats {
231    fn new() -> Self {
232        Self {
233            parse_hits: AtomicU64::new(0),
234            parse_misses: AtomicU64::new(0),
235            optimize_hits: AtomicU64::new(0),
236            optimize_misses: AtomicU64::new(0),
237            evictions: AtomicU64::new(0),
238        }
239    }
240
241    /// 解析缓存命中次数
242    pub fn parse_hits(&self) -> u64 {
243        self.parse_hits.load(Ordering::Relaxed)
244    }
245
246    /// 解析缓存未命中次数
247    pub fn parse_misses(&self) -> u64 {
248        self.parse_misses.load(Ordering::Relaxed)
249    }
250
251    /// 优化缓存命中次数
252    pub fn optimize_hits(&self) -> u64 {
253        self.optimize_hits.load(Ordering::Relaxed)
254    }
255
256    /// 优化缓存未命中次数
257    pub fn optimize_misses(&self) -> u64 {
258        self.optimize_misses.load(Ordering::Relaxed)
259    }
260
261    /// LRU 淘汰次数
262    pub fn evictions(&self) -> u64 {
263        self.evictions.load(Ordering::Relaxed)
264    }
265
266    /// 解析缓存命中率(0.0 ~ 1.0)
267    pub fn parse_hit_rate(&self) -> f64 {
268        let hits = self.parse_hits();
269        let misses = self.parse_misses();
270        let total = hits + misses;
271        if total == 0 {
272            0.0
273        } else {
274            hits as f64 / total as f64
275        }
276    }
277
278    /// 优化缓存命中率(0.0 ~ 1.0)
279    pub fn optimize_hit_rate(&self) -> f64 {
280        let hits = self.optimize_hits();
281        let misses = self.optimize_misses();
282        let total = hits + misses;
283        if total == 0 {
284            0.0
285        } else {
286            hits as f64 / total as f64
287        }
288    }
289}
290
291impl Default for PlanCacheStats {
292    fn default() -> Self {
293        Self::new()
294    }
295}
296
297// ─── PlanCacheStatsSnapshot ──────────────────────────────────────
298
299/// 统计快照(用于读取一致性视图)
300#[derive(Debug, Clone)]
301pub struct PlanCacheStatsSnapshot {
302    /// 解析缓存命中次数
303    pub parse_hits: u64,
304    /// 解析缓存未命中次数
305    pub parse_misses: u64,
306    /// 优化缓存命中次数
307    pub optimize_hits: u64,
308    /// 优化缓存未命中次数
309    pub optimize_misses: u64,
310    /// LRU 淘汰次数
311    pub evictions: u64,
312    /// 当前缓存条目数
313    pub size: usize,
314    /// 解析缓存命中率
315    pub parse_hit_rate: f64,
316    /// 优化缓存命中率
317    pub optimize_hit_rate: f64,
318}
319
320// ─── LruOrder64 ──────────────────────────────────────────────────
321
322/// u64 键的 LRU 双向链表(arena 实现,O(1) touch/remove/lru_key)
323struct LruOrder64 {
324    nodes: Vec<LruNode64>,
325    free_list: Vec<usize>,
326    index: HashMap<u64, usize>,
327    head: Option<usize>,
328    tail: Option<usize>,
329}
330
331struct LruNode64 {
332    key: u64,
333    prev: Option<usize>,
334    next: Option<usize>,
335}
336
337impl LruOrder64 {
338    fn new() -> Self {
339        Self {
340            nodes: Vec::new(),
341            free_list: Vec::new(),
342            index: HashMap::new(),
343            head: None,
344            tail: None,
345        }
346    }
347
348    fn touch(&mut self, key: u64) {
349        if let Some(&idx) = self.index.get(&key) {
350            self.unlink(idx);
351            self.link_tail(idx);
352        } else {
353            let idx = self.alloc_node(key);
354            self.link_tail(idx);
355            self.index.insert(key, idx);
356        }
357    }
358
359    fn remove(&mut self, key: u64) {
360        if let Some(idx) = self.index.remove(&key) {
361            self.unlink(idx);
362            self.free_node(idx);
363        }
364    }
365
366    fn lru_key(&self) -> Option<u64> {
367        self.head.map(|idx| self.nodes[idx].key)
368    }
369
370    fn clear(&mut self) {
371        self.nodes.clear();
372        self.free_list.clear();
373        self.index.clear();
374        self.head = None;
375        self.tail = None;
376    }
377
378    fn len(&self) -> usize {
379        self.index.len()
380    }
381
382    fn alloc_node(&mut self, key: u64) -> usize {
383        if let Some(idx) = self.free_list.pop() {
384            self.nodes[idx] = LruNode64 {
385                key,
386                prev: None,
387                next: None,
388            };
389            idx
390        } else {
391            self.nodes.push(LruNode64 {
392                key,
393                prev: None,
394                next: None,
395            });
396            self.nodes.len() - 1
397        }
398    }
399
400    fn free_node(&mut self, idx: usize) {
401        self.free_list.push(idx);
402    }
403
404    fn unlink(&mut self, idx: usize) {
405        let prev = self.nodes[idx].prev;
406        let next = self.nodes[idx].next;
407        match prev {
408            Some(p) => self.nodes[p].next = next,
409            None => self.head = next,
410        }
411        match next {
412            Some(n) => self.nodes[n].prev = prev,
413            None => self.tail = prev,
414        }
415        self.nodes[idx].prev = None;
416        self.nodes[idx].next = None;
417    }
418
419    fn link_tail(&mut self, idx: usize) {
420        match self.tail {
421            Some(t) => {
422                self.nodes[idx].prev = Some(t);
423                self.nodes[t].next = Some(idx);
424            }
425            None => {
426                self.head = Some(idx);
427            }
428        }
429        self.tail = Some(idx);
430    }
431}
432
433// ─── PlanCache ───────────────────────────────────────────────────
434
435/// 查询计划缓存
436///
437/// 双缓存架构:
438/// - `parse_cache`:SQL → AST(解析结果缓存)
439/// - `optimize_cache`:SQL → 优化分析(优化结果缓存)
440///
441/// 锁顺序约定:parse_cache → optimize_cache → access_order → table_index → stats
442/// (按此顺序加锁,避免死锁)
443pub struct PlanCache {
444    /// 解析缓存(hash → entry)
445    parse_cache: RwLock<HashMap<u64, PlanCacheEntry>>,
446    /// 优化缓存(hash → entry)
447    optimize_cache: RwLock<HashMap<u64, PlanCacheEntry>>,
448    /// LRU 访问顺序(arena 双向链表)
449    access_order: RwLock<LruOrder64>,
450    /// 表级失效索引(table → `Vec<hash>`)
451    table_index: RwLock<HashMap<String, Vec<u64>>>,
452    /// 统计计数器
453    stats: PlanCacheStats,
454    /// 最大缓存条目数
455    max_size: usize,
456    /// 默认 TTL
457    default_ttl: Option<Duration>,
458}
459
460impl PlanCache {
461    /// 创建新的查询计划缓存
462    ///
463    /// - `max_size`:最大缓存条目数(LRU 淘汰)
464    /// - `default_ttl`:默认 TTL(None 表示永不过期)
465    pub fn new(max_size: usize, default_ttl: Option<Duration>) -> Self {
466        Self {
467            parse_cache: RwLock::new(HashMap::new()),
468            optimize_cache: RwLock::new(HashMap::new()),
469            access_order: RwLock::new(LruOrder64::new()),
470            table_index: RwLock::new(HashMap::new()),
471            stats: PlanCacheStats::new(),
472            max_size,
473            default_ttl,
474        }
475    }
476
477    /// 获取或解析 SQL
478    ///
479    /// 命中缓存时返回 AST + stats.parse_hits++ + LRU touch。
480    /// 未命中时解析 SQL + 存入缓存 + stats.parse_misses++。
481    pub fn get_or_parse(&self, sql: &str) -> Result<Arc<Statement>, String> {
482        let key = PlanCacheKey::from_sql(sql);
483
484        {
485            let cache = self.parse_cache.read();
486            if let Some(entry) = cache.get(&key.hash) {
487                if !entry.is_expired() {
488                    if let Some(ast) = &entry.ast {
489                        self.stats.parse_hits.fetch_add(1, Ordering::Relaxed);
490                        self.access_order.write().touch(key.hash);
491                        return Ok(ast.clone());
492                    }
493                }
494            }
495        }
496
497        self.stats.parse_misses.fetch_add(1, Ordering::Relaxed);
498
499        let dialect = GenericDialect {};
500        let statements = Parser::parse_sql(&dialect, sql).map_err(|e| e.to_string())?;
501        if statements.is_empty() {
502            return Err("empty SQL".to_string());
503        }
504        let ast = Arc::new(statements.into_iter().next().unwrap());
505        let tables = SqlNormalizer::extract_tables(sql);
506
507        {
508            let mut access_order = self.access_order.write();
509            let mut cache = self.parse_cache.write();
510            let mut table_index = self.table_index.write();
511
512            if cache.len() >= self.max_size {
513                if let Some(lru_hash) = access_order.lru_key() {
514                    access_order.remove(lru_hash);
515                    cache.remove(&lru_hash);
516                    self.optimize_cache.write().remove(&lru_hash);
517                    Self::remove_from_table_index(&mut table_index, lru_hash);
518                    self.stats.evictions.fetch_add(1, Ordering::Relaxed);
519                }
520            }
521
522            let entry = PlanCacheEntry {
523                ast: Some(ast.clone()),
524                analysis: None,
525                created_at: Instant::now(),
526                tables: tables.clone(),
527                ttl: self.default_ttl,
528            };
529            cache.insert(key.hash, entry);
530            access_order.touch(key.hash);
531
532            for table in &tables {
533                table_index.entry(table.clone()).or_default().push(key.hash);
534            }
535        }
536
537        Ok(ast)
538    }
539
540    /// 获取或优化 SQL
541    ///
542    /// 命中缓存时返回优化分析 + stats.optimize_hits++ + LRU touch。
543    /// 未命中时返回 None + stats.optimize_misses++(调用方应执行优化后调用 `store_optimize`)。
544    pub fn get_or_optimize(&self, sql: &str) -> Option<Arc<String>> {
545        let key = PlanCacheKey::from_sql(sql);
546
547        {
548            let cache = self.optimize_cache.read();
549            if let Some(entry) = cache.get(&key.hash) {
550                if !entry.is_expired() {
551                    if let Some(analysis) = &entry.analysis {
552                        self.stats.optimize_hits.fetch_add(1, Ordering::Relaxed);
553                        self.access_order.write().touch(key.hash);
554                        return Some(analysis.clone());
555                    }
556                }
557            }
558        }
559
560        self.stats.optimize_misses.fetch_add(1, Ordering::Relaxed);
561        None
562    }
563
564    /// 存储优化结果
565    ///
566    /// 在 `get_or_optimize` 返回 None 后,调用方执行优化并存储结果。
567    pub fn store_optimize(&self, sql: &str, analysis: Arc<String>) {
568        let key = PlanCacheKey::from_sql(sql);
569        let tables = SqlNormalizer::extract_tables(sql);
570
571        let mut access_order = self.access_order.write();
572        let mut cache = self.optimize_cache.write();
573        let mut table_index = self.table_index.write();
574
575        if cache.len() >= self.max_size {
576            if let Some(lru_hash) = access_order.lru_key() {
577                access_order.remove(lru_hash);
578                cache.remove(&lru_hash);
579                self.parse_cache.write().remove(&lru_hash);
580                Self::remove_from_table_index(&mut table_index, lru_hash);
581                self.stats.evictions.fetch_add(1, Ordering::Relaxed);
582            }
583        }
584
585        let entry = PlanCacheEntry {
586            ast: None,
587            analysis: Some(analysis),
588            created_at: Instant::now(),
589            tables: tables.clone(),
590            ttl: self.default_ttl,
591        };
592        cache.insert(key.hash, entry);
593        access_order.touch(key.hash);
594
595        for table in &tables {
596            table_index.entry(table.clone()).or_default().push(key.hash);
597        }
598    }
599
600    /// 表级精确失效
601    ///
602    /// 失效所有引用指定表的缓存条目,返回失效条目数。
603    pub fn invalidate_table(&self, table: &str) -> usize {
604        let mut table_index = self.table_index.write();
605        let keys = table_index.remove(table).unwrap_or_default();
606
607        if keys.is_empty() {
608            return 0;
609        }
610
611        let count = keys.len();
612        let mut access_order = self.access_order.write();
613        let mut parse_cache = self.parse_cache.write();
614        let mut optimize_cache = self.optimize_cache.write();
615
616        for &hash in &keys {
617            access_order.remove(hash);
618            parse_cache.remove(&hash);
619            optimize_cache.remove(&hash);
620        }
621
622        for remaining_keys in table_index.values_mut() {
623            remaining_keys.retain(|k| !keys.contains(k));
624        }
625
626        self.stats
627            .evictions
628            .fetch_add(count as u64, Ordering::Relaxed);
629        count
630    }
631
632    /// 全量清空缓存
633    pub fn invalidate_all(&self) {
634        self.parse_cache.write().clear();
635        self.optimize_cache.write().clear();
636        self.access_order.write().clear();
637        self.table_index.write().clear();
638    }
639
640    /// 获取统计快照
641    pub fn stats(&self) -> PlanCacheStatsSnapshot {
642        let size = self.access_order.read().len();
643        PlanCacheStatsSnapshot {
644            parse_hits: self.stats.parse_hits(),
645            parse_misses: self.stats.parse_misses(),
646            optimize_hits: self.stats.optimize_hits(),
647            optimize_misses: self.stats.optimize_misses(),
648            evictions: self.stats.evictions(),
649            size,
650            parse_hit_rate: self.stats.parse_hit_rate(),
651            optimize_hit_rate: self.stats.optimize_hit_rate(),
652        }
653    }
654
655    /// 当前缓存大小
656    pub fn size(&self) -> usize {
657        self.access_order.read().len()
658    }
659
660    fn remove_from_table_index(table_index: &mut HashMap<String, Vec<u64>>, hash: u64) {
661        for keys in table_index.values_mut() {
662            keys.retain(|k| *k != hash);
663        }
664    }
665}
666
667impl Default for PlanCache {
668    fn default() -> Self {
669        Self::new(1024, None)
670    }
671}
672
673// ─── 单元测试 ─────────────────────────────────────────────────────
674
675#[cfg(test)]
676mod tests {
677    use super::*;
678
679    #[test]
680    fn test_sql_normalizer_basic() {
681        let sql1 = "SELECT * FROM users WHERE id = ?";
682        let sql2 = "select * from users where id = ?";
683        let n1 = SqlNormalizer::normalize(sql1);
684        let n2 = SqlNormalizer::normalize(sql2);
685        assert_eq!(n1, n2, "大小写差异应归一化");
686    }
687
688    #[test]
689    fn test_sql_normalizer_whitespace() {
690        let sql1 = "SELECT   *   FROM   users";
691        let sql2 = "SELECT * FROM users";
692        let n1 = SqlNormalizer::normalize(sql1);
693        let n2 = SqlNormalizer::normalize(sql2);
694        assert_eq!(n1, n2, "空白差异应归一化");
695    }
696
697    #[test]
698    fn test_sql_normalizer_different_semantics() {
699        let sql1 = "SELECT * FROM users WHERE id = ?";
700        let sql2 = "SELECT * FROM orders WHERE id = ?";
701        let n1 = SqlNormalizer::normalize(sql1);
702        let n2 = SqlNormalizer::normalize(sql2);
703        assert_ne!(n1, n2, "不同表名应产生不同归一化");
704    }
705
706    #[test]
707    fn test_sql_normalizer_parse_error_fallback() {
708        let sql = "this is not valid sql !!!";
709        let normalized = SqlNormalizer::normalize(sql);
710        assert!(!normalized.is_empty());
711    }
712
713    #[test]
714    fn test_extract_tables_select() {
715        let sql = "SELECT * FROM users JOIN orders ON users.id = orders.user_id";
716        let tables = SqlNormalizer::extract_tables(sql);
717        assert!(tables.contains(&"users".to_string()));
718        assert!(tables.contains(&"orders".to_string()));
719    }
720
721    #[test]
722    fn test_extract_tables_insert() {
723        let sql = "INSERT INTO products (name) VALUES (?)";
724        let tables = SqlNormalizer::extract_tables(sql);
725        assert!(tables.contains(&"products".to_string()));
726    }
727
728    #[test]
729    fn test_extract_tables_update() {
730        let sql = "UPDATE products SET name = ? WHERE id = ?";
731        let tables = SqlNormalizer::extract_tables(sql);
732        assert!(tables.contains(&"products".to_string()));
733    }
734
735    #[test]
736    fn test_extract_tables_delete() {
737        let sql = "DELETE FROM products WHERE id = ?";
738        let tables = SqlNormalizer::extract_tables(sql);
739        assert!(
740            tables.iter().any(|t| t.contains("products")),
741            "应包含 products 表,实际: {:?}",
742            tables
743        );
744    }
745
746    #[test]
747    fn test_plan_cache_key_same_sql() {
748        let k1 = PlanCacheKey::from_sql("SELECT * FROM users WHERE id = ?");
749        let k2 = PlanCacheKey::from_sql("select * from users where id = ?");
750        assert_eq!(k1.hash, k2.hash, "相同 SQL 模板应产生相同 hash");
751    }
752
753    #[test]
754    fn test_plan_cache_key_different_sql() {
755        let k1 = PlanCacheKey::from_sql("SELECT * FROM users");
756        let k2 = PlanCacheKey::from_sql("SELECT * FROM orders");
757        assert_ne!(k1.hash, k2.hash, "不同 SQL 应产生不同 hash");
758    }
759
760    #[test]
761    fn test_plan_cache_key_no_sensitive_data() {
762        let key = PlanCacheKey::from_sql("SELECT * FROM users WHERE password = ?");
763        assert!(
764            !key.sql_normalized.contains("secret123"),
765            "参数化查询缓存键不应包含参数值"
766        );
767    }
768
769    #[test]
770    fn test_plan_cache_new() {
771        let cache = PlanCache::new(1024, None);
772        assert_eq!(cache.size(), 0);
773        let stats = cache.stats();
774        assert_eq!(stats.size, 0);
775        assert_eq!(stats.parse_hits, 0);
776        assert_eq!(stats.parse_misses, 0);
777    }
778
779    #[test]
780    fn test_plan_cache_default() {
781        let cache = PlanCache::default();
782        assert_eq!(cache.size(), 0);
783    }
784
785    #[test]
786    fn test_get_or_parse_hit() {
787        let cache = PlanCache::new(100, None);
788        let sql = "SELECT * FROM users WHERE id = ?";
789
790        let ast1 = cache.get_or_parse(sql).expect("parse");
791        assert_eq!(cache.stats().parse_misses, 1, "首次应 miss");
792
793        let ast2 = cache.get_or_parse(sql).expect("parse");
794        assert_eq!(cache.stats().parse_hits, 1, "第二次应 hit");
795        assert_eq!(cache.stats().parse_misses, 1, "misses 不变");
796
797        assert!(Arc::ptr_eq(&ast1, &ast2), "命中应返回相同 Arc");
798    }
799
800    #[test]
801    fn test_get_or_parse_different_params_same_template() {
802        let cache = PlanCache::new(100, None);
803        let sql1 = "SELECT * FROM users WHERE id = ?";
804        let sql2 = "select * from users where id = ?";
805
806        cache.get_or_parse(sql1).expect("parse");
807        cache.get_or_parse(sql2).expect("parse");
808
809        assert_eq!(cache.stats().parse_hits, 1, "相同模板不同写法应命中");
810        assert_eq!(cache.stats().parse_misses, 1);
811    }
812
813    #[test]
814    fn test_get_or_parse_different_sql() {
815        let cache = PlanCache::new(100, None);
816        cache.get_or_parse("SELECT * FROM users").expect("parse");
817        cache.get_or_parse("SELECT * FROM orders").expect("parse");
818
819        assert_eq!(cache.stats().parse_misses, 2, "不同 SQL 应各 miss 一次");
820        assert_eq!(cache.stats().parse_hits, 0);
821    }
822
823    #[test]
824    fn test_get_or_optimize_miss_then_store_then_hit() {
825        let cache = PlanCache::new(100, None);
826        let sql = "SELECT * FROM users WHERE id = ?";
827
828        assert!(cache.get_or_optimize(sql).is_none(), "首次应 miss");
829        assert_eq!(cache.stats().optimize_misses, 1);
830
831        cache.store_optimize(sql, Arc::new("optimized plan".to_string()));
832
833        let result = cache.get_or_optimize(sql);
834        assert!(result.is_some(), "存储后应 hit");
835        assert_eq!(cache.stats().optimize_hits, 1);
836        assert_eq!(*result.unwrap().as_ref(), "optimized plan");
837    }
838
839    #[test]
840    fn test_invalidate_table_precise() {
841        let cache = PlanCache::new(100, None);
842        cache.get_or_parse("SELECT * FROM users").expect("parse");
843        cache.get_or_parse("SELECT * FROM orders").expect("parse");
844        assert_eq!(cache.size(), 2);
845
846        let evicted = cache.invalidate_table("users");
847        assert_eq!(evicted, 1, "应失效 1 条");
848        assert_eq!(cache.size(), 1, "应剩余 1 条");
849
850        let stats = cache.stats();
851        assert!(stats.parse_hits == 0, "orders 缓存应不受影响");
852        cache.get_or_parse("SELECT * FROM orders").expect("parse");
853        assert_eq!(cache.stats().parse_hits, 1, "orders 应命中缓存");
854    }
855
856    #[test]
857    fn test_invalidate_table_nonexistent() {
858        let cache = PlanCache::new(100, None);
859        cache.get_or_parse("SELECT * FROM users").expect("parse");
860        let evicted = cache.invalidate_table("nonexistent");
861        assert_eq!(evicted, 0, "不存在的表应返回 0");
862        assert_eq!(cache.size(), 1, "缓存不应受影响");
863    }
864
865    #[test]
866    fn test_invalidate_all() {
867        let cache = PlanCache::new(100, None);
868        cache.get_or_parse("SELECT * FROM users").expect("parse");
869        cache.get_or_parse("SELECT * FROM orders").expect("parse");
870        assert_eq!(cache.size(), 2);
871
872        cache.invalidate_all();
873        assert_eq!(cache.size(), 0, "全量清空后 size 应为 0");
874    }
875
876    #[test]
877    fn test_lru_eviction() {
878        let cache = PlanCache::new(3, None);
879        cache.get_or_parse("SELECT * FROM t1").expect("parse");
880        cache.get_or_parse("SELECT * FROM t2").expect("parse");
881        cache.get_or_parse("SELECT * FROM t3").expect("parse");
882        assert_eq!(cache.size(), 3);
883
884        cache.get_or_parse("SELECT * FROM t4").expect("parse");
885        assert_eq!(cache.size(), 3, "max_size=3 应保持 3 条");
886        assert!(cache.stats().evictions >= 1, "应有淘汰");
887
888        cache.get_or_parse("SELECT * FROM t1").expect("parse");
889        assert!(cache.stats().parse_misses >= 4, "t1 被淘汰后应重新 miss");
890    }
891
892    #[test]
893    fn test_lru_eviction_max_size_1() {
894        let cache = PlanCache::new(1, None);
895        cache.get_or_parse("SELECT * FROM t1").expect("parse");
896        assert_eq!(cache.size(), 1);
897
898        cache.get_or_parse("SELECT * FROM t2").expect("parse");
899        assert_eq!(cache.size(), 1, "max_size=1 应保持 1 条");
900
901        cache.get_or_parse("SELECT * FROM t1").expect("parse");
902        assert!(cache.stats().parse_misses >= 3, "t1 应被淘汰后重新 miss");
903    }
904
905    #[test]
906    fn test_stats_hit_rate() {
907        let cache = PlanCache::new(100, None);
908        let sql = "SELECT * FROM users";
909
910        cache.get_or_parse(sql).expect("parse");
911        cache.get_or_parse(sql).expect("parse");
912        cache.get_or_parse(sql).expect("parse");
913
914        let stats = cache.stats();
915        assert_eq!(stats.parse_hits, 2);
916        assert_eq!(stats.parse_misses, 1);
917        assert!((stats.parse_hit_rate - (2.0 / 3.0)).abs() < 0.001);
918    }
919
920    #[test]
921    fn test_stats_hit_rate_empty() {
922        let cache = PlanCache::new(100, None);
923        let stats = cache.stats();
924        assert_eq!(stats.parse_hit_rate, 0.0, "空缓存命中率应为 0.0");
925        assert_eq!(stats.optimize_hit_rate, 0.0);
926    }
927
928    #[test]
929    fn test_ttl_expiration() {
930        let cache = PlanCache::new(100, Some(Duration::from_nanos(1)));
931        let sql = "SELECT * FROM users";
932
933        cache.get_or_parse(sql).expect("parse");
934        std::thread::sleep(Duration::from_millis(10));
935
936        cache.get_or_parse(sql).expect("parse");
937        assert!(cache.stats().parse_misses >= 2, "TTL 过期后应重新 miss");
938    }
939
940    #[test]
941    fn test_plan_cache_concurrent_same_sql() {
942        use std::sync::Arc;
943        use std::thread;
944
945        let cache = Arc::new(PlanCache::new(100, None));
946        let sql = "SELECT * FROM users WHERE id = ?";
947        let mut handles = Vec::new();
948
949        for _ in 0..10 {
950            let cache = cache.clone();
951            handles.push(thread::spawn(move || {
952                cache.get_or_parse(sql).expect("parse");
953            }));
954        }
955
956        for h in handles {
957            h.join().expect("thread");
958        }
959
960        assert!(cache.size() >= 1, "并发后应至少有 1 条缓存");
961        assert!(
962            cache.stats().parse_misses + cache.stats().parse_hits >= 10,
963            "应有 10 次访问记录"
964        );
965    }
966}