1use crate::audit::AuditError;
21use crate::event::Event;
22use crate::integrity::IntegrityError;
23use crate::snapshot::Snapshot;
24use async_trait::async_trait;
25use thiserror::Error;
26
27#[derive(Debug, Clone, Copy, PartialEq, Eq)]
29pub enum VersionCheck {
30 New,
32 Expected(i64),
34 Auto,
36}
37
38impl VersionCheck {
39 pub fn version(&self) -> Option<i64> {
40 match self {
41 VersionCheck::New => Some(0),
42 VersionCheck::Expected(v) => Some(*v),
43 VersionCheck::Auto => None,
44 }
45 }
46}
47
48#[derive(Debug, Error)]
50pub enum EventStoreError {
51 #[error("Concurrency conflict: expected version {expected}, but aggregate is at version {actual} (aggregate_id: {aggregate_id})")]
53 ConcurrencyConflict {
54 aggregate_id: String,
55 expected: i64,
56 actual: i64,
57 },
58
59 #[error("Aggregate not found: {aggregate_id}")]
60 AggregateNotFound { aggregate_id: String },
61
62 #[error(
63 "Invalid event sequence: expected {expected}, got {actual} (aggregate_id: {aggregate_id})"
64 )]
65 InvalidSequence {
66 aggregate_id: String,
67 expected: i64,
68 actual: i64,
69 },
70
71 #[error(
74 "Audit metadata validation failed for event {event_index} (aggregate_id: {aggregate_id}): {source}"
75 )]
76 InvalidAudit {
77 aggregate_id: String,
78 event_index: usize,
79 #[source]
80 source: AuditError,
81 },
82
83 #[error("Integrity validation failed: {source}")]
86 Integrity {
87 #[from]
88 source: IntegrityError,
89 },
90
91 #[error("Database error: {message}")]
92 DatabaseError { message: String },
93
94 #[error("Serialization error: {message}")]
95 SerializationError { message: String },
96
97 #[error("I/O error: {0}")]
98 IoError(#[from] std::io::Error),
99
100 #[error("snapshots not supported by this store")]
103 Unsupported,
104
105 #[error("Event store error: {message}")]
106 Other { message: String },
107}
108
109impl EventStoreError {
110 pub fn database(message: impl Into<String>) -> Self {
111 EventStoreError::DatabaseError {
112 message: message.into(),
113 }
114 }
115
116 pub fn serialization(message: impl Into<String>) -> Self {
117 EventStoreError::SerializationError {
118 message: message.into(),
119 }
120 }
121
122 pub fn other(message: impl Into<String>) -> Self {
123 EventStoreError::Other {
124 message: message.into(),
125 }
126 }
127
128 pub fn invalid_audit(
129 aggregate_id: impl Into<String>,
130 event_index: usize,
131 source: AuditError,
132 ) -> Self {
133 EventStoreError::InvalidAudit {
134 aggregate_id: aggregate_id.into(),
135 event_index,
136 source,
137 }
138 }
139}
140
141pub type EventStoreResult<T> = Result<T, EventStoreError>;
143
144pub fn validate_audit_batch(aggregate_id: &str, events: &[Event]) -> EventStoreResult<()> {
147 for (idx, ev) in events.iter().enumerate() {
148 ev.audit
149 .validate()
150 .map_err(|e| EventStoreError::invalid_audit(aggregate_id, idx, e))?;
151 }
152 Ok(())
153}
154
155#[async_trait]
160pub trait EventStore: Send + Sync {
161 async fn append(
166 &self,
167 aggregate_id: &str,
168 version_check: VersionCheck,
169 events: Vec<Event>,
170 ) -> EventStoreResult<()>;
171
172 async fn append_to(
178 &self,
179 aggregate_type: &str,
180 aggregate_id: &str,
181 version_check: VersionCheck,
182 events: Vec<Event>,
183 ) -> EventStoreResult<()> {
184 let _ = aggregate_type;
185 self.append(aggregate_id, version_check, events).await
186 }
187
188 async fn load(&self, aggregate_id: &str) -> EventStoreResult<Vec<Event>>;
189
190 async fn load_stream(
191 &self,
192 aggregate_type: &str,
193 aggregate_id: &str,
194 ) -> EventStoreResult<Vec<Event>> {
195 let _ = aggregate_type;
196 self.load(aggregate_id).await
197 }
198
199 async fn load_from(
200 &self,
201 aggregate_id: &str,
202 from_sequence: i64,
203 ) -> EventStoreResult<Vec<Event>>;
204
205 async fn load_stream_from(
206 &self,
207 aggregate_type: &str,
208 aggregate_id: &str,
209 from_sequence: i64,
210 ) -> EventStoreResult<Vec<Event>> {
211 let _ = aggregate_type;
212 self.load_from(aggregate_id, from_sequence).await
213 }
214
215 async fn stream_all(&self, from_position: i64) -> EventStoreResult<Vec<Event>>;
216
217 async fn get_version(&self, aggregate_id: &str) -> EventStoreResult<i64>;
218
219 async fn get_stream_version(
220 &self,
221 aggregate_type: &str,
222 aggregate_id: &str,
223 ) -> EventStoreResult<i64> {
224 let _ = aggregate_type;
225 self.get_version(aggregate_id).await
226 }
227
228 async fn save_snapshot(&self, snapshot: &Snapshot) -> EventStoreResult<()> {
234 let _ = snapshot;
235 Err(EventStoreError::Unsupported)
236 }
237
238 async fn load_snapshot(&self, aggregate_id: &str) -> EventStoreResult<Option<Snapshot>> {
243 let _ = aggregate_id;
244 Ok(None)
245 }
246
247 async fn load_snapshot_for(
248 &self,
249 aggregate_type: &str,
250 aggregate_id: &str,
251 ) -> EventStoreResult<Option<Snapshot>> {
252 let _ = aggregate_type;
253 self.load_snapshot(aggregate_id).await
254 }
255}
256
257#[cfg(any(test, feature = "test-utils"))]
262mod in_memory {
263 use super::*;
264 use std::collections::HashMap;
265 use std::sync::Arc;
266 use tokio::sync::Mutex as TokioMutex;
267
268 #[derive(Clone, Default)]
274 pub struct InMemoryEventStore {
275 events: Arc<TokioMutex<Vec<Event>>>,
276 snapshots: Arc<TokioMutex<HashMap<(String, String), Snapshot>>>,
277 }
278
279 impl InMemoryEventStore {
280 pub fn new() -> Self {
281 Self::default()
282 }
283 }
284
285 #[async_trait]
286 impl EventStore for InMemoryEventStore {
287 async fn append(
288 &self,
289 aggregate_id: &str,
290 version_check: VersionCheck,
291 events: Vec<Event>,
292 ) -> EventStoreResult<()> {
293 validate_audit_batch(aggregate_id, &events)?;
294
295 let mut store = self.events.lock().await;
296
297 let current_version = store
298 .iter()
299 .filter(|e| e.aggregate_id == aggregate_id)
300 .map(|e| e.sequence)
301 .max()
302 .unwrap_or(0);
303
304 if let Some(expected) = version_check.version() {
305 if current_version != expected {
306 return Err(EventStoreError::ConcurrencyConflict {
307 aggregate_id: aggregate_id.to_string(),
308 expected,
309 actual: current_version,
310 });
311 }
312 }
313
314 store.extend(events);
315 Ok(())
316 }
317
318 async fn append_to(
319 &self,
320 aggregate_type: &str,
321 aggregate_id: &str,
322 version_check: VersionCheck,
323 events: Vec<Event>,
324 ) -> EventStoreResult<()> {
325 validate_audit_batch(aggregate_id, &events)?;
326 let mut store = self.events.lock().await;
327 let current_version = store
328 .iter()
329 .filter(|event| {
330 event.aggregate_type == aggregate_type && event.aggregate_id == aggregate_id
331 })
332 .map(|event| event.sequence)
333 .max()
334 .unwrap_or(0);
335
336 if let Some(expected) = version_check.version() {
337 if current_version != expected {
338 return Err(EventStoreError::ConcurrencyConflict {
339 aggregate_id: aggregate_id.to_string(),
340 expected,
341 actual: current_version,
342 });
343 }
344 }
345
346 store.extend(events);
347 Ok(())
348 }
349
350 async fn load(&self, aggregate_id: &str) -> EventStoreResult<Vec<Event>> {
351 let store = self.events.lock().await;
352 Ok(store
353 .iter()
354 .filter(|e| e.aggregate_id == aggregate_id)
355 .cloned()
356 .collect())
357 }
358
359 async fn load_stream(
360 &self,
361 aggregate_type: &str,
362 aggregate_id: &str,
363 ) -> EventStoreResult<Vec<Event>> {
364 let store = self.events.lock().await;
365 Ok(store
366 .iter()
367 .filter(|event| {
368 event.aggregate_type == aggregate_type && event.aggregate_id == aggregate_id
369 })
370 .cloned()
371 .collect())
372 }
373
374 async fn load_from(
375 &self,
376 aggregate_id: &str,
377 from_sequence: i64,
378 ) -> EventStoreResult<Vec<Event>> {
379 let store = self.events.lock().await;
380 Ok(store
381 .iter()
382 .filter(|e| e.aggregate_id == aggregate_id && e.sequence >= from_sequence)
383 .cloned()
384 .collect())
385 }
386
387 async fn load_stream_from(
388 &self,
389 aggregate_type: &str,
390 aggregate_id: &str,
391 from_sequence: i64,
392 ) -> EventStoreResult<Vec<Event>> {
393 let store = self.events.lock().await;
394 Ok(store
395 .iter()
396 .filter(|event| {
397 event.aggregate_type == aggregate_type
398 && event.aggregate_id == aggregate_id
399 && event.sequence >= from_sequence
400 })
401 .cloned()
402 .collect())
403 }
404
405 async fn stream_all(&self, from_position: i64) -> EventStoreResult<Vec<Event>> {
406 let store = self.events.lock().await;
407 Ok(store.iter().skip(from_position as usize).cloned().collect())
408 }
409
410 async fn get_version(&self, aggregate_id: &str) -> EventStoreResult<i64> {
411 let store = self.events.lock().await;
412 Ok(store
413 .iter()
414 .filter(|e| e.aggregate_id == aggregate_id)
415 .map(|e| e.sequence)
416 .max()
417 .unwrap_or(0))
418 }
419
420 async fn get_stream_version(
421 &self,
422 aggregate_type: &str,
423 aggregate_id: &str,
424 ) -> EventStoreResult<i64> {
425 let store = self.events.lock().await;
426 Ok(store
427 .iter()
428 .filter(|event| {
429 event.aggregate_type == aggregate_type && event.aggregate_id == aggregate_id
430 })
431 .map(|event| event.sequence)
432 .max()
433 .unwrap_or(0))
434 }
435
436 async fn save_snapshot(&self, snapshot: &Snapshot) -> EventStoreResult<()> {
437 let mut snapshots = self.snapshots.lock().await;
438 snapshots.insert(
439 (
440 snapshot.aggregate_type.clone(),
441 snapshot.aggregate_id.clone(),
442 ),
443 snapshot.clone(),
444 );
445 Ok(())
446 }
447
448 async fn load_snapshot(&self, aggregate_id: &str) -> EventStoreResult<Option<Snapshot>> {
449 let snapshots = self.snapshots.lock().await;
450 Ok(snapshots
451 .iter()
452 .find(|((_, id), _)| id == aggregate_id)
453 .map(|(_, snapshot)| snapshot.clone()))
454 }
455
456 async fn load_snapshot_for(
457 &self,
458 aggregate_type: &str,
459 aggregate_id: &str,
460 ) -> EventStoreResult<Option<Snapshot>> {
461 let snapshots = self.snapshots.lock().await;
462 Ok(snapshots
463 .get(&(aggregate_type.to_string(), aggregate_id.to_string()))
464 .cloned())
465 }
466 }
467}
468
469#[cfg(any(test, feature = "test-utils"))]
470pub use in_memory::InMemoryEventStore;
471
472#[cfg(test)]
473mod tests {
474 use super::*;
475 use crate::audit::AuditMetadata;
476 use crate::event::Event;
477 use serde_json::json;
478
479 #[test]
480 fn test_version_check_new() {
481 assert_eq!(VersionCheck::New.version(), Some(0));
482 }
483
484 #[test]
485 fn test_version_check_expected() {
486 assert_eq!(VersionCheck::Expected(5).version(), Some(5));
487 }
488
489 #[test]
490 fn test_version_check_auto() {
491 assert_eq!(VersionCheck::Auto.version(), None);
492 }
493
494 #[test]
495 fn test_error_messages() {
496 let error = EventStoreError::ConcurrencyConflict {
497 aggregate_id: "user-123".to_string(),
498 expected: 5,
499 actual: 6,
500 };
501 let msg = error.to_string();
502 assert!(msg.contains("expected version 5"));
503 assert!(msg.contains("aggregate is at version 6"));
504 assert!(msg.contains("user-123"));
505
506 assert!(EventStoreError::database("X").to_string().contains("X"));
507 assert!(EventStoreError::serialization("Y")
508 .to_string()
509 .contains("Y"));
510 }
511
512 #[test]
513 fn test_validate_audit_batch_rejects_pending() {
514 let mut e = Event::new("User", "u1", 1, "X", json!({}));
515 e.audit = AuditMetadata::pending();
516 let err = validate_audit_batch("u1", &[e]).unwrap_err();
517 assert!(matches!(
518 err,
519 EventStoreError::InvalidAudit { event_index: 0, .. }
520 ));
521 }
522
523 #[test]
524 fn test_validate_audit_batch_passes_stamped() {
525 let e =
526 Event::new("User", "u1", 1, "X", json!({})).with_audit(AuditMetadata::test_default());
527 validate_audit_batch("u1", &[e]).expect("stamped audit must pass");
528 }
529
530 #[tokio::test]
531 async fn test_in_memory_store_rejects_pending_audit() {
532 let store = InMemoryEventStore::new();
533 let e = Event::new("User", "u1", 1, "X", json!({})); let err = store
535 .append("u1", VersionCheck::New, vec![e])
536 .await
537 .unwrap_err();
538 assert!(matches!(err, EventStoreError::InvalidAudit { .. }));
539 }
540
541 #[tokio::test]
542 async fn test_in_memory_store_persists_stamped_event() {
543 let store = InMemoryEventStore::new();
544 let e =
545 Event::new("User", "u1", 1, "X", json!({})).with_audit(AuditMetadata::test_default());
546 store
547 .append("u1", VersionCheck::New, vec![e])
548 .await
549 .unwrap();
550 let loaded = store.load("u1").await.unwrap();
551 assert_eq!(loaded.len(), 1);
552 }
553
554 #[tokio::test]
555 async fn test_in_memory_snapshot_save_then_load() {
556 let store = InMemoryEventStore::new();
557 let snap = Snapshot::new("u1", "User", 3, json!({ "name": "Alice" }));
558 store.save_snapshot(&snap).await.unwrap();
559 let loaded = store.load_snapshot("u1").await.unwrap();
560 assert_eq!(loaded, Some(snap));
561 }
562
563 #[tokio::test]
564 async fn test_in_memory_load_snapshot_unknown_returns_none() {
565 let store = InMemoryEventStore::new();
566 assert_eq!(store.load_snapshot("missing").await.unwrap(), None);
567 }
568
569 #[tokio::test]
570 async fn test_in_memory_snapshot_overwrites_on_resave() {
571 let store = InMemoryEventStore::new();
572 store
573 .save_snapshot(&Snapshot::new("u1", "User", 3, json!({ "v": 3 })))
574 .await
575 .unwrap();
576 let newer = Snapshot::new("u1", "User", 9, json!({ "v": 9 }));
577 store.save_snapshot(&newer).await.unwrap();
578 let loaded = store.load_snapshot("u1").await.unwrap().unwrap();
579 assert_eq!(loaded.version, 9);
580 assert_eq!(loaded.state["v"], 9);
581 }
582}