Skip to main content

distributed/hashmap_repo/
repository.rs

1#![expect(
2    clippy::manual_async_fn,
3    reason = "async trait impls return impl Future + Send to preserve public Send bounds"
4)]
5
6use std::collections::{HashMap, HashSet};
7use std::future::Future;
8use std::sync::{Arc, RwLock};
9
10use crate::entity::{
11    Entity, EventRecord, EventRecordError, BITCODE_PAYLOAD_CODEC, BITCODE_PAYLOAD_CODEC_VERSION,
12};
13use crate::outbox::OutboxMessage;
14use crate::read_model::in_memory::apply_read_model_write_plan;
15use crate::read_model::{
16    InMemoryReadModelStore, ReadModelAdapterCapabilities, ReadModelCommitOutcome, ReadModelError,
17    ReadModelLoadGraph, ReadModelLoadRequest, ReadModelQueryCapabilities, ReadModelWritePlan,
18};
19use crate::repository::{
20    CommitBatch, GetStream, InboxStore, PreparedEventAppend, ReadModelWritePlanStore,
21    RelationalReadModelQueryStore, RepositoryError, SnapshotStore, SnapshotWrite, StreamIdentity,
22    StreamWrite, TransactionalCommit,
23};
24use crate::snapshot::{InMemorySnapshotStore, SnapshotRecord};
25
26/// In-memory repository implementation using HashMap.
27///
28/// This repository is cheap to clone because it uses `Arc<RwLock<...>>`
29/// internally - cloning creates another handle to the same storage.
30/// Also includes an embedded `InMemoryReadModelStore` for read model storage.
31#[derive(Clone)]
32pub struct HashMapRepository {
33    event_store: Arc<RwLock<HashMap<String, Vec<EventRecord>>>>,
34    outbox_store: Arc<RwLock<HashMap<String, OutboxMessage>>>,
35    model_store: InMemoryReadModelStore,
36    snapshot_store: InMemorySnapshotStore,
37    /// Consumer inbox: the set of recorded `(consumer, message_id)` receipts.
38    inbox_store: Arc<RwLock<HashSet<(String, String)>>>,
39}
40
41/// In-memory outbox table handle.
42#[derive(Clone)]
43pub struct HashMapOutboxStore {
44    pub(crate) storage: Arc<RwLock<HashMap<String, OutboxMessage>>>,
45}
46
47impl Default for HashMapRepository {
48    fn default() -> Self {
49        Self::new()
50    }
51}
52
53impl HashMapRepository {
54    /// Create a new empty repository.
55    pub fn new() -> Self {
56        HashMapRepository {
57            event_store: Arc::new(RwLock::new(HashMap::new())),
58            outbox_store: Arc::new(RwLock::new(HashMap::new())),
59            model_store: InMemoryReadModelStore::new(),
60            snapshot_store: InMemorySnapshotStore::new(),
61            inbox_store: Arc::new(RwLock::new(HashSet::new())),
62        }
63    }
64
65    #[cfg(test)]
66    pub(crate) fn outbox_storage(&self) -> &RwLock<HashMap<String, OutboxMessage>> {
67        self.outbox_store.as_ref()
68    }
69
70    /// Access the in-memory outbox table handle.
71    pub fn outbox_store(&self) -> HashMapOutboxStore {
72        HashMapOutboxStore {
73            storage: Arc::clone(&self.outbox_store),
74        }
75    }
76
77    /// Access the embedded read model store directly.
78    pub fn model_store(&self) -> &InMemoryReadModelStore {
79        &self.model_store
80    }
81
82    /// Access the embedded snapshot store directly.
83    pub fn snapshot_store(&self) -> &InMemorySnapshotStore {
84        &self.snapshot_store
85    }
86
87    /// Whether a consumer inbox receipt for `(consumer, message_id)` is recorded.
88    pub fn inbox_contains(&self, consumer: &str, message_id: &str) -> bool {
89        self.inbox_store
90            .read()
91            .map(|set| set.contains(&(consumer.to_string(), message_id.to_string())))
92            .unwrap_or(false)
93    }
94}
95
96impl GetStream for HashMapRepository {
97    fn get_stream<'a>(
98        &'a self,
99        identity: &'a StreamIdentity,
100    ) -> impl Future<Output = Result<Option<Entity>, RepositoryError>> + Send + 'a {
101        async move {
102            let storage = self
103                .event_store
104                .read()
105                .map_err(|_| RepositoryError::LockPoisoned("async stream read"))?;
106
107            if let Some(events) = storage.get(&identity.storage_key()) {
108                let mut entity = Entity::new();
109                entity.set_id(identity.aggregate_id());
110                entity.load_from_history(events.clone());
111                Ok(Some(entity))
112            } else {
113                Ok(None)
114            }
115        }
116    }
117
118    fn get_streams<'a>(
119        &'a self,
120        identities: &'a [StreamIdentity],
121    ) -> impl Future<Output = Result<Vec<Entity>, RepositoryError>> + Send + 'a {
122        async move {
123            let mut entities = Vec::with_capacity(identities.len());
124            for identity in identities {
125                if let Some(entity) = self.get_stream(identity).await? {
126                    entities.push(entity);
127                }
128            }
129            Ok(entities)
130        }
131    }
132}
133
134impl TransactionalCommit for HashMapRepository {
135    fn commit_batch<'a>(
136        &'a self,
137        batch: CommitBatch<'a>,
138    ) -> impl Future<Output = Result<(), RepositoryError>> + Send + 'a {
139        async move {
140            reject_duplicate_streams(&batch.streams)?;
141            validate_entity_id_matches_identity(&batch.streams)?;
142            let prepared = batch
143                .streams
144                .iter()
145                .map(PreparedEventAppend::from_stream_write)
146                .collect::<Vec<_>>();
147            validate_prepared_appends(&prepared)?;
148            for write in &batch.snapshots {
149                validate_snapshot_write(write)?;
150            }
151            reject_duplicate_outbox_messages(&batch.outbox_messages)?;
152
153            let mut storage = self
154                .event_store
155                .write()
156                .map_err(|_| RepositoryError::LockPoisoned("async stream write"))?;
157            let mut relational_rows = self
158                .model_store
159                .relational_rows
160                .write()
161                .map_err(|_| RepositoryError::LockPoisoned("async read model write"))?;
162            let mut snapshot_storage = self
163                .snapshot_store
164                .storage
165                .write()
166                .map_err(|_| RepositoryError::LockPoisoned("async snapshot write"))?;
167            let mut outbox_storage = self
168                .outbox_store
169                .write()
170                .map_err(|_| RepositoryError::LockPoisoned("async outbox write"))?;
171            let mut inbox_storage = self
172                .inbox_store
173                .write()
174                .map_err(|_| RepositoryError::LockPoisoned("async inbox write"))?;
175
176            let mut staged_events = storage.clone();
177            let mut staged_rows = relational_rows.clone();
178            let mut staged_snapshots = snapshot_storage.clone();
179            let mut staged_outbox = outbox_storage.clone();
180            let mut staged_inbox = inbox_storage.clone();
181
182            for append in &prepared {
183                let stored_len =
184                    stored_stream_version(staged_events.get(&append.identity.storage_key()));
185                if stored_len != append.expected_version {
186                    return Err(RepositoryError::ConcurrentWrite {
187                        id: append.identity.to_string(),
188                        expected: append.expected_version,
189                        actual: stored_len,
190                    });
191                }
192            }
193
194            for append in prepared {
195                let stored = staged_events
196                    .entry(append.identity.storage_key())
197                    .or_insert_with(Vec::new);
198                stored.extend(append.events);
199            }
200
201            for plan in batch.read_model_plans {
202                apply_read_model_write_plan(plan, &mut staged_rows)?;
203            }
204
205            for write in batch.snapshots {
206                match write {
207                    SnapshotWrite::Save { identity, record } => {
208                        staged_snapshots.insert(identity.storage_key(), record);
209                    }
210                }
211            }
212
213            for message in batch.outbox_messages {
214                let id = message.id().to_string();
215                if staged_outbox.contains_key(&id) {
216                    return Err(RepositoryError::DuplicateOutboxMessageInBatch { id });
217                }
218                staged_outbox.insert(id, message);
219            }
220
221            // Inbox receipts gate effectively-once: a receipt that already exists
222            // (committed or duplicated in this batch) rolls the whole batch back so
223            // effects are not double-applied.
224            for receipt in batch.inbox_receipts {
225                receipt.validate()?;
226                let key = (receipt.consumer.clone(), receipt.message_id.clone());
227                if !staged_inbox.insert(key) {
228                    return Err(RepositoryError::DuplicateInboxReceipt {
229                        consumer: receipt.consumer,
230                        message_id: receipt.message_id,
231                    });
232                }
233            }
234
235            *storage = staged_events;
236            *relational_rows = staged_rows;
237            *snapshot_storage = staged_snapshots;
238            *outbox_storage = staged_outbox;
239            *inbox_storage = staged_inbox;
240
241            for stream in batch.streams {
242                stream.entity.mark_committed();
243            }
244
245            Ok(())
246        }
247    }
248}
249
250impl InboxStore for HashMapRepository {
251    fn inbox_contains<'a>(
252        &'a self,
253        consumer: &'a str,
254        message_id: &'a str,
255    ) -> impl Future<Output = Result<bool, RepositoryError>> + Send + 'a {
256        async move { Ok(self.inbox_contains(consumer, message_id)) }
257    }
258}
259
260fn reject_duplicate_streams(streams: &[StreamWrite<'_>]) -> Result<(), RepositoryError> {
261    let mut seen = HashSet::with_capacity(streams.len());
262    for stream in streams {
263        let key = stream.identity.storage_key();
264        if !seen.insert(key) {
265            return Err(RepositoryError::DuplicateStreamInBatch {
266                id: stream.identity.to_string(),
267            });
268        }
269    }
270    Ok(())
271}
272
273fn reject_duplicate_outbox_messages(messages: &[OutboxMessage]) -> Result<(), RepositoryError> {
274    let mut seen = HashSet::with_capacity(messages.len());
275    for message in messages {
276        crate::outbox::validate_outbox_message_table_write(message)
277            .map_err(|err| RepositoryError::Model(err.to_string()))?;
278        let id = message.id();
279        if id.trim().is_empty() {
280            return Err(RepositoryError::Model(
281                "outbox message id must not be empty".into(),
282            ));
283        }
284        if message.event_type.trim().is_empty() {
285            return Err(RepositoryError::Model(format!(
286                "outbox message `{id}` event type must not be empty"
287            )));
288        }
289        if !seen.insert(id.to_string()) {
290            return Err(RepositoryError::DuplicateOutboxMessageInBatch { id: id.into() });
291        }
292    }
293    Ok(())
294}
295
296fn validate_entity_id_matches_identity(streams: &[StreamWrite<'_>]) -> Result<(), RepositoryError> {
297    for stream in streams {
298        if stream.entity.id() != stream.identity.aggregate_id() {
299            return Err(RepositoryError::Model(format!(
300                "stream identity `{}` does not match entity id `{}`",
301                stream.identity,
302                stream.entity.id()
303            )));
304        }
305    }
306    Ok(())
307}
308
309fn validate_prepared_appends(appends: &[PreparedEventAppend]) -> Result<(), RepositoryError> {
310    for append in appends {
311        for (offset, event) in append.events.iter().enumerate() {
312            validate_supported_event_codec(event)?;
313            let expected_sequence = append.expected_version + offset as u64 + 1;
314            if event.sequence != expected_sequence {
315                return Err(RepositoryError::Model(format!(
316                    "event `{}` for stream `{}` has sequence {}, expected {}",
317                    event.event_name, append.identity, event.sequence, expected_sequence
318                )));
319            }
320        }
321    }
322    Ok(())
323}
324
325fn validate_supported_event_codec(event: &EventRecord) -> Result<(), RepositoryError> {
326    if event.payload_codec != BITCODE_PAYLOAD_CODEC
327        || event.payload_codec_version != BITCODE_PAYLOAD_CODEC_VERSION
328    {
329        return Err(EventRecordError::unsupported_codec(
330            &event.payload_codec,
331            event.payload_codec_version,
332        )
333        .into());
334    }
335    Ok(())
336}
337
338fn validate_snapshot_write(write: &SnapshotWrite) -> Result<(), RepositoryError> {
339    match write {
340        SnapshotWrite::Save { identity, record } => validate_snapshot_identity(identity, record),
341    }
342}
343
344fn validate_snapshot_identity(
345    identity: &StreamIdentity,
346    record: &SnapshotRecord,
347) -> Result<(), RepositoryError> {
348    record.validate_for_identity(identity)
349}
350
351fn stored_stream_version(events: Option<&Vec<EventRecord>>) -> u64 {
352    // A missing stream has committed version 0; the first appended event will
353    // occupy sequence 1.
354    events.map_or(0, |events| events.len() as u64)
355}
356
357impl ReadModelWritePlanStore for HashMapRepository {
358    fn read_model_capabilities(&self) -> ReadModelAdapterCapabilities {
359        self.model_store.read_model_capabilities()
360    }
361
362    fn commit_write_plan(
363        &self,
364        plan: ReadModelWritePlan,
365    ) -> impl Future<Output = Result<ReadModelCommitOutcome, ReadModelError>> + Send + '_ {
366        self.model_store.commit_write_plan(plan)
367    }
368}
369
370impl RelationalReadModelQueryStore for HashMapRepository {
371    fn read_model_query_capabilities(&self) -> ReadModelQueryCapabilities {
372        self.model_store.read_model_query_capabilities()
373    }
374
375    fn load_graph(
376        &self,
377        request: ReadModelLoadRequest,
378    ) -> impl Future<Output = Result<ReadModelLoadGraph, ReadModelError>> + Send + '_ {
379        self.model_store.load_graph(request)
380    }
381}
382
383impl SnapshotStore for HashMapRepository {
384    fn get_snapshot<'a>(
385        &'a self,
386        identity: &'a StreamIdentity,
387    ) -> impl Future<Output = Result<Option<SnapshotRecord>, RepositoryError>> + Send + 'a {
388        async move {
389            let storage = self
390                .snapshot_store
391                .storage
392                .read()
393                .map_err(|_| RepositoryError::LockPoisoned("async snapshot read"))?;
394            Ok(storage.get(&identity.storage_key()).cloned())
395        }
396    }
397
398    fn save_snapshot<'a>(
399        &'a self,
400        identity: &'a StreamIdentity,
401        record: SnapshotRecord,
402    ) -> impl Future<Output = Result<(), RepositoryError>> + Send + 'a {
403        async move {
404            validate_snapshot_identity(identity, &record)?;
405            let mut storage = self
406                .snapshot_store
407                .storage
408                .write()
409                .map_err(|_| RepositoryError::LockPoisoned("async snapshot write"))?;
410            storage.insert(identity.storage_key(), record);
411            Ok(())
412        }
413    }
414
415    fn delete_snapshot<'a>(
416        &'a self,
417        identity: &'a StreamIdentity,
418    ) -> impl Future<Output = Result<bool, RepositoryError>> + Send + 'a {
419        async move {
420            let mut storage = self
421                .snapshot_store
422                .storage
423                .write()
424                .map_err(|_| RepositoryError::LockPoisoned("async snapshot write"))?;
425            Ok(storage.remove(&identity.storage_key()).is_some())
426        }
427    }
428}
429
430#[cfg(test)]
431mod tests {
432    use super::*;
433
434    fn identity(id: &str) -> StreamIdentity {
435        StreamIdentity::new("test.aggregate", id).unwrap()
436    }
437
438    async fn commit_one(
439        repo: &HashMapRepository,
440        entity: &mut Entity,
441    ) -> Result<(), RepositoryError> {
442        let id = entity.id().to_string();
443        repo.commit_batch(CommitBatch::new(vec![StreamWrite::new(
444            identity(&id),
445            entity,
446        )]))
447        .await
448    }
449
450    #[test]
451    fn new() {
452        let repo = HashMapRepository::new();
453        assert!(repo.event_store.read().unwrap().is_empty());
454    }
455
456    #[tokio::test]
457    async fn single_entity_commit() {
458        let repo = HashMapRepository::new();
459        let id = "test_id";
460        let mut entity = Entity::with_id(id);
461
462        entity.digest("test_event", &("arg1", "arg2")).unwrap();
463
464        commit_one(&repo, &mut entity).await.unwrap();
465
466        let fetched_entity = repo.get_stream(&identity(id)).await.unwrap().unwrap();
467        assert_eq!(fetched_entity.id(), id);
468        assert_eq!(fetched_entity.events(), entity.events());
469    }
470
471    #[tokio::test]
472    async fn multiple_entity_commit() {
473        let repo = HashMapRepository::new();
474
475        let mut entity1 = Entity::with_id("id_1");
476        entity1.digest("event1", &"arg1").unwrap();
477
478        let mut entity2 = Entity::with_id("id_2");
479        entity2.digest("event2", &"arg2").unwrap();
480
481        repo.commit_batch(CommitBatch::new(vec![
482            StreamWrite::new(identity("id_1"), &mut entity1),
483            StreamWrite::new(identity("id_2"), &mut entity2),
484        ]))
485        .await
486        .unwrap();
487
488        let all_entities: Vec<Entity> = repo
489            .get_streams(&[identity("id_1"), identity("id_2")])
490            .await
491            .unwrap();
492        assert_eq!(all_entities.len(), 2);
493    }
494
495    #[tokio::test]
496    async fn duplicate_stream_ids_rejected_before_write() {
497        let repo = HashMapRepository::new();
498
499        let mut entity1 = Entity::with_id("same-id");
500        entity1.digest("event1", &"arg1").unwrap();
501
502        let mut entity2 = Entity::with_id("same-id");
503        entity2.digest("event2", &"arg2").unwrap();
504
505        let err = repo
506            .commit_batch(CommitBatch::new(vec![
507                StreamWrite::new(identity("same-id"), &mut entity1),
508                StreamWrite::new(identity("same-id"), &mut entity2),
509            ]))
510            .await
511            .unwrap_err();
512        assert_eq!(
513            err,
514            RepositoryError::DuplicateStreamInBatch {
515                id: identity("same-id").to_string(),
516            }
517        );
518
519        assert!(repo
520            .get_stream(&identity("same-id"))
521            .await
522            .unwrap()
523            .is_none());
524        assert_eq!(entity1.committed_version(), 0);
525        assert_eq!(entity2.committed_version(), 0);
526        assert_eq!(entity1.new_events().len(), 1);
527        assert_eq!(entity2.new_events().len(), 1);
528    }
529
530    #[tokio::test]
531    async fn inbox_receipts_record_dedupe_and_roll_back_atomically() {
532        use crate::repository::InboxReceipt;
533        let repo = HashMapRepository::new();
534
535        let mut batch = CommitBatch::empty();
536        batch.inbox_receipts.push(InboxReceipt::new("proj", "m1"));
537        repo.commit_batch(batch).await.unwrap();
538        assert!(repo.inbox_contains("proj", "m1"));
539        assert!(!repo.inbox_contains("proj", "m2"));
540
541        // A batch with a duplicate (m1) and a fresh receipt (m2) rolls back whole.
542        let mut dup = CommitBatch::empty();
543        dup.inbox_receipts.push(InboxReceipt::new("proj", "m1"));
544        dup.inbox_receipts.push(InboxReceipt::new("proj", "m2"));
545        let err = repo.commit_batch(dup).await.unwrap_err();
546        assert!(
547            matches!(err, RepositoryError::DuplicateInboxReceipt { ref message_id, .. } if message_id == "m1"),
548            "got {err:?}"
549        );
550        assert!(
551            !repo.inbox_contains("proj", "m2"),
552            "the duplicate rolled the whole batch back"
553        );
554
555        // An empty receipt field is rejected (parity with the SQL CHECK).
556        let mut invalid = CommitBatch::empty();
557        invalid.inbox_receipts.push(InboxReceipt::new("", "m3"));
558        assert!(matches!(
559            repo.commit_batch(invalid).await.unwrap_err(),
560            RepositoryError::InvalidInboxReceipt { .. }
561        ));
562    }
563}