1use super::error::SignalError;
37use super::signal::Signal;
38use async_trait::async_trait;
39use parking_lot::RwLock;
40use serde::{Deserialize, Serialize};
41use std::collections::VecDeque;
42use std::sync::Arc;
43use std::time::SystemTime;
44
45#[derive(Debug, Clone, Serialize, Deserialize)]
49pub struct StoredSignal<T> {
50 pub id: u64,
52 pub signal_name: String,
54 pub timestamp: SystemTime,
56 pub payload: T,
58}
59
60impl<T> StoredSignal<T> {
61 pub fn new(id: u64, signal_name: String, payload: T) -> Self {
63 Self {
64 id,
65 signal_name,
66 timestamp: SystemTime::now(),
67 payload,
68 }
69 }
70}
71
72#[async_trait]
112pub trait SignalStore<T: Send + Sync + 'static>: Send + Sync {
113 async fn store(&self, signal: StoredSignal<T>) -> Result<(), SignalError>;
115
116 async fn retrieve(&self, id: u64) -> Result<Option<StoredSignal<T>>, SignalError>;
118
119 async fn list(&self, limit: usize, offset: usize) -> Result<Vec<StoredSignal<T>>, SignalError>;
121
122 async fn count(&self) -> Result<u64, SignalError>;
124
125 async fn clear(&self) -> Result<(), SignalError>;
127}
128
129pub struct MemoryStore<T> {
146 signals: Arc<RwLock<VecDeque<StoredSignal<T>>>>,
147 next_id: Arc<RwLock<u64>>,
148 max_size: usize,
149}
150
151impl<T> MemoryStore<T> {
152 pub fn new() -> Self {
154 Self {
155 signals: Arc::new(RwLock::new(VecDeque::new())),
156 next_id: Arc::new(RwLock::new(1)),
157 max_size: usize::MAX,
158 }
159 }
160
161 pub fn with_max_size(max_size: usize) -> Self {
165 Self {
166 signals: Arc::new(RwLock::new(VecDeque::new())),
167 next_id: Arc::new(RwLock::new(1)),
168 max_size,
169 }
170 }
171
172 pub fn max_size(&self) -> usize {
174 self.max_size
175 }
176}
177
178impl<T> Default for MemoryStore<T> {
179 fn default() -> Self {
180 Self::new()
181 }
182}
183
184impl<T> Clone for MemoryStore<T> {
185 fn clone(&self) -> Self {
186 Self {
187 signals: Arc::clone(&self.signals),
188 next_id: Arc::clone(&self.next_id),
189 max_size: self.max_size,
190 }
191 }
192}
193
194#[async_trait]
195impl<T: Send + Sync + Clone + 'static> SignalStore<T> for MemoryStore<T> {
196 async fn store(&self, signal: StoredSignal<T>) -> Result<(), SignalError> {
197 let mut signals = self.signals.write();
198
199 if signals.len() >= self.max_size {
201 signals.pop_front();
202 }
203
204 signals.push_back(signal);
205 Ok(())
206 }
207
208 async fn retrieve(&self, id: u64) -> Result<Option<StoredSignal<T>>, SignalError> {
209 let signals = self.signals.read();
210 Ok(signals.iter().find(|s| s.id == id).cloned())
211 }
212
213 async fn list(&self, limit: usize, offset: usize) -> Result<Vec<StoredSignal<T>>, SignalError> {
214 let signals = self.signals.read();
215 Ok(signals.iter().skip(offset).take(limit).cloned().collect())
216 }
217
218 async fn count(&self) -> Result<u64, SignalError> {
219 Ok(self.signals.read().len() as u64)
220 }
221
222 async fn clear(&self) -> Result<(), SignalError> {
223 self.signals.write().clear();
224 Ok(())
225 }
226}
227
228pub struct PersistentSignal<T: Send + Sync + 'static> {
252 signal: Signal<T>,
253 store: Arc<dyn SignalStore<T>>,
254 signal_name: String,
255 next_id: Arc<RwLock<u64>>,
256}
257
258impl<T: Send + Sync + Clone + 'static> PersistentSignal<T> {
259 pub fn new<S>(signal: Signal<T>, store: S) -> Self
266 where
267 S: SignalStore<T> + 'static,
268 {
269 let signal_name = format!("persistent_{}", std::any::type_name::<T>());
270
271 Self {
272 signal,
273 store: Arc::new(store),
274 signal_name,
275 next_id: Arc::new(RwLock::new(1)),
276 }
277 }
278
279 pub async fn send(&self, instance: T) -> Result<(), SignalError> {
300 let stored_instance = instance.clone();
302
303 let id = {
305 let mut next_id = self.next_id.write();
306 let id = *next_id;
307 *next_id += 1;
308 id
309 };
310
311 let stored_signal = StoredSignal::new(id, self.signal_name.clone(), stored_instance);
312
313 self.store.store(stored_signal).await?;
315
316 self.signal.send(instance).await
318 }
319
320 pub async fn retrieve(&self, id: u64) -> Result<Option<StoredSignal<T>>, SignalError> {
322 self.store.retrieve(id).await
323 }
324
325 pub async fn list(
327 &self,
328 limit: usize,
329 offset: usize,
330 ) -> Result<Vec<StoredSignal<T>>, SignalError> {
331 self.store.list(limit, offset).await
332 }
333
334 pub async fn count(&self) -> Result<u64, SignalError> {
336 self.store.count().await
337 }
338
339 pub async fn clear(&self) -> Result<(), SignalError> {
341 self.store.clear().await
342 }
343
344 pub fn signal(&self) -> &Signal<T> {
346 &self.signal
347 }
348
349 pub fn store(&self) -> Arc<dyn SignalStore<T>> {
351 Arc::clone(&self.store)
352 }
353}
354
355impl<T: Send + Sync + Clone + 'static> Clone for PersistentSignal<T> {
356 fn clone(&self) -> Self {
357 Self {
358 signal: self.signal.clone(),
359 store: Arc::clone(&self.store),
360 signal_name: self.signal_name.clone(),
361 next_id: Arc::clone(&self.next_id),
362 }
363 }
364}
365
366#[cfg(test)]
367mod tests {
368 use super::*;
369 use crate::signals::SignalName;
370 use std::sync::atomic::{AtomicUsize, Ordering};
371
372 #[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
373 struct TestEvent {
374 id: i32,
375 message: String,
376 }
377
378 #[tokio::test]
379 async fn test_memory_store_basic() {
380 let store = MemoryStore::new();
381 let event = StoredSignal::new(
382 1,
383 "test".to_string(),
384 TestEvent {
385 id: 1,
386 message: "Hello".to_string(),
387 },
388 );
389
390 store.store(event.clone()).await.unwrap();
392
393 let retrieved = store.retrieve(1).await.unwrap();
395 assert!(retrieved.is_some());
396 assert_eq!(retrieved.unwrap().payload.id, 1);
397
398 assert_eq!(store.count().await.unwrap(), 1);
400
401 let list = store.list(10, 0).await.unwrap();
403 assert_eq!(list.len(), 1);
404 assert_eq!(list[0].payload.message, "Hello");
405 }
406
407 #[tokio::test]
408 async fn test_memory_store_max_size() {
409 let store = MemoryStore::with_max_size(3);
410
411 for i in 1..=5 {
413 let event = StoredSignal::new(
414 i,
415 "test".to_string(),
416 TestEvent {
417 id: i as i32,
418 message: format!("Event {}", i),
419 },
420 );
421 store.store(event).await.unwrap();
422 }
423
424 assert_eq!(store.count().await.unwrap(), 3);
426
427 let list = store.list(10, 0).await.unwrap();
429 assert_eq!(list[0].id, 3);
430 assert_eq!(list[1].id, 4);
431 assert_eq!(list[2].id, 5);
432 }
433
434 #[tokio::test]
435 async fn test_memory_store_clear() {
436 let store = MemoryStore::new();
437
438 for i in 1..=3 {
439 let event = StoredSignal::new(
440 i,
441 "test".to_string(),
442 TestEvent {
443 id: i as i32,
444 message: "test".to_string(),
445 },
446 );
447 store.store(event).await.unwrap();
448 }
449
450 assert_eq!(store.count().await.unwrap(), 3);
451
452 store.clear().await.unwrap();
453 assert_eq!(store.count().await.unwrap(), 0);
454 }
455
456 #[tokio::test]
457 async fn test_persistent_signal_send_and_store() {
458 let signal = Signal::<TestEvent>::new(SignalName::custom("test_persistent"));
459 let store = MemoryStore::new();
460 let persistent = PersistentSignal::new(signal.clone(), store.clone());
461
462 let counter = Arc::new(AtomicUsize::new(0));
463 let counter_clone = Arc::clone(&counter);
464
465 signal.connect(move |_event| {
466 let counter = Arc::clone(&counter_clone);
467 async move {
468 counter.fetch_add(1, Ordering::SeqCst);
469 Ok(())
470 }
471 });
472
473 let event = TestEvent {
474 id: 42,
475 message: "Test event".to_string(),
476 };
477
478 persistent.send(event.clone()).await.unwrap();
480
481 assert_eq!(counter.load(Ordering::SeqCst), 1);
485
486 assert_eq!(store.count().await.unwrap(), 1);
488 let stored = store.retrieve(1).await.unwrap();
489 assert!(stored.is_some());
490 assert_eq!(stored.unwrap().payload.id, 42);
491 }
492
493 #[tokio::test]
494 async fn test_persistent_signal_list_pagination() {
495 let signal = Signal::<TestEvent>::new(SignalName::custom("test_pagination"));
496 let store = MemoryStore::new();
497 let persistent = PersistentSignal::new(signal, store);
498
499 for i in 1..=10 {
501 let event = TestEvent {
502 id: i,
503 message: format!("Event {}", i),
504 };
505 persistent.send(event).await.unwrap();
506 }
507
508 let page1 = persistent.list(5, 0).await.unwrap();
510 assert_eq!(page1.len(), 5);
511 assert_eq!(page1[0].payload.id, 1);
512
513 let page2 = persistent.list(5, 5).await.unwrap();
514 assert_eq!(page2.len(), 5);
515 assert_eq!(page2[0].payload.id, 6);
516
517 assert_eq!(persistent.count().await.unwrap(), 10);
519 }
520
521 #[tokio::test]
522 async fn test_persistent_signal_clear() {
523 let signal = Signal::<TestEvent>::new(SignalName::custom("test_clear"));
524 let store = MemoryStore::new();
525 let persistent = PersistentSignal::new(signal, store);
526
527 for i in 1..=5 {
528 persistent
529 .send(TestEvent {
530 id: i,
531 message: "test".to_string(),
532 })
533 .await
534 .unwrap();
535 }
536
537 assert_eq!(persistent.count().await.unwrap(), 5);
538
539 persistent.clear().await.unwrap();
540 assert_eq!(persistent.count().await.unwrap(), 0);
541 }
542}