Skip to main content

weavatrix_memory/store/
in_memory.rs

1use crate::{
2    AppendReceipt, EventId, EventMetadata, EventStore, ExpectedVersion, MemoryError, NewEvent,
3    Result, StoredEvent, StreamId,
4};
5use std::collections::{HashMap, HashSet};
6
7mod validation;
8
9#[derive(Debug, Clone)]
10pub struct InMemoryStore<E> {
11    events: Vec<StoredEvent<E>>,
12    streams: HashMap<StreamId, Vec<usize>>,
13    event_ids: HashSet<EventId>,
14}
15
16impl<E> Default for InMemoryStore<E> {
17    fn default() -> Self {
18        Self {
19            events: Vec::new(),
20            streams: HashMap::new(),
21            event_ids: HashSet::new(),
22        }
23    }
24}
25
26impl<E: Clone> InMemoryStore<E> {
27    pub(crate) fn prepare_append(
28        &self,
29        stream: &StreamId,
30        expected: ExpectedVersion,
31        events: &[NewEvent<E>],
32    ) -> Result<Vec<StoredEvent<E>>> {
33        let actual = self.stream_version(stream);
34        validation::expected(stream, expected, actual)?;
35        validation::unique_ids(&self.event_ids, events)?;
36
37        let start_version = actual.map_or(Ok(0), |version| {
38            version.checked_add(1).ok_or(MemoryError::CapacityOverflow)
39        })?;
40        let start_position =
41            u64::try_from(self.events.len()).map_err(|_| MemoryError::CapacityOverflow)?;
42        let mut committed = Vec::with_capacity(events.len());
43
44        for (offset, event) in events.iter().enumerate() {
45            let offset = u64::try_from(offset).map_err(|_| MemoryError::CapacityOverflow)?;
46            let stream_version = start_version
47                .checked_add(offset)
48                .ok_or(MemoryError::CapacityOverflow)?;
49            let global_position = start_position
50                .checked_add(offset)
51                .ok_or(MemoryError::CapacityOverflow)?;
52            committed.push(StoredEvent {
53                metadata: EventMetadata {
54                    id: event.id.clone(),
55                    stream_id: stream.clone(),
56                    stream_version,
57                    global_position,
58                    event_type: event.event_type.clone(),
59                    occurred_at: event.occurred_at,
60                    recorded_at: event.recorded_at,
61                    agent_id: event.agent_id.clone(),
62                    session_id: event.session_id.clone(),
63                    correlation_id: event.correlation_id.clone(),
64                    causation_id: event.causation_id.clone(),
65                },
66                payload: event.payload.clone(),
67            });
68        }
69
70        Ok(committed)
71    }
72
73    pub(crate) fn prepare_append_owned(
74        &self,
75        stream: &StreamId,
76        expected: ExpectedVersion,
77        events: Vec<NewEvent<E>>,
78    ) -> Result<Vec<StoredEvent<E>>> {
79        let actual = self.stream_version(stream);
80        validation::expected(stream, expected, actual)?;
81        validation::unique_ids(&self.event_ids, &events)?;
82
83        let start_version = actual.map_or(Ok(0), |version| {
84            version.checked_add(1).ok_or(MemoryError::CapacityOverflow)
85        })?;
86        let start_position =
87            u64::try_from(self.events.len()).map_err(|_| MemoryError::CapacityOverflow)?;
88        let mut committed = Vec::with_capacity(events.len());
89
90        for (offset, event) in events.into_iter().enumerate() {
91            let offset = u64::try_from(offset).map_err(|_| MemoryError::CapacityOverflow)?;
92            committed.push(StoredEvent {
93                metadata: EventMetadata {
94                    id: event.id,
95                    stream_id: stream.clone(),
96                    stream_version: start_version
97                        .checked_add(offset)
98                        .ok_or(MemoryError::CapacityOverflow)?,
99                    global_position: start_position
100                        .checked_add(offset)
101                        .ok_or(MemoryError::CapacityOverflow)?,
102                    event_type: event.event_type,
103                    occurred_at: event.occurred_at,
104                    recorded_at: event.recorded_at,
105                    agent_id: event.agent_id,
106                    session_id: event.session_id,
107                    correlation_id: event.correlation_id,
108                    causation_id: event.causation_id,
109                },
110                payload: event.payload,
111            });
112        }
113
114        Ok(committed)
115    }
116
117    pub(crate) fn commit_prepared(&mut self, committed: &[StoredEvent<E>]) {
118        let Some(first) = committed.first() else {
119            return;
120        };
121        debug_assert!(
122            committed
123                .iter()
124                .all(|event| event.metadata.stream_id == first.metadata.stream_id)
125        );
126        let Self {
127            events,
128            streams,
129            event_ids,
130        } = self;
131        events.reserve(committed.len());
132        event_ids.reserve(committed.len());
133        let positions = streams.entry(first.metadata.stream_id.clone()).or_default();
134        positions.reserve(committed.len());
135        for event in committed {
136            event_ids.insert(event.metadata.id.clone());
137            positions.push(events.len());
138            events.push(event.clone());
139        }
140    }
141
142    pub(crate) fn commit_prepared_owned(&mut self, committed: Vec<StoredEvent<E>>) {
143        let Some(first) = committed.first() else {
144            return;
145        };
146        debug_assert!(
147            committed
148                .iter()
149                .all(|event| event.metadata.stream_id == first.metadata.stream_id)
150        );
151        let Self {
152            events,
153            streams,
154            event_ids,
155        } = self;
156        events.reserve(committed.len());
157        event_ids.reserve(committed.len());
158        let positions = streams.entry(first.metadata.stream_id.clone()).or_default();
159        positions.reserve(committed.len());
160        for event in committed {
161            event_ids.insert(event.metadata.id.clone());
162            positions.push(events.len());
163            events.push(event);
164        }
165    }
166
167    pub(crate) fn restore(events: Vec<StoredEvent<E>>) -> Result<Self> {
168        let mut store = Self::default();
169        store.events.reserve(events.len());
170        store.event_ids.reserve(events.len());
171        for event in events {
172            let expected_position =
173                u64::try_from(store.events.len()).map_err(|_| MemoryError::CapacityOverflow)?;
174            if event.metadata.global_position != expected_position {
175                return Err(MemoryError::InvalidReplay {
176                    reason: format!(
177                        "global position {}, expected {expected_position}",
178                        event.metadata.global_position
179                    ),
180                });
181            }
182            let expected_version = store
183                .stream_version(&event.metadata.stream_id)
184                .map_or(Ok(0), |version| {
185                    version.checked_add(1).ok_or(MemoryError::CapacityOverflow)
186                })?;
187            if event.metadata.stream_version != expected_version {
188                return Err(MemoryError::InvalidReplay {
189                    reason: format!(
190                        "stream {} version {}, expected {expected_version}",
191                        event.metadata.stream_id, event.metadata.stream_version
192                    ),
193                });
194            }
195            if !store.event_ids.insert(event.metadata.id.clone()) {
196                return Err(MemoryError::DuplicateEvent {
197                    id: event.metadata.id.to_string(),
198                });
199            }
200            let positions = store
201                .streams
202                .entry(event.metadata.stream_id.clone())
203                .or_default();
204            positions.push(store.events.len());
205            store.events.push(event);
206        }
207        Ok(store)
208    }
209}
210
211impl<E: Clone> EventStore<E> for InMemoryStore<E> {
212    fn append(
213        &mut self,
214        stream: &StreamId,
215        expected: ExpectedVersion,
216        events: &[NewEvent<E>],
217    ) -> Result<Vec<StoredEvent<E>>> {
218        let committed = self.prepare_append(stream, expected, events)?;
219        self.commit_prepared(&committed);
220        Ok(committed)
221    }
222
223    fn append_owned(
224        &mut self,
225        stream: &StreamId,
226        expected: ExpectedVersion,
227        events: Vec<NewEvent<E>>,
228    ) -> Result<Vec<StoredEvent<E>>> {
229        let committed = self.prepare_append_owned(stream, expected, events)?;
230        self.commit_prepared(&committed);
231        Ok(committed)
232    }
233
234    fn append_owned_receipt(
235        &mut self,
236        stream: &StreamId,
237        expected: ExpectedVersion,
238        events: Vec<NewEvent<E>>,
239    ) -> Result<AppendReceipt> {
240        let committed = self.prepare_append_owned(stream, expected, events)?;
241        let receipt = AppendReceipt::from_events(&committed);
242        self.commit_prepared_owned(committed);
243        Ok(receipt)
244    }
245
246    fn load_stream(&self, stream: &StreamId, after: Option<u64>) -> Vec<StoredEvent<E>> {
247        self.streams
248            .get(stream)
249            .into_iter()
250            .flatten()
251            .map(|index| &self.events[*index])
252            .filter(|event| after.is_none_or(|cursor| event.metadata.stream_version > cursor))
253            .cloned()
254            .collect()
255    }
256
257    fn load_all(&self, after: Option<u64>, limit: usize) -> Vec<StoredEvent<E>> {
258        self.events
259            .iter()
260            .filter(|event| after.is_none_or(|cursor| event.metadata.global_position > cursor))
261            .take(limit)
262            .cloned()
263            .collect()
264    }
265
266    fn stream_version(&self, stream: &StreamId) -> Option<u64> {
267        self.streams
268            .get(stream)
269            .and_then(|positions| positions.last())
270            .map(|index| self.events[*index].metadata.stream_version)
271    }
272
273    fn len(&self) -> usize {
274        self.events.len()
275    }
276}