Skip to main content

weavatrix_memory/store/
in_memory.rs

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