sz_rust_orm_facade/data_scope/
cache.rs1use 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#[async_trait]
15pub trait DeptTreeProvider: Send + Sync {
16 async fn sub_depts(&self, dept_id: i64) -> Result<Vec<i64>, DataScopeError>;
18}
19
20struct CachedEntry {
22 depts: Vec<i64>,
23 inserted_at: Instant,
24}
25
26pub struct DeptTreeCache {
28 provider: Arc<dyn DeptTreeProvider>,
29 cache: DashMap<i64, CachedEntry>,
30 ttl: Duration,
31}
32
33impl DeptTreeCache {
34 pub fn new(provider: Arc<dyn DeptTreeProvider>, ttl: Duration) -> Self {
36 Self {
37 provider,
38 cache: DashMap::new(),
39 ttl,
40 }
41 }
42
43 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 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 pub fn invalidate(&self, dept_id: i64) {
94 self.cache.remove(&dept_id);
95 }
96
97 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}