Skip to main content

sz_rust_orm_facade/data_scope/
cache.rs

1//! 部门树缓存 — DeptTreeCache
2//!
3//! 使用 DashMap 分片锁实现高并发读,TTL 过期自动刷新。
4//! 递归展开算法维护 HashSet 检测循环引用。
5
6use crate::data_scope::error::DataScopeError;
7use async_trait::async_trait;
8use dashmap::DashMap;
9use std::collections::{HashSet, VecDeque};
10use std::sync::Arc;
11use std::time::{Duration, Instant};
12
13/// 部门树数据源 trait
14#[async_trait]
15pub trait DeptTreeProvider: Send + Sync {
16    /// 获取指定部门的直接子部门 ID 列表
17    async fn sub_depts(&self, dept_id: i64) -> Result<Vec<i64>, DataScopeError>;
18}
19
20/// 缓存条目
21struct CachedEntry {
22    depts: Vec<i64>,
23    inserted_at: Instant,
24}
25
26/// 部门树缓存
27pub struct DeptTreeCache {
28    provider: Arc<dyn DeptTreeProvider>,
29    cache: DashMap<i64, CachedEntry>,
30    ttl: Duration,
31}
32
33impl DeptTreeCache {
34    /// 创建缓存
35    pub fn new(provider: Arc<dyn DeptTreeProvider>, ttl: Duration) -> Self {
36        Self {
37            provider,
38            cache: DashMap::new(),
39            ttl,
40        }
41    }
42
43    /// 获取指定部门及其所有子部门 ID 列表(含自身)
44    ///
45    /// 缓存命中(键存在且未过期)直接返回;未命中调用 provider 递归展开。
46    pub async fn get_with_sub(&self, dept_id: i64) -> Result<Vec<i64>, DataScopeError> {
47        if let Some(entry) = self.cache.get(&dept_id) {
48            if entry.inserted_at.elapsed() < self.ttl {
49                return Ok(entry.depts.clone());
50            }
51        }
52
53        let result = self.expand_with_sub(dept_id).await?;
54        self.cache.insert(
55            dept_id,
56            CachedEntry {
57                depts: result.clone(),
58                inserted_at: Instant::now(),
59            },
60        );
61        Ok(result)
62    }
63
64    /// 递归展开部门树(BFS + 循环引用检测)
65    async fn expand_with_sub(&self, dept_id: i64) -> Result<Vec<i64>, DataScopeError> {
66        let mut result = vec![dept_id];
67        let mut visited = HashSet::new();
68        visited.insert(dept_id);
69        let mut queue = VecDeque::new();
70        queue.push_back(dept_id);
71
72        while let Some(current) = queue.pop_front() {
73            let subs = self.provider.sub_depts(current).await?;
74            for sub in subs {
75                if visited.contains(&sub) {
76                    tracing::warn!(
77                        target: "data_scope",
78                        "circular reference detected: dept {} -> dept {} (skipping)",
79                        current, sub
80                    );
81                    continue;
82                }
83                visited.insert(sub);
84                result.push(sub);
85                queue.push_back(sub);
86            }
87        }
88
89        Ok(result)
90    }
91
92    /// 失效单键
93    pub fn invalidate(&self, dept_id: i64) {
94        self.cache.remove(&dept_id);
95    }
96
97    /// 全量失效
98    pub fn invalidate_all(&self) {
99        self.cache.clear();
100    }
101}
102
103#[cfg(test)]
104mod tests {
105    use super::*;
106
107    struct MockProvider {
108        tree: std::collections::HashMap<i64, Vec<i64>>,
109    }
110
111    #[async_trait]
112    impl DeptTreeProvider for MockProvider {
113        async fn sub_depts(&self, dept_id: i64) -> Result<Vec<i64>, DataScopeError> {
114            Ok(self.tree.get(&dept_id).cloned().unwrap_or_default())
115        }
116    }
117
118    #[tokio::test]
119    async fn test_cache_hit() {
120        let provider = Arc::new(MockProvider {
121            tree: [(5, vec![6, 7]), (6, vec![8]), (7, vec![])]
122                .into_iter()
123                .collect(),
124        });
125        let cache = DeptTreeCache::new(provider, Duration::from_secs(300));
126        let result1 = cache.get_with_sub(5).await.unwrap();
127        assert_eq!(result1, vec![5, 6, 7, 8]);
128        let result2 = cache.get_with_sub(5).await.unwrap();
129        assert_eq!(result2, result1);
130    }
131
132    #[tokio::test]
133    async fn test_circular_reference() {
134        let provider = Arc::new(MockProvider {
135            tree: [(1, vec![2]), (2, vec![1])].into_iter().collect(),
136        });
137        let cache = DeptTreeCache::new(provider, Duration::from_secs(300));
138        let result = cache.get_with_sub(1).await.unwrap();
139        assert!(result.contains(&1));
140        assert!(result.contains(&2));
141    }
142
143    #[tokio::test]
144    async fn test_invalidate() {
145        let provider = Arc::new(MockProvider {
146            tree: [(5, vec![6])].into_iter().collect(),
147        });
148        let cache = DeptTreeCache::new(provider, Duration::from_secs(300));
149        let _ = cache.get_with_sub(5).await.unwrap();
150        cache.invalidate(5);
151        let result = cache.get_with_sub(5).await.unwrap();
152        assert_eq!(result, vec![5, 6]);
153    }
154
155    #[tokio::test]
156    async fn test_invalidate_all() {
157        let provider = Arc::new(MockProvider {
158            tree: [(5, vec![6]), (10, vec![11])].into_iter().collect(),
159        });
160        let cache = DeptTreeCache::new(provider, Duration::from_secs(300));
161        let _ = cache.get_with_sub(5).await.unwrap();
162        let _ = cache.get_with_sub(10).await.unwrap();
163        cache.invalidate_all();
164        let _ = cache.get_with_sub(5).await.unwrap();
165    }
166
167    #[tokio::test]
168    async fn test_cache_ttl_expired_requeries() {
169        let provider = Arc::new(MockProvider {
170            tree: [(5, vec![6])].into_iter().collect(),
171        });
172        let cache = DeptTreeCache::new(provider, Duration::from_millis(1));
173        let r1 = cache.get_with_sub(5).await.unwrap();
174        assert_eq!(r1, vec![5, 6]);
175        tokio::time::sleep(Duration::from_millis(20)).await;
176        let r2 = cache.get_with_sub(5).await.unwrap();
177        assert_eq!(r2, vec![5, 6]);
178    }
179}