Skip to main content

weavatrix_memory/store/in_memory/
mod.rs

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