Skip to main content

driven/state/
shared_rules.rs

1//! Shared Rule Storage
2//!
3//! Thread-safe rule storage with zero-copy access.
4
5use std::collections::HashMap;
6use std::sync::{Arc, RwLock};
7
8use super::DirtyBits;
9
10/// Reference to a shared rule
11#[derive(Debug, Clone)]
12pub struct RuleRef {
13    /// Rule ID
14    pub id: u32,
15    /// Rule content
16    content: Arc<[u8]>,
17    /// Version number
18    pub version: u64,
19}
20
21impl RuleRef {
22    /// Create a new rule reference
23    pub fn new(id: u32, content: Vec<u8>) -> Self {
24        Self {
25            id,
26            content: content.into(),
27            version: 1,
28        }
29    }
30
31    /// Get content as bytes
32    pub fn as_bytes(&self) -> &[u8] {
33        &self.content
34    }
35
36    /// Get content as string (if valid UTF-8)
37    pub fn as_str(&self) -> Option<&str> {
38        std::str::from_utf8(&self.content).ok()
39    }
40
41    /// Get content length
42    pub fn len(&self) -> usize {
43        self.content.len()
44    }
45
46    /// Check if empty
47    pub fn is_empty(&self) -> bool {
48        self.content.is_empty()
49    }
50}
51
52/// Shared rule storage with dirty tracking
53#[derive(Debug)]
54pub struct SharedRules {
55    /// Rules by ID
56    rules: RwLock<HashMap<u32, RuleRef>>,
57    /// Dirty bit tracker
58    dirty: DirtyBits,
59    /// Next rule ID
60    next_id: std::sync::atomic::AtomicU32,
61    /// Global version
62    version: std::sync::atomic::AtomicU64,
63}
64
65impl SharedRules {
66    /// Create new shared storage
67    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    /// Insert a new rule
77    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    /// Get a rule by ID
92    pub fn get(&self, id: u32) -> Option<RuleRef> {
93        self.rules.read().unwrap().get(&id).cloned()
94    }
95
96    /// Update a rule
97    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    /// Remove a rule
115    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    /// Get all rule IDs
126    pub fn ids(&self) -> Vec<u32> {
127        self.rules.read().unwrap().keys().copied().collect()
128    }
129
130    /// Get rule count
131    pub fn len(&self) -> usize {
132        self.rules.read().unwrap().len()
133    }
134
135    /// Check if empty
136    pub fn is_empty(&self) -> bool {
137        self.rules.read().unwrap().is_empty()
138    }
139
140    /// Get dirty tracker
141    pub fn dirty(&self) -> &DirtyBits {
142        &self.dirty
143    }
144
145    /// Get current version
146    pub fn version(&self) -> u64 {
147        self.version.load(std::sync::atomic::Ordering::SeqCst)
148    }
149
150    /// Check if changes exist
151    pub fn has_changes(&self) -> bool {
152        self.dirty.has_changes()
153    }
154
155    /// Mark as synced
156    pub fn mark_synced(&self) {
157        self.dirty.mark_synced();
158    }
159
160    /// Get dirty rule IDs
161    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    /// Clear all rules
173    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
185/// Thread-safe shared rules handle
186pub type SharedRulesHandle = Arc<SharedRules>;
187
188/// Create a new shared rules handle
189pub 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}