oxicode_sdk/coordination/
shared_memory.rs1use parking_lot::RwLock;
4use serde::{Deserialize, Serialize};
5use std::collections::HashMap;
6use tokio::sync::broadcast;
7
8#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
10pub struct MemoryKey {
11 pub namespace: String,
13 pub key: String,
15}
16
17impl MemoryKey {
18 pub fn new(namespace: impl Into<String>, key: impl Into<String>) -> Self {
20 Self {
21 namespace: namespace.into(),
22 key: key.into(),
23 }
24 }
25}
26
27#[derive(Debug, Clone, Serialize, Deserialize)]
29pub struct MemoryEntry {
30 pub value: serde_json::Value,
32 pub version: u64,
34 pub modified_at_ms: u64,
36 pub modified_by: String,
38}
39
40#[derive(Debug, Clone, Serialize, Deserialize)]
42pub enum MemoryEvent {
43 Written {
45 key: MemoryKey,
47 version: u64,
49 author: String,
51 },
52 Deleted {
54 key: MemoryKey,
56 },
57}
58
59pub struct SharedMemory {
64 data: RwLock<HashMap<MemoryKey, MemoryEntry>>,
65 tx: broadcast::Sender<MemoryEvent>,
66}
67
68impl SharedMemory {
69 pub fn new() -> Self {
71 let (tx, _) = broadcast::channel(256);
72 Self {
73 data: RwLock::new(HashMap::new()),
74 tx,
75 }
76 }
77
78 pub fn read(&self, key: &MemoryKey) -> Option<serde_json::Value> {
80 self.data.read().get(key).map(|e| e.value.clone())
81 }
82
83 pub fn read_entry(&self, key: &MemoryKey) -> Option<MemoryEntry> {
85 self.data.read().get(key).cloned()
86 }
87
88 pub fn write(
93 &self,
94 key: &MemoryKey,
95 value: serde_json::Value,
96 author: &str,
97 expected_version: Option<u64>,
98 ) -> Result<u64, crate::error::SdkError> {
99 let mut data = self.data.write();
100
101 if let Some(expected) = expected_version {
102 if let Some(entry) = data.get(key) {
103 if entry.version != expected {
104 return Err(crate::error::SdkError::VersionConflict {
105 key: format!("{}:{}", key.namespace, key.key),
106 expected,
107 current: entry.version,
108 });
109 }
110 } else if expected != 0 {
111 return Err(crate::error::SdkError::VersionConflict {
112 key: format!("{}:{}", key.namespace, key.key),
113 expected,
114 current: 0,
115 });
116 }
117 }
118
119 let current_version = data.get(key).map(|e| e.version).unwrap_or(0);
120 let new_version = current_version + 1;
121
122 data.insert(
123 key.clone(),
124 MemoryEntry {
125 value,
126 version: new_version,
127 modified_at_ms: now_ms(),
128 modified_by: author.to_string(),
129 },
130 );
131
132 let _ = self.tx.send(MemoryEvent::Written {
133 key: key.clone(),
134 version: new_version,
135 author: author.to_string(),
136 });
137
138 Ok(new_version)
139 }
140
141 pub fn increment(&self, key: &MemoryKey, delta: i64, author: &str) -> i64 {
143 let mut data = self.data.write();
144 let entry = data.entry(key.clone()).or_insert(MemoryEntry {
145 value: serde_json::json!(0),
146 version: 0,
147 modified_at_ms: 0,
148 modified_by: String::new(),
149 });
150
151 let current = entry.value.as_i64().unwrap_or(0);
152 let new_val = current + delta;
153 entry.value = serde_json::json!(new_val);
154 entry.version += 1;
155 entry.modified_at_ms = now_ms();
156 entry.modified_by = author.to_string();
157 new_val
158 }
159
160 pub fn delete(&self, key: &MemoryKey) -> bool {
162 let removed = self.data.write().remove(key).is_some();
163 if removed {
164 let _ = self.tx.send(MemoryEvent::Deleted { key: key.clone() });
165 }
166 removed
167 }
168
169 pub fn list_namespace(&self, namespace: &str) -> Vec<MemoryKey> {
171 self.data
172 .read()
173 .keys()
174 .filter(|k| k.namespace == namespace)
175 .cloned()
176 .collect()
177 }
178
179 pub fn subscribe(&self) -> broadcast::Receiver<MemoryEvent> {
181 self.tx.subscribe()
182 }
183}
184
185impl Default for SharedMemory {
186 fn default() -> Self {
187 Self::new()
188 }
189}
190
191fn now_ms() -> u64 {
192 std::time::SystemTime::now()
193 .duration_since(std::time::UNIX_EPOCH)
194 .map(|d| d.as_millis() as u64)
195 .unwrap_or(0)
196}
197
198#[cfg(test)]
199mod tests {
200 use super::*;
201
202 #[test]
203 fn write_and_read() {
204 let mem = SharedMemory::new();
205 let key = MemoryKey::new("ns", "counter");
206 mem.write(&key, serde_json::json!(42), "agent-1", None)
207 .unwrap();
208 assert_eq!(mem.read(&key), Some(serde_json::json!(42)));
209 }
210
211 #[test]
212 fn version_increments() {
213 let mem = SharedMemory::new();
214 let key = MemoryKey::new("ns", "val");
215 let v1 = mem.write(&key, serde_json::json!("a"), "a1", None).unwrap();
216 let v2 = mem.write(&key, serde_json::json!("b"), "a2", None).unwrap();
217 assert_eq!(v1, 1);
218 assert_eq!(v2, 2);
219 let entry = mem.read_entry(&key).unwrap();
220 assert_eq!(entry.version, 2);
221 }
222
223 #[test]
224 fn optimistic_lock_success() {
225 let mem = SharedMemory::new();
226 let key = MemoryKey::new("ns", "val");
227 let v1 = mem.write(&key, serde_json::json!("a"), "a1", None).unwrap();
228 let v2 = mem
229 .write(&key, serde_json::json!("b"), "a2", Some(v1))
230 .unwrap();
231 assert_eq!(v2, 2);
232 }
233
234 #[test]
235 fn optimistic_lock_conflict() {
236 let mem = SharedMemory::new();
237 let key = MemoryKey::new("ns", "val");
238 let _v1 = mem.write(&key, serde_json::json!("a"), "a1", None).unwrap();
239 let result = mem.write(&key, serde_json::json!("b"), "a2", Some(99));
240 assert!(result.is_err());
241 match result.unwrap_err() {
242 crate::error::SdkError::VersionConflict {
243 expected, current, ..
244 } => {
245 assert_eq!(expected, 99);
246 assert_eq!(current, 1);
247 }
248 _ => panic!("Expected VersionConflict"),
249 }
250 }
251
252 #[test]
253 fn atomic_increment() {
254 let mem = SharedMemory::new();
255 let key = MemoryKey::new("ns", "counter");
256 assert_eq!(mem.increment(&key, 5, "a1"), 5);
257 assert_eq!(mem.increment(&key, 3, "a2"), 8);
258 assert_eq!(mem.read(&key), Some(serde_json::json!(8)));
259 }
260
261 #[test]
262 fn delete_key() {
263 let mem = SharedMemory::new();
264 let key = MemoryKey::new("ns", "val");
265 mem.write(&key, serde_json::json!(1), "a1", None).unwrap();
266 assert!(mem.delete(&key));
267 assert!(mem.read(&key).is_none());
268 assert!(!mem.delete(&key)); }
270
271 #[test]
272 fn list_namespace() {
273 let mem = SharedMemory::new();
274 mem.write(
275 &MemoryKey::new("reviews", "a"),
276 serde_json::json!(1),
277 "a1",
278 None,
279 )
280 .unwrap();
281 mem.write(
282 &MemoryKey::new("reviews", "b"),
283 serde_json::json!(2),
284 "a1",
285 None,
286 )
287 .unwrap();
288 mem.write(
289 &MemoryKey::new("other", "c"),
290 serde_json::json!(3),
291 "a1",
292 None,
293 )
294 .unwrap();
295 let keys = mem.list_namespace("reviews");
296 assert_eq!(keys.len(), 2);
297 }
298
299 #[test]
300 fn subscribe_events() {
301 let mem = SharedMemory::new();
302 let mut rx = mem.subscribe();
303 let key = MemoryKey::new("ns", "val");
304 mem.write(&key, serde_json::json!(1), "a1", None).unwrap();
305 let event = rx.try_recv().unwrap();
306 assert!(matches!(event, MemoryEvent::Written { .. }));
307 }
308}