driven/state/
shared_rules.rs1use std::collections::HashMap;
6use std::sync::{Arc, RwLock};
7
8use super::DirtyBits;
9
10#[derive(Debug, Clone)]
12pub struct RuleRef {
13 pub id: u32,
15 content: Arc<[u8]>,
17 pub version: u64,
19}
20
21impl RuleRef {
22 pub fn new(id: u32, content: Vec<u8>) -> Self {
24 Self {
25 id,
26 content: content.into(),
27 version: 1,
28 }
29 }
30
31 pub fn as_bytes(&self) -> &[u8] {
33 &self.content
34 }
35
36 pub fn as_str(&self) -> Option<&str> {
38 std::str::from_utf8(&self.content).ok()
39 }
40
41 pub fn len(&self) -> usize {
43 self.content.len()
44 }
45
46 pub fn is_empty(&self) -> bool {
48 self.content.is_empty()
49 }
50}
51
52#[derive(Debug)]
54pub struct SharedRules {
55 rules: RwLock<HashMap<u32, RuleRef>>,
57 dirty: DirtyBits,
59 next_id: std::sync::atomic::AtomicU32,
61 version: std::sync::atomic::AtomicU64,
63}
64
65impl SharedRules {
66 pub fn new() -> Self {
68 Self {
69 rules: RwLock::new(HashMap::new()),
70 dirty: DirtyBits::new(),
71 next_id: std::sync::atomic::AtomicU32::new(1),
72 version: std::sync::atomic::AtomicU64::new(1),
73 }
74 }
75
76 pub fn insert(&self, content: Vec<u8>) -> u32 {
78 let id = self
79 .next_id
80 .fetch_add(1, std::sync::atomic::Ordering::SeqCst);
81 let rule = RuleRef::new(id, content);
82
83 self.rules.write().unwrap().insert(id, rule);
84 self.dirty.dirty_standard((id % 64) as u8);
85 self.version
86 .fetch_add(1, std::sync::atomic::Ordering::SeqCst);
87
88 id
89 }
90
91 pub fn get(&self, id: u32) -> Option<RuleRef> {
93 self.rules.read().unwrap().get(&id).cloned()
94 }
95
96 pub fn update(&self, id: u32, content: Vec<u8>) -> bool {
98 let mut rules = self.rules.write().unwrap();
99 if let Some(rule) = rules.get_mut(&id) {
100 *rule = RuleRef {
101 id,
102 content: content.into(),
103 version: rule.version + 1,
104 };
105 self.dirty.dirty_standard((id % 64) as u8);
106 self.version
107 .fetch_add(1, std::sync::atomic::Ordering::SeqCst);
108 true
109 } else {
110 false
111 }
112 }
113
114 pub fn remove(&self, id: u32) -> Option<RuleRef> {
116 let removed = self.rules.write().unwrap().remove(&id);
117 if removed.is_some() {
118 self.dirty.dirty_standard((id % 64) as u8);
119 self.version
120 .fetch_add(1, std::sync::atomic::Ordering::SeqCst);
121 }
122 removed
123 }
124
125 pub fn ids(&self) -> Vec<u32> {
127 self.rules.read().unwrap().keys().copied().collect()
128 }
129
130 pub fn len(&self) -> usize {
132 self.rules.read().unwrap().len()
133 }
134
135 pub fn is_empty(&self) -> bool {
137 self.rules.read().unwrap().is_empty()
138 }
139
140 pub fn dirty(&self) -> &DirtyBits {
142 &self.dirty
143 }
144
145 pub fn version(&self) -> u64 {
147 self.version.load(std::sync::atomic::Ordering::SeqCst)
148 }
149
150 pub fn has_changes(&self) -> bool {
152 self.dirty.has_changes()
153 }
154
155 pub fn mark_synced(&self) {
157 self.dirty.mark_synced();
158 }
159
160 pub fn dirty_ids(&self) -> Vec<u32> {
162 let rules = self.rules.read().unwrap();
163 let dirty_indices = self.dirty.standards.dirty_indices();
164
165 rules
166 .keys()
167 .filter(|&&id| dirty_indices.contains(&((id % 64) as u8)))
168 .copied()
169 .collect()
170 }
171
172 pub fn clear(&self) {
174 self.rules.write().unwrap().clear();
175 self.dirty.mark_synced();
176 }
177}
178
179impl Default for SharedRules {
180 fn default() -> Self {
181 Self::new()
182 }
183}
184
185pub type SharedRulesHandle = Arc<SharedRules>;
187
188pub fn create_shared_rules() -> SharedRulesHandle {
190 Arc::new(SharedRules::new())
191}
192
193#[cfg(test)]
194mod tests {
195 use super::*;
196
197 #[test]
198 fn test_insert_and_get() {
199 let rules = SharedRules::new();
200
201 let id = rules.insert(b"rule content".to_vec());
202 let rule = rules.get(id).unwrap();
203
204 assert_eq!(rule.as_bytes(), b"rule content");
205 assert_eq!(rule.version, 1);
206 }
207
208 #[test]
209 fn test_update() {
210 let rules = SharedRules::new();
211
212 let id = rules.insert(b"original".to_vec());
213 rules.update(id, b"updated".to_vec());
214
215 let rule = rules.get(id).unwrap();
216 assert_eq!(rule.as_bytes(), b"updated");
217 assert_eq!(rule.version, 2);
218 }
219
220 #[test]
221 fn test_dirty_tracking() {
222 let rules = SharedRules::new();
223
224 assert!(!rules.has_changes());
225
226 rules.insert(b"test".to_vec());
227 assert!(rules.has_changes());
228
229 rules.mark_synced();
230 assert!(!rules.has_changes());
231 }
232
233 #[test]
234 fn test_thread_safety() {
235 let rules = Arc::new(SharedRules::new());
236
237 let handles: Vec<_> = (0..10)
238 .map(|i| {
239 let rules = rules.clone();
240 std::thread::spawn(move || {
241 rules.insert(format!("rule {}", i).into_bytes());
242 })
243 })
244 .collect();
245
246 for handle in handles {
247 handle.join().unwrap();
248 }
249
250 assert_eq!(rules.len(), 10);
251 }
252}