1use super::{
2 InMemoryRepository, MemoryAccessEvent, MemoryChangeResult, MemoryChangeSet, MemoryNamespace,
3 MemoryNamespaceChangeToken, MemoryNamespaceSnapshot, MemoryNode, MemoryQuery,
4 MemoryQueryResult, MemoryRepository, MemoryRepositoryError, MemorySnapshotRequest,
5 MemoryUsageSummary,
6};
7use serde::{Deserialize, Serialize};
8use sha2::{Digest, Sha256};
9use std::path::{Path, PathBuf};
10use std::sync::Arc;
11use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
12use tokio::sync::Mutex;
13
14const JOURNAL_FILE: &str = "memory-v2.journal";
15const LOCK_FILE: &str = "memory-v2.lock";
16const JOURNAL_VERSION: u8 = 1;
17const MAX_JOURNAL_RECORD_BYTES: usize = 16 * 1024 * 1024;
18
19#[derive(Debug, Clone)]
25pub struct FileMemoryRepository {
26 state: Arc<FileRepositoryState>,
27}
28
29#[derive(Debug)]
30struct FileRepositoryState {
31 inner: InMemoryRepository,
32 writer: Mutex<JournalWriter>,
33 _lock_file: std::fs::File,
34}
35
36#[derive(Debug)]
37struct JournalWriter {
38 file: tokio::fs::File,
39 poisoned: bool,
40}
41
42#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
43#[serde(tag = "recordType", rename_all = "snake_case")]
44enum JournalRecord {
45 Change { change_set: MemoryChangeSet },
46 Admission { event: MemoryAccessEvent },
47 Use { event: MemoryAccessEvent },
48}
49
50#[derive(Debug, Serialize, Deserialize)]
51#[serde(rename_all = "camelCase")]
52struct JournalEnvelope {
53 version: u8,
54 checksum: String,
55 record: JournalRecord,
56}
57
58impl FileMemoryRepository {
59 pub async fn open(root: impl AsRef<Path>) -> Result<Self, MemoryRepositoryError> {
60 let root = root.as_ref().to_path_buf();
61 tokio::fs::create_dir_all(&root)
62 .await
63 .map_err(|error| persistence("create repository directory", error))?;
64
65 let lock_file = acquire_lock(root.join(LOCK_FILE)).await?;
66 let journal_path = root.join(JOURNAL_FILE);
67 let (records, valid_bytes, file_bytes) = read_journal(&journal_path).await?;
68 if valid_bytes < file_bytes {
69 truncate_torn_tail(&journal_path, valid_bytes).await?;
70 }
71
72 let inner = InMemoryRepository::new();
73 for record in records {
74 replay(&inner, record).await?;
75 }
76
77 let file = tokio::fs::OpenOptions::new()
78 .create(true)
79 .append(true)
80 .open(&journal_path)
81 .await
82 .map_err(|error| persistence("open journal for append", error))?;
83 Ok(Self {
84 state: Arc::new(FileRepositoryState {
85 inner,
86 writer: Mutex::new(JournalWriter {
87 file,
88 poisoned: false,
89 }),
90 _lock_file: lock_file,
91 }),
92 })
93 }
94}
95
96impl FileRepositoryState {
97 async fn persist_access(
98 &self,
99 event: MemoryAccessEvent,
100 access_kind: AccessKind,
101 ) -> Result<(), MemoryRepositoryError> {
102 let mut writer = self.writer.lock().await;
103 writer.ensure_healthy()?;
104 let replayed = match access_kind {
105 AccessKind::Admission => self.inner.preview_admission(&event).await?,
106 AccessKind::Use => self.inner.preview_use(&event).await?,
107 };
108 if replayed {
109 return Ok(());
110 }
111
112 let record = match access_kind {
113 AccessKind::Admission => JournalRecord::Admission {
114 event: event.clone(),
115 },
116 AccessKind::Use => JournalRecord::Use {
117 event: event.clone(),
118 },
119 };
120 writer.append(record).await?;
121 let published = match access_kind {
122 AccessKind::Admission => self.inner.record_admission(event).await,
123 AccessKind::Use => self.inner.record_use(event).await,
124 };
125 if let Err(error) = published {
126 writer.poisoned = true;
127 return Err(MemoryRepositoryError::Persistence {
128 operation: "publish synced access record".into(),
129 message: error.to_string(),
130 });
131 }
132 Ok(())
133 }
134
135 async fn apply_change(
136 &self,
137 change_set: MemoryChangeSet,
138 ) -> Result<MemoryChangeResult, MemoryRepositoryError> {
139 let mut writer = self.writer.lock().await;
140 writer.ensure_healthy()?;
141 let (expected, replayed) = self.inner.preview_apply(&change_set).await?;
142 if replayed {
143 return Ok(expected);
144 }
145
146 writer
147 .append(JournalRecord::Change {
148 change_set: change_set.clone(),
149 })
150 .await?;
151 match self.inner.apply(change_set).await {
152 Ok(actual) if actual == expected => Ok(actual),
153 Ok(actual) => {
154 writer.poisoned = true;
155 Err(MemoryRepositoryError::Persistence {
156 operation: "publish synced change set".into(),
157 message: format!(
158 "preview result diverged from published result: expected {expected:?}, actual {actual:?}"
159 ),
160 })
161 }
162 Err(error) => {
163 writer.poisoned = true;
164 Err(MemoryRepositoryError::Persistence {
165 operation: "publish synced change set".into(),
166 message: error.to_string(),
167 })
168 }
169 }
170 }
171}
172
173#[derive(Debug, Clone, Copy)]
174enum AccessKind {
175 Admission,
176 Use,
177}
178
179#[async_trait::async_trait]
180impl MemoryRepository for FileMemoryRepository {
181 async fn apply(
182 &self,
183 change_set: MemoryChangeSet,
184 ) -> Result<MemoryChangeResult, MemoryRepositoryError> {
185 let state = self.state.clone();
186 join_transaction(
187 tokio::spawn(async move { state.apply_change(change_set).await }),
188 "apply change set",
189 )
190 .await
191 }
192
193 async fn get(
194 &self,
195 namespace: &MemoryNamespace,
196 node_id: &str,
197 ) -> Result<Option<MemoryNode>, MemoryRepositoryError> {
198 self.state.inner.get(namespace, node_id).await
199 }
200
201 async fn query(&self, query: MemoryQuery) -> Result<MemoryQueryResult, MemoryRepositoryError> {
202 self.state.inner.query(query).await
203 }
204
205 async fn snapshot_namespace(
206 &self,
207 request: MemorySnapshotRequest,
208 ) -> Result<MemoryNamespaceSnapshot, MemoryRepositoryError> {
209 self.state.inner.snapshot_namespace(request).await
210 }
211
212 async fn namespace_change_token(
213 &self,
214 namespace: &MemoryNamespace,
215 ) -> Result<Option<MemoryNamespaceChangeToken>, MemoryRepositoryError> {
216 self.state.inner.namespace_change_token(namespace).await
217 }
218
219 async fn record_admission(
220 &self,
221 event: MemoryAccessEvent,
222 ) -> Result<(), MemoryRepositoryError> {
223 let state = self.state.clone();
224 join_transaction(
225 tokio::spawn(async move { state.persist_access(event, AccessKind::Admission).await }),
226 "record admission",
227 )
228 .await
229 }
230
231 async fn record_use(&self, event: MemoryAccessEvent) -> Result<(), MemoryRepositoryError> {
232 let state = self.state.clone();
233 join_transaction(
234 tokio::spawn(async move { state.persist_access(event, AccessKind::Use).await }),
235 "record use",
236 )
237 .await
238 }
239
240 async fn usage_summary(
241 &self,
242 namespace: &MemoryNamespace,
243 node_id: &str,
244 ) -> Result<MemoryUsageSummary, MemoryRepositoryError> {
245 self.state.inner.usage_summary(namespace, node_id).await
246 }
247}
248
249async fn join_transaction<T: Send + 'static>(
250 handle: tokio::task::JoinHandle<Result<T, MemoryRepositoryError>>,
251 operation: &str,
252) -> Result<T, MemoryRepositoryError> {
253 handle
254 .await
255 .map_err(|error| MemoryRepositoryError::Persistence {
256 operation: operation.into(),
257 message: format!("transaction task failed: {error}"),
258 })?
259}
260
261impl JournalWriter {
262 fn ensure_healthy(&self) -> Result<(), MemoryRepositoryError> {
263 if self.poisoned {
264 return Err(MemoryRepositoryError::Persistence {
265 operation: "append journal".into(),
266 message: "writer is poisoned; drop and reopen the repository".into(),
267 });
268 }
269 Ok(())
270 }
271
272 async fn append(&mut self, record: JournalRecord) -> Result<(), MemoryRepositoryError> {
273 self.ensure_healthy()?;
274 let encoded = encode_record(record)?;
275 if encoded.len() > MAX_JOURNAL_RECORD_BYTES {
276 return Err(MemoryRepositoryError::LimitExceeded {
277 resource: "journal record bytes".into(),
278 limit: MAX_JOURNAL_RECORD_BYTES,
279 actual: encoded.len(),
280 });
281 }
282 if let Err(error) = self.file.write_all(&encoded).await {
283 self.poisoned = true;
284 return Err(persistence("write journal record", error));
285 }
286 if let Err(error) = self.file.write_all(b"\n").await {
287 self.poisoned = true;
288 return Err(persistence("write journal delimiter", error));
289 }
290 if let Err(error) = self.file.flush().await {
291 self.poisoned = true;
292 return Err(persistence("flush journal record", error));
293 }
294 if let Err(error) = self.file.sync_data().await {
295 self.poisoned = true;
296 return Err(persistence("sync journal record", error));
297 }
298 Ok(())
299 }
300}
301
302async fn acquire_lock(path: PathBuf) -> Result<std::fs::File, MemoryRepositoryError> {
303 tokio::task::spawn_blocking(move || {
304 let file = std::fs::OpenOptions::new()
305 .create(true)
306 .truncate(false)
307 .read(true)
308 .write(true)
309 .open(path)?;
310 fs2::FileExt::try_lock_exclusive(&file)?;
311 Ok::<_, std::io::Error>(file)
312 })
313 .await
314 .map_err(|error| MemoryRepositoryError::Persistence {
315 operation: "join repository lock task".into(),
316 message: error.to_string(),
317 })?
318 .map_err(|error| persistence("acquire exclusive repository lock", error))
319}
320
321async fn read_journal(
322 path: &Path,
323) -> Result<(Vec<JournalRecord>, u64, u64), MemoryRepositoryError> {
324 let file = match tokio::fs::File::open(path).await {
325 Ok(file) => file,
326 Err(error) if error.kind() == std::io::ErrorKind::NotFound => {
327 return Ok((Vec::new(), 0, 0));
328 }
329 Err(error) => return Err(persistence("open journal for recovery", error)),
330 };
331 let file_bytes = file
332 .metadata()
333 .await
334 .map_err(|error| persistence("read journal metadata", error))?
335 .len();
336 let mut reader = BufReader::new(file);
337 let mut records = Vec::new();
338 let mut valid_bytes = 0_u64;
339 let mut line = Vec::new();
340
341 loop {
342 line.clear();
343 let bytes_read = reader
344 .read_until(b'\n', &mut line)
345 .await
346 .map_err(|error| persistence("read journal record", error))?;
347 if bytes_read == 0 {
348 break;
349 }
350 if line.len() > MAX_JOURNAL_RECORD_BYTES + 1 {
351 return Err(MemoryRepositoryError::LimitExceeded {
352 resource: "journal record bytes".into(),
353 limit: MAX_JOURNAL_RECORD_BYTES,
354 actual: line.len(),
355 });
356 }
357 if line.last() != Some(&b'\n') {
358 break;
359 }
360 line.pop();
361 if line.is_empty() {
362 return Err(MemoryRepositoryError::Persistence {
363 operation: "decode journal record".into(),
364 message: "empty journal record".into(),
365 });
366 }
367 records.push(decode_record(&line)?);
368 valid_bytes += bytes_read as u64;
369 }
370 Ok((records, valid_bytes, file_bytes))
371}
372
373async fn truncate_torn_tail(path: &Path, valid_bytes: u64) -> Result<(), MemoryRepositoryError> {
374 let file = tokio::fs::OpenOptions::new()
375 .write(true)
376 .open(path)
377 .await
378 .map_err(|error| persistence("open torn journal for repair", error))?;
379 file.set_len(valid_bytes)
380 .await
381 .map_err(|error| persistence("truncate torn journal tail", error))?;
382 file.sync_data()
383 .await
384 .map_err(|error| persistence("sync repaired journal", error))
385}
386
387async fn replay(
388 repository: &InMemoryRepository,
389 record: JournalRecord,
390) -> Result<(), MemoryRepositoryError> {
391 let result = match record {
392 JournalRecord::Change { change_set } => repository.apply(change_set).await.map(|_| ()),
393 JournalRecord::Admission { event } => repository.record_admission(event).await,
394 JournalRecord::Use { event } => repository.record_use(event).await,
395 };
396 result.map_err(|error| MemoryRepositoryError::Persistence {
397 operation: "replay journal record".into(),
398 message: error.to_string(),
399 })
400}
401
402fn encode_record(record: JournalRecord) -> Result<Vec<u8>, MemoryRepositoryError> {
403 let checksum = record_checksum(&record)?;
404 serde_json::to_vec(&JournalEnvelope {
405 version: JOURNAL_VERSION,
406 checksum,
407 record,
408 })
409 .map_err(|error| MemoryRepositoryError::Persistence {
410 operation: "encode journal record".into(),
411 message: error.to_string(),
412 })
413}
414
415fn decode_record(bytes: &[u8]) -> Result<JournalRecord, MemoryRepositoryError> {
416 let envelope = serde_json::from_slice::<JournalEnvelope>(bytes).map_err(|error| {
417 MemoryRepositoryError::Persistence {
418 operation: "decode journal record".into(),
419 message: error.to_string(),
420 }
421 })?;
422 if envelope.version != JOURNAL_VERSION {
423 return Err(MemoryRepositoryError::Persistence {
424 operation: "decode journal record".into(),
425 message: format!("unsupported journal version: {}", envelope.version),
426 });
427 }
428 let actual = record_checksum(&envelope.record)?;
429 if actual != envelope.checksum {
430 return Err(MemoryRepositoryError::Persistence {
431 operation: "verify journal checksum".into(),
432 message: format!(
433 "checksum mismatch: expected {}, actual {actual}",
434 envelope.checksum
435 ),
436 });
437 }
438 Ok(envelope.record)
439}
440
441fn record_checksum(record: &JournalRecord) -> Result<String, MemoryRepositoryError> {
442 let bytes = serde_json::to_vec(record).map_err(|error| MemoryRepositoryError::Persistence {
443 operation: "encode journal checksum payload".into(),
444 message: error.to_string(),
445 })?;
446 let digest = Sha256::digest(bytes);
447 Ok(digest.iter().map(|byte| format!("{byte:02x}")).collect())
448}
449
450fn persistence(operation: &str, error: impl std::fmt::Display) -> MemoryRepositoryError {
451 MemoryRepositoryError::Persistence {
452 operation: operation.into(),
453 message: error.to_string(),
454 }
455}