1use super::change_engine::{stage_change_set, StagedChange};
2use super::query::query_nodes;
3use super::snapshot::snapshot_from_map;
4use super::validation::{validate_required_text, MAX_IDENTIFIER_BYTES};
5use super::{
6 MemoryAccessEvent, MemoryChangeResult, MemoryChangeSet, MemoryNamespace,
7 MemoryNamespaceChangeToken, MemoryNamespaceSnapshot, MemoryNode, MemoryQuery,
8 MemoryQueryResult, MemoryRepository, MemoryRepositoryError, MemoryRepositorySnapshot,
9 MemorySnapshotRequest, MemoryUsageSummary,
10};
11use std::collections::BTreeMap;
12use tokio::sync::RwLock;
13
14#[derive(Debug, Clone)]
15struct AppliedChange {
16 change_set: MemoryChangeSet,
17 result: MemoryChangeResult,
18}
19
20#[derive(Debug, Default)]
21struct RepositoryState {
22 nodes: BTreeMap<MemoryNamespace, BTreeMap<String, MemoryNode>>,
23 namespace_change_sequences: BTreeMap<MemoryNamespace, u64>,
24 applied_changes: BTreeMap<MemoryNamespace, BTreeMap<String, AppliedChange>>,
25 admissions: BTreeMap<MemoryNamespace, BTreeMap<String, MemoryAccessEvent>>,
26 uses: BTreeMap<MemoryNamespace, BTreeMap<String, MemoryAccessEvent>>,
27 usage: BTreeMap<MemoryNamespace, BTreeMap<String, MemoryUsageSummary>>,
28}
29
30#[derive(Debug, Default)]
32pub struct InMemoryRepository {
33 state: RwLock<RepositoryState>,
34}
35
36impl InMemoryRepository {
37 pub fn new() -> Self {
38 Self::default()
39 }
40
41 pub async fn snapshot(&self) -> MemoryRepositorySnapshot {
43 let state = self.state.read().await;
44 MemoryRepositorySnapshot {
45 nodes: state
46 .nodes
47 .values()
48 .flat_map(BTreeMap::values)
49 .cloned()
50 .collect(),
51 admissions: state
52 .admissions
53 .values()
54 .flat_map(BTreeMap::values)
55 .cloned()
56 .collect(),
57 uses: state
58 .uses
59 .values()
60 .flat_map(BTreeMap::values)
61 .cloned()
62 .collect(),
63 }
64 }
65
66 pub(crate) async fn preview_apply(
67 &self,
68 change_set: &MemoryChangeSet,
69 ) -> Result<(MemoryChangeResult, bool), MemoryRepositoryError> {
70 let state = self.state.read().await;
71 let prepared = prepare_change(&state, change_set)?;
72 Ok((prepared.result, prepared.staged.is_none()))
73 }
74
75 pub(crate) async fn preview_admission(
76 &self,
77 event: &MemoryAccessEvent,
78 ) -> Result<bool, MemoryRepositoryError> {
79 self.preview_access(event, AccessKind::Admission).await
80 }
81
82 pub(crate) async fn preview_use(
83 &self,
84 event: &MemoryAccessEvent,
85 ) -> Result<bool, MemoryRepositoryError> {
86 self.preview_access(event, AccessKind::Use).await
87 }
88
89 async fn preview_access(
90 &self,
91 event: &MemoryAccessEvent,
92 kind: AccessKind,
93 ) -> Result<bool, MemoryRepositoryError> {
94 event.validate()?;
95 let state = self.state.read().await;
96 validate_access(&state, event, kind)
97 }
98
99 async fn record_access(
100 &self,
101 event: MemoryAccessEvent,
102 kind: AccessKind,
103 ) -> Result<(), MemoryRepositoryError> {
104 event.validate()?;
105 let mut state = self.state.write().await;
106 if validate_access(&state, &event, kind)? {
107 return Ok(());
108 }
109
110 let current_summary = state
111 .usage
112 .get(&event.namespace)
113 .and_then(|nodes| nodes.get(&event.node_id))
114 .copied()
115 .unwrap_or_default();
116 let mut next_summary = current_summary;
117 match kind {
118 AccessKind::Admission => {
119 next_summary.admissions =
120 next_summary.admissions.checked_add(1).ok_or_else(|| {
121 MemoryRepositoryError::invariant("admission counter overflow")
122 })?;
123 }
124 AccessKind::Use => {
125 next_summary.uses = next_summary
126 .uses
127 .checked_add(1)
128 .ok_or_else(|| MemoryRepositoryError::invariant("use counter overflow"))?;
129 }
130 }
131
132 match kind {
133 AccessKind::Admission => &mut state.admissions,
134 AccessKind::Use => &mut state.uses,
135 }
136 .entry(event.namespace.clone())
137 .or_default()
138 .insert(event.id.clone(), event.clone());
139 state
140 .usage
141 .entry(event.namespace)
142 .or_default()
143 .insert(event.node_id, next_summary);
144 Ok(())
145 }
146}
147
148#[derive(Debug, Clone, Copy)]
149enum AccessKind {
150 Admission,
151 Use,
152}
153
154#[async_trait::async_trait]
155impl MemoryRepository for InMemoryRepository {
156 async fn apply(
157 &self,
158 change_set: MemoryChangeSet,
159 ) -> Result<MemoryChangeResult, MemoryRepositoryError> {
160 let mut state = self.state.write().await;
161 let PreparedChange {
162 result,
163 staged,
164 next_change_sequence,
165 } = prepare_change(&state, &change_set)?;
166 let Some(mut staged) = staged else {
167 return Ok(result);
168 };
169 let next_change_sequence = next_change_sequence.ok_or_else(|| {
170 MemoryRepositoryError::invariant(
171 "novel memory change did not reserve a namespace change sequence",
172 )
173 })?;
174 let namespace = change_set.namespace.clone();
175 let namespace_nodes = state.nodes.entry(namespace.clone()).or_default();
176 for node_id in &staged.changed {
177 let node = staged.nodes.remove(node_id).ok_or_else(|| {
178 MemoryRepositoryError::invariant(format!(
179 "changed memory node {node_id} was not staged"
180 ))
181 })?;
182 namespace_nodes.insert(node_id.clone(), node);
183 }
184 state
185 .applied_changes
186 .entry(namespace.clone())
187 .or_default()
188 .insert(
189 change_set.idempotency_key.clone(),
190 AppliedChange {
191 change_set,
192 result: result.clone(),
193 },
194 );
195 state
196 .namespace_change_sequences
197 .insert(namespace, next_change_sequence);
198 Ok(result)
199 }
200
201 async fn get(
202 &self,
203 namespace: &MemoryNamespace,
204 node_id: &str,
205 ) -> Result<Option<MemoryNode>, MemoryRepositoryError> {
206 namespace.validate()?;
207 validate_required_text("nodeId", node_id, MAX_IDENTIFIER_BYTES)?;
208 let state = self.state.read().await;
209 Ok(state
210 .nodes
211 .get(namespace)
212 .and_then(|nodes| nodes.get(node_id))
213 .cloned())
214 }
215
216 async fn query(&self, query: MemoryQuery) -> Result<MemoryQueryResult, MemoryRepositoryError> {
217 query.validate()?;
218 let state = self.state.read().await;
219 Ok(query_nodes(state.nodes.get(&query.namespace), &query))
220 }
221
222 async fn snapshot_namespace(
223 &self,
224 request: MemorySnapshotRequest,
225 ) -> Result<MemoryNamespaceSnapshot, MemoryRepositoryError> {
226 request.validate()?;
227 let state = self.state.read().await;
228 snapshot_from_map(state.nodes.get(&request.namespace), request)
229 }
230
231 async fn namespace_change_token(
232 &self,
233 namespace: &MemoryNamespace,
234 ) -> Result<Option<MemoryNamespaceChangeToken>, MemoryRepositoryError> {
235 namespace.validate()?;
236 let state = self.state.read().await;
237 let sequence = state
238 .namespace_change_sequences
239 .get(namespace)
240 .copied()
241 .unwrap_or(0);
242 Ok(Some(MemoryNamespaceChangeToken::new(sequence)))
243 }
244
245 async fn record_admission(
246 &self,
247 event: MemoryAccessEvent,
248 ) -> Result<(), MemoryRepositoryError> {
249 self.record_access(event, AccessKind::Admission).await
250 }
251
252 async fn record_use(&self, event: MemoryAccessEvent) -> Result<(), MemoryRepositoryError> {
253 self.record_access(event, AccessKind::Use).await
254 }
255
256 async fn usage_summary(
257 &self,
258 namespace: &MemoryNamespace,
259 node_id: &str,
260 ) -> Result<MemoryUsageSummary, MemoryRepositoryError> {
261 namespace.validate()?;
262 validate_required_text("nodeId", node_id, MAX_IDENTIFIER_BYTES)?;
263 let state = self.state.read().await;
264 if state
265 .nodes
266 .get(namespace)
267 .and_then(|nodes| nodes.get(node_id))
268 .is_none()
269 {
270 return Err(MemoryRepositoryError::NodeNotFound {
271 node_id: node_id.to_owned(),
272 });
273 }
274 Ok(state
275 .usage
276 .get(namespace)
277 .and_then(|nodes| nodes.get(node_id))
278 .copied()
279 .unwrap_or_default())
280 }
281}
282
283struct PreparedChange {
284 result: MemoryChangeResult,
285 staged: Option<StagedChange>,
286 next_change_sequence: Option<u64>,
287}
288
289fn prepare_change(
290 state: &RepositoryState,
291 change_set: &MemoryChangeSet,
292) -> Result<PreparedChange, MemoryRepositoryError> {
293 change_set.validate_shape()?;
294 if let Some(applied) = state
295 .applied_changes
296 .get(&change_set.namespace)
297 .and_then(|changes| changes.get(&change_set.idempotency_key))
298 {
299 return if applied.change_set == *change_set {
300 Ok(PreparedChange {
301 result: applied.result.clone(),
302 staged: None,
303 next_change_sequence: None,
304 })
305 } else {
306 Err(MemoryRepositoryError::IdempotencyConflict {
307 key: change_set.idempotency_key.clone(),
308 })
309 };
310 }
311
312 let staged = stage_change_set(change_set, state.nodes.get(&change_set.namespace))?;
313 let current_change_sequence = state
314 .namespace_change_sequences
315 .get(&change_set.namespace)
316 .copied()
317 .unwrap_or(0);
318 let next_change_sequence = current_change_sequence
319 .checked_add(1)
320 .ok_or_else(|| MemoryRepositoryError::invariant("namespace change sequence overflowed"))?;
321 let result = MemoryChangeResult {
322 idempotency_key: change_set.idempotency_key.clone(),
323 occurred_at: change_set.occurred_at,
324 nodes: staged
325 .changed
326 .iter()
327 .filter_map(|node_id| staged.nodes.get(node_id).cloned())
328 .collect(),
329 };
330 Ok(PreparedChange {
331 result,
332 staged: Some(staged),
333 next_change_sequence: Some(next_change_sequence),
334 })
335}
336
337fn validate_access(
338 state: &RepositoryState,
339 event: &MemoryAccessEvent,
340 kind: AccessKind,
341) -> Result<bool, MemoryRepositoryError> {
342 let node = state
343 .nodes
344 .get(&event.namespace)
345 .and_then(|nodes| nodes.get(&event.node_id))
346 .ok_or_else(|| MemoryRepositoryError::NodeNotFound {
347 node_id: event.node_id.clone(),
348 })?;
349 let revision_updated_at = node
350 .revision_updated_at(event.node_revision)
351 .ok_or_else(|| MemoryRepositoryError::NodeRevisionNotFound {
352 node_id: event.node_id.clone(),
353 revision: event.node_revision,
354 })?;
355 if matches!(kind, AccessKind::Admission)
356 && (node.revision != event.node_revision || node.status != super::MemoryStatus::Active)
357 {
358 return Err(MemoryRepositoryError::AdmissionNotAllowed {
359 node_id: event.node_id.clone(),
360 revision: event.node_revision,
361 current_revision: node.revision,
362 current_status: node.status.as_str().into(),
363 });
364 }
365 if event.occurred_at < revision_updated_at {
366 return Err(MemoryRepositoryError::invalid(
367 "accessEvent.occurredAt",
368 "must not precede the referenced node revision",
369 ));
370 }
371
372 let existing = match kind {
373 AccessKind::Admission => state
374 .admissions
375 .get(&event.namespace)
376 .and_then(|events| events.get(&event.id)),
377 AccessKind::Use => state
378 .uses
379 .get(&event.namespace)
380 .and_then(|events| events.get(&event.id)),
381 };
382 match existing {
383 Some(existing) if existing == event => Ok(true),
384 Some(_) => Err(MemoryRepositoryError::IdempotencyConflict {
385 key: event.id.clone(),
386 }),
387 None => Ok(false),
388 }
389}