Skip to main content

a3s_memory/repository/
file.rs

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/// Durable local repository backed by a checksummed write-ahead journal.
20///
21/// One live instance owns a repository directory at a time. Mutations are
22/// validated before they are appended and synced; only then are they published
23/// to the in-memory read view. A restart deterministically replays the journal.
24#[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}